25#ifndef TVM_RELAX_EXPR_FUNCTOR_H_
26#define TVM_RELAX_EXPR_FUNCTOR_H_
35#include <unordered_map>
51template <
typename FType>
55#define EXPR_FUNCTOR_DEFAULT \
57 return VisitExprDefault_(op, std::forward<Args>(args)...); \
60#define EXPR_FUNCTOR_DISABLED \
62 TVM_FFI_THROW(TypeError) << "Relax does not support " << op->GetTypeKey() << " expressions"; \
66#define RELAX_EXPR_FUNCTOR_DISPATCH(OP) \
67 vtable.template set_dispatch<OP>([](const ffi::ObjectRef& n, TSelf* self, Args... args) { \
68 return self->VisitExpr_(static_cast<const OP*>(n.get()), std::forward<Args>(args)...); \
71#define PY_EXPR_VISITOR_DEFAULT(N, PY_FUNC, DEFAULT_FUNC) \
73 if (PY_FUNC != nullptr) \
79#define PY_EXPR_MUTATOR_DEFAULT(N, PY_FUNC, DEFAULT_FUNC, RET_TYPE) \
81 if (PY_FUNC != nullptr) { \
82 RET_TYPE ret = PY_FUNC(N).cast<RET_TYPE>(); \
85 return DEFAULT_FUNC; \
89#define PY_EXPR_VISITOR_DISPATCH(OP, PY_FUNC) \
90 vtable.template set_dispatch<OP>([](const ffi::ObjectRef& n, TSelf* self) { \
91 if (self->PY_FUNC != nullptr) \
94 self->VisitExpr_(static_cast<const OP*>(n.get())); \
97#define PY_EXPR_MUTATOR_DISPATCH(OP, PY_FUNC) \
98 vtable.template set_dispatch<OP>([](const ffi::ObjectRef& n, TSelf* self) { \
99 if (self->PY_FUNC != nullptr) { \
100 Expr expr = self->PY_FUNC(n).cast<Expr>(); \
103 return self->VisitExpr_(static_cast<const OP*>(n.get())); \
107#define PY_EXPR_MUTATOR_VISIT_EXPR_POST_ORDER_DISPATCH(OP) \
108 post_order_vtable.template set_dispatch<OP>([](const ffi::ObjectRef& n, TSelf* self) { \
109 return self->VisitExprPostOrder_(static_cast<const OP*>(n.get())); \
112template <
typename R,
typename... Args>
138 <<
"Found null pointer node while traversing AST. The previous pass may "
139 "have generated invalid data.";
140 static FType vtable = InitVTable();
141 if (vtable.can_dispatch(
n)) {
142 return vtable(
n,
this, std::forward<Args>(args)...);
144 return VisitExprFallback_(
n.
get(), std::forward<Args>(args)...);
200 static FType InitVTable() {
388 class DefaultTypeFieldVisitor :
public TypeVisitor {
390 explicit DefaultTypeFieldVisitor(
ExprVisitor* parent);
393 void VisitTypeExprField(
const Expr& expr)
final;
394 void VisitTypeExprField(
const PrimExpr& expr)
final;
402 DefaultTypeFieldVisitor default_tyfield_visitor_{
this};
521 class DefaultTypeFieldMutator :
public TypeMutator {
526 Expr VisitTypeExprField(
const Expr& expr)
final;
535 DefaultTypeFieldMutator default_tyfield_mutator_{
this};
631 ffi::Optional<ffi::Array<Var>> params = std::nullopt);
663 template <
typename T>
681 std::unordered_map<Var, Var, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
var_remap_;
The utility for constructing Relax binding blocks.
Call corresponds to callable invocation.
Definition expr.h:440
Base type of all the expressions.
Definition base_expr.h:300
Managed reference to ExprNode.
Definition base_expr.h:335
Constant floating point literals in the program.
Definition expr.h:550
Global variable that lives in the top-level module.
Definition expr.h:397
Constant integer literals in the program.
Definition expr.h:487
A dynamically dispatched functor on the type of the first argument.
Definition node_functor.h:62
Primitive Op(builtin intrinsics)
Definition op.h:94
Typed reference/view over any Expr whose ExprNode::ty is PrimType.
Definition base_expr.h:401
Definition source_map.h:111
Load a value from an indexed expression source.
Definition expr.h:109
Get the index-th field out of a tuple.
Definition expr.h:76
Tuple container.
Definition expr.h:48
Type is the base type of all types.
Definition base_expr.h:52
Managed reference to TypeNode.
Definition base_expr.h:77
A local variable in the IR.
Definition expr.h:355
Managed reference to VarNode.
Definition expr.h:372
RAII wrapper function to enter and exit a context object similar to python's with syntax.
Definition with_context.h:59
ContextType * get()
Definition with_context.h:81
a + b
Definition expr.h:124
a && b
Definition expr.h:421
Create a vector where all the elements are value.
Definition vector_expr.h:73
Cast value from one data type to another.
Definition expr.h:80
a / b in the C semnatics.
Definition expr.h:182
a == b
Definition expr.h:313
Floor division, floor(a/b)
Definition expr.h:221
The remainder of the floordiv.
Definition expr.h:239
a >= b
Definition expr.h:403
a > b
Definition expr.h:385
a < b
Definition expr.h:349
Let binding. Bind var to value then evaluate body.
Definition expr.h:537
max(a, b)
Definition expr.h:275
min(a, b)
Definition expr.h:257
a % b in the C semnatics.
Definition expr.h:203
a * b
Definition expr.h:161
a != b
Definition expr.h:331
a || b
Definition expr.h:447
Construct a vector with lanes elements where its i-th element equals base + i * stride....
Definition vector_expr.h:42
return true_value if condition is true, otherwise return false_value.
Definition expr.h:503
Shuffle instruction. vec = concat(vectors) result = (vec[indices[0]], vec[indices[1]] ....
Definition vector_expr.h:105
ffi::String constants, only used in asserts.
Definition expr.h:53
a - b
Definition expr.h:142
Definition block_builder.h:257
static BlockBuilder Create(ffi::Optional< IRModule > ctx_mod)
Create a BlockBuilder.
Constant tensor.
Definition expr.h:93
Represent a data type constant.
Definition expr.h:162
A sub-type of the variable node used to mark dataflow variables from normal visible "function local" ...
Definition expr.h:69
Definition expr_functor.h:113
virtual R VisitExpr_(const TupleNode *op, Args... args)
Definition expr_functor.h:150
virtual R VisitExpr(const Expr &n, Args... args)
The functor call.
Definition expr_functor.h:136
virtual R VisitExpr_(const ShapeExprNode *op, Args... args)
Definition expr_functor.h:153
virtual R VisitExpr_(const TensorLoadNode *op, Args... args)
Definition expr_functor.h:158
virtual R VisitExpr_(const FunctionNode *op, Args... args)
Definition expr_functor.h:156
virtual R VisitExpr_(const ExternFuncNode *op, Args... args)
Definition expr_functor.h:154
virtual R VisitExpr_(const GlobalVarNode *op, Args... args)
Definition expr_functor.h:155
virtual R VisitExpr_(const prim::LetNode *op, Args...) final
Definition expr_functor.h:159
virtual R VisitExpr_(const DataflowVarNode *op, Args... args)
Definition expr_functor.h:152
R operator()(const Expr &n, Args... args)
Same as call.
Definition expr_functor.h:129
virtual R VisitExpr_(const VarNode *op, Args... args)
Definition expr_functor.h:151
virtual R VisitExpr_(const CallNode *op, Args... args)
Definition expr_functor.h:157
virtual R VisitExpr_(const ConstantNode *op, Args... args)
Definition expr_functor.h:149
virtual ~ExprFunctor()
virtual destructor
Definition expr_functor.h:122
A dynamical functor that dispatches on in the first Expr argument. You can use this as a more powerfu...
Definition expr_functor.h:52
A mutator works in unnormalized form.
Definition expr_functor.h:415
Expr VisitExpr_(const TupleNode *op) override
Expr VisitExpr_(const prim::LTNode *op) override
Expr VisitExpr_(const tvm::IntImmNode *op) override
Expr VisitExpr_(const prim::ModNode *op) override
Expr VisitExpr_(const prim::SelectNode *op) override
Expr VisitExpr_(const prim::NotNode *op) override
Expr VisitExpr_(const prim::EQNode *op) override
Expr VisitExpr_(const ConstantNode *op) override
Expr VisitExpr_(const VarNode *op) override
virtual Type VisitExprDepTypeField(const Type &ty)
Visit ty that may recursively contain Expr/PrimExpr.
Expr VisitExpr_(const prim::GENode *op) override
Expr VisitExpr_(const prim::StringImmNode *op) override
Expr VisitExpr_(const prim::SubNode *op) override
Expr VisitExpr_(const prim::FloorModNode *op) override
Expr VisitExpr_(const GlobalVarNode *op) override
Expr VisitExpr_(const TensorLoadNode *op) override
Expr VisitExpr_(const prim::AddNode *op) override
Expr VisitExpr_(const prim::FloorDivNode *op) override
Expr VisitExpr_(const prim::CastNode *op) override
Expr VisitExpr_(const SeqExprNode *op) override
virtual BindingBlock VisitBindingBlock(const BindingBlock &block)
Mutate BindingBlock.
Expr VisitExpr_(const ShapeExprNode *op) override
Expr VisitExpr_(const prim::MaxNode *op) override
Expr VisitExpr_(const prim::GTNode *op) override
Expr VisitExpr_(const ExternFuncNode *op) override
Expr VisitExpr_(const DataTypeImmNode *op) override
Expr VisitExpr_(const tvm::FloatImmNode *op) override
Expr VisitExpr_(const prim::MulNode *op) override
Expr VisitExpr_(const CallNode *op) override
Expr VisitExpr_(const OpNode *op) override
Expr VisitExpr_(const IfNode *op) override
Expr VisitExpr_(const prim::LENode *op) override
Expr VisitExpr_(const StringImmNode *op) override
Expr VisitExpr_(const prim::OrNode *op) override
Expr VisitExprFallback_(const ExprNode *op) override
Expr VisitExpr_(const prim::BroadcastNode *op) override
Expr VisitExpr_(const TupleGetItemNode *op) override
Expr VisitExpr_(const prim::AndNode *op) override
Expr VisitExpr_(const prim::NENode *op) override
Expr VisitExpr_(const prim::DivNode *op) override
Expr VisitExpr(const Expr &expr) override
Expr VisitExpr_(const prim::ShuffleNode *op) override
Expr VisitExpr_(const prim::MinNode *op) override
Expr VisitExpr_(const FunctionNode *op) override
Expr VisitExpr_(const prim::RampNode *op) override
virtual PrimExpr VisitTypePrimExprField(const PrimExpr &expr)
Used to visit the PrimExpr inside of dependent type fields.
bool VisitAndCheckTypeFieldUnchanged(const ffi::ObjectRef &ty)
Check whether VisitExprDepTypeField change ty.
Definition expr_functor.h:504
Expr VisitExpr_(const DataflowVarNode *op) override
A mutator works in normal form.
Definition expr_functor.h:545
std::unordered_map< Var, Var, ffi::ObjectPtrHash, ffi::ObjectPtrEqual > var_remap_
Remap a var to a new var in use-site.
Definition expr_functor.h:681
virtual void VisitBinding_(const VarBindingNode *binding, const DataTypeImmNode *val)
Expr VisitExpr_(const VarNode *op) override
Var WithType(Var var, Type ty)
Create a new var with specified type if the original var's shape or type does not match with the spec...
virtual BindingBlock VisitBindingBlock_(const DataflowBlockNode *block)
void ReEmitBinding(const VarBindingNode *binding, Expr new_value)
Try to remit binding and bind it to a new_value.
virtual void VisitBinding_(const VarBindingNode *binding, const GlobalVarNode *val)
virtual void VisitBinding_(const VarBindingNode *binding, const CallNode *val)
Expr VisitExpr_(const FunctionNode *op) override
Expr VisitWithNewScope(const Expr &body_expr, ffi::Optional< ffi::Array< Var > > params=std::nullopt)
Rewrite the expr with a new scope, used in a Function's body.
virtual void VisitBinding_(const VarBindingNode *binding, const FunctionNode *val)
virtual BindingBlock VisitBindingBlock_(const BindingBlockNode *block)
Expr VisitExpr(const Expr &expr) override
ffi::Optional< Expr > LookupBinding(const Var &var)
Look up the value bound to a variable.
Expr VisitExpr_(const IfNode *op) override
Expr VisitExpr_(const SeqExprNode *op) override
virtual BindingBlock VisitBindingBlock(const BindingBlock &block) override
Generic dispatcher for binding blocks.
virtual void VisitBinding_(const VarBindingNode *binding, const VarNode *val)
virtual void VisitBinding_(const VarBindingNode *binding, const TupleGetItemNode *val)
virtual Var VisitVarDef(const Var &var)
Generic dispatcher for rewriting the var definition site.
virtual void VisitBinding_(const VarBindingNode *binding, const ShapeExprNode *val)
virtual Var VisitVarDef_(const DataflowVarNode *var)
virtual void VisitBinding_(const VarBindingNode *binding, const ExternFuncNode *val)
virtual void VisitBinding_(const VarBindingNode *binding, const ConstantNode *val)
Expr VisitExprPostOrder_(const T *op)
Post-order rewrite a node and normalize.
Definition expr_functor.h:664
virtual void VisitBinding_(const VarBindingNode *binding, const StringImmNode *val)
virtual void VisitBinding_(const VarBindingNode *binding, const DataflowVarNode *val)
virtual void VisitBinding(const Binding &binding)
Generic dispatcher for bindings.
virtual void VisitBinding_(const VarBindingNode *binding, const TupleNode *val)
Expr VisitExpr_(const DataflowVarNode *op) override
Expr VisitWithInnerScope(const Expr &body_expr)
Rewrite the expr with a new scope, used in the branches of If.
virtual void VisitBinding_(const VarBindingNode *binding, const SeqExprNode *val)
virtual void VisitBinding_(const MatchCastNode *binding)
virtual void VisitBinding_(const VarBindingNode *binding, const IfNode *val)
BlockBuilder builder_
Internal block builder to emit bindings during rewriting.
Definition expr_functor.h:678
virtual Var VisitVarDef_(const VarNode *var)
ExprMutator(ffi::Optional< IRModule > mod=std::nullopt)
Definition expr_functor.h:549
virtual void VisitBinding_(const VarBindingNode *binding, const OpNode *val)
virtual void VisitBinding_(const VarBindingNode *binding)
virtual void VisitBinding_(const VarBindingNode *binding, const ExprNode *val)
A simple visitor wrapper around ExprFunctor. Recursively visit the content.
Definition expr_functor.h:255
void VisitExpr_(const prim::BroadcastNode *op) override
virtual void VisitTypePrimExprField(const PrimExpr &expr)
void VisitExpr_(const prim::FloorModNode *op) override
void VisitExpr_(const prim::EQNode *op) override
void VisitExpr_(const prim::GTNode *op) override
void VisitExpr_(const GlobalVarNode *op) override
void VisitExprFallback_(const ExprNode *op) override
void VisitExpr_(const prim::FloorDivNode *op) override
void VisitExpr_(const StringImmNode *op) override
void VisitExpr_(const prim::AddNode *op) override
void VisitExpr_(const TensorLoadNode *op) override
virtual void VisitBinding_(const VarBindingNode *binding, const OpNode *val)
void VisitExpr_(const SeqExprNode *op) override
virtual void VisitBinding_(const VarBindingNode *binding, const TupleNode *val)
void VisitExpr_(const CallNode *op) override
void VisitExpr_(const prim::ModNode *op) override
void VisitExpr_(const IfNode *op) override
virtual void VisitBinding_(const VarBindingNode *binding, const GlobalVarNode *val)
void VisitExpr_(const TupleGetItemNode *op) override
void VisitExpr(const Expr &expr) override
Generic dispatcher for Expr.
void VisitExpr_(const prim::RampNode *op) override
virtual void VisitBindingBlock_(const DataflowBlockNode *block)
virtual void VisitSpan(const Span &span)
void VisitExpr_(const prim::AndNode *op) override
void VisitExpr_(const tvm::FloatImmNode *op) override
virtual void VisitBinding_(const VarBindingNode *binding, const ExprNode *val)
virtual void VisitBinding_(const VarBindingNode *binding, const StringImmNode *val)
virtual void VisitBinding_(const VarBindingNode *binding, const ShapeExprNode *val)
void VisitExpr_(const prim::StringImmNode *op) override
void VisitExpr_(const OpNode *op) override
void VisitExpr_(const prim::MinNode *op) override
virtual void VisitVarDef(const Var &var)
Generic dispatcher for visiting the var definition site.
void VisitExpr_(const tvm::IntImmNode *op) override
virtual void VisitBinding_(const VarBindingNode *binding, const IfNode *val)
virtual void VisitBinding_(const VarBindingNode *binding, const CallNode *val)
virtual void VisitBinding_(const VarBindingNode *binding, const DataflowVarNode *val)
virtual void VisitBinding_(const MatchCastNode *binding)
void VisitExpr_(const prim::NotNode *op) override
void VisitExpr_(const prim::ShuffleNode *op) override
void VisitExpr_(const DataTypeImmNode *op) override
virtual void VisitBinding_(const VarBindingNode *binding, const TupleGetItemNode *val)
void VisitExpr_(const FunctionNode *op) override
void VisitExpr_(const VarNode *op) override
virtual void VisitBinding_(const VarBindingNode *binding, const VarNode *val)
void VisitExpr_(const prim::LTNode *op) override
void VisitExpr_(const prim::NENode *op) override
virtual void VisitBinding_(const VarBindingNode *binding, const ConstantNode *val)
void VisitExpr_(const ExternFuncNode *op) override
virtual void VisitBindingBlock(const BindingBlock &block)
Generic dispatcher for binding blocks.
void VisitExpr_(const prim::MulNode *op) override
void VisitExpr_(const DataflowVarNode *op) override
virtual void VisitVarDef_(const DataflowVarNode *var)
virtual void VisitBinding_(const VarBindingNode *binding, const DataTypeImmNode *val)
void VisitExpr_(const prim::SubNode *op) override
virtual void VisitVarDef_(const VarNode *var)
virtual void VisitBinding_(const VarBindingNode *binding, const SeqExprNode *val)
virtual void VisitExprDepTypeField(const Type &ty)
Visit ty may recursively contain Expr/PrimExpr.
void VisitExpr_(const ShapeExprNode *op) override
void VisitExpr_(const prim::CastNode *op) override
void VisitExpr_(const prim::MaxNode *op) override
void VisitExpr_(const prim::LENode *op) override
void VisitExpr_(const prim::DivNode *op) override
virtual void VisitBinding_(const VarBindingNode *binding, const ExternFuncNode *val)
void VisitExpr_(const TupleNode *op) override
void VisitExpr_(const prim::SelectNode *op) override
virtual void VisitBinding_(const VarBindingNode *binding)
void VisitExpr_(const prim::OrNode *op) override
virtual void VisitBinding(const Binding &binding)
Generic dispatcher for bindings.
void VisitExpr_(const ConstantNode *op) override
virtual void VisitBinding_(const VarBindingNode *binding, const FunctionNode *val)
void VisitExpr_(const prim::GENode *op) override
virtual void VisitBindingBlock_(const BindingBlockNode *block)
The extern function, which can represent packed function.
Definition expr.h:540
Function type information.
Definition type.h:266
A Relax function.
Definition expr.h:447
Condition expression.
Definition expr.h:400
Runtime-match the value to the type.
Definition expr.h:234
A sequence of blocks followed by an expression.
Definition expr.h:336
A shape expression which allows users to construct a shape containing PrimExpr.
Definition expr.h:47
Represent a string literal constant.
Definition expr.h:130
TypeMutator that mutates Relax type nodes.
Definition type_functor.h:136
A type visitor.
Definition type_functor.h:117
void PostOrderVisit(const Expr &node, std::function< void(const Expr &)> fvisit)
An object that builds and maintains block scope and StmtSref mapping for Dependence analysis.
Definition analyzer.h:40
Defines the Functor data structures.
#define EXPR_FUNCTOR_DEFAULT
Definition expr_functor.h:55
#define RELAX_EXPR_FUNCTOR_DISPATCH(OP)
Definition expr_functor.h:66
#define EXPR_FUNCTOR_DISABLED
Definition expr_functor.h:60
Relax types, including the richer dependent Relax type nodes.
a <= b
Definition expr.h:367
Functors and visitors for Relax type nodes.