summaryrefslogtreecommitdiff
path: root/source/slang/slang-ast-expr.h
diff options
context:
space:
mode:
authorSai Praveen Bangaru <31557731+saipraveenb25@users.noreply.github.com>2022-10-20 14:22:00 -0400
committerGitHub <noreply@github.com>2022-10-20 11:22:00 -0700
commit1093218d6f0e114eb9fa52d60ca525bf9dd9f98a (patch)
treee85158637680f783caaf7f4433a6844398cd8f7b /source/slang/slang-ast-expr.h
parent576c8407e60143682cd40c68101c6eae8563ca3d (diff)
Modified the new type system to support generic differentiable types … (#2413)
* Modified the new type system to support generic differentiable types and added support for differentiating overloaded functions. * Changed a few asserts to release asserts to avoid unreferenced variable errors * Fixed a naming issue with TypeWitnessBreadcumb::Flavor::Decl * Added logic to avoid tracking differentiable types if the module does not use auto-diff or define differentiable types. * Moved the auto-diff passes to after the specialization step, added a more complex generics test * Added a generics stress test and fixed AST-side logic. IR side needs some more work * Added differential getter and setter logic, fixed multiple issues with DifferentiableTypeDictionary, added support for loops and conditions * Changed differential getters to use pointer types, added getter type checking * Fixed some bugs related to diff type registration and differential getters * Removed some superfluous code * Removed some more unused code. * Fixed an issue with witness substitution * Minor fix Co-authored-by: Yong He <yonghe@outlook.com>
Diffstat (limited to 'source/slang/slang-ast-expr.h')
-rw-r--r--source/slang/slang-ast-expr.h24
1 files changed, 22 insertions, 2 deletions
diff --git a/source/slang/slang-ast-expr.h b/source/slang/slang-ast-expr.h
index 13d687da0..e0a55cc29 100644
--- a/source/slang/slang-ast-expr.h
+++ b/source/slang/slang-ast-expr.h
@@ -38,6 +38,18 @@ class VarExpr : public DeclRefExpr
SLANG_AST_CLASS(VarExpr)
};
+class DifferentiableDeclRefExpr : public Expr
+{
+ SLANG_AST_CLASS(DifferentiableDeclRefExpr)
+
+ // Inner decl ref expr that references a differentiable expression.
+ Expr* inner = nullptr;
+
+ // Information on getters and setters if available.
+ Expr* setterExpr = nullptr;
+ Expr* getterExpr = nullptr;
+};
+
// An expression that references an overloaded set of declarations
// having the same name.
class OverloadedExpr : public Expr
@@ -428,13 +440,21 @@ class OpenRefExpr : public Expr
Expr* innerExpr = nullptr;
};
+ /// Base class for higher-order function application
+ /// Eg: foo(fn) where fn is a function expression.
+ ///
+class HigherOrderInvokeExpr : public Expr
+{
+ SLANG_ABSTRACT_AST_CLASS(HigherOrderInvokeExpr)
+ Expr* baseFunction;
+};
+
/// An expression of the form `__jvp(fn)` to access the
/// forward-mode derivative version of the function `fn`
///
-class JVPDifferentiateExpr: public Expr
+class JVPDifferentiateExpr: public HigherOrderInvokeExpr
{
SLANG_AST_CLASS(JVPDifferentiateExpr)
- Expr* baseFunction;
};
/// A type expression of the form `__TaggedUnion(A, ...)`.