tvm
Loading...
Searching...
No Matches
expr_functor.h
Go to the documentation of this file.
1/*
2 * Licensed to the Apache Software Foundation (ASF) under one
3 * or more contributor license agreements. See the NOTICE file
4 * distributed with this work for additional information
5 * regarding copyright ownership. The ASF licenses this file
6 * to you under the Apache License, Version 2.0 (the
7 * "License"); you may not use this file except in compliance
8 * with the License. You may obtain a copy of the License at
9 *
10 * http://www.apache.org/licenses/LICENSE-2.0
11 *
12 * Unless required by applicable law or agreed to in writing,
13 * software distributed under the License is distributed on an
14 * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15 * KIND, either express or implied. See the License for the
16 * specific language governing permissions and limitations
17 * under the License.
18 */
19
25#ifndef TVM_RELAX_EXPR_FUNCTOR_H_
26#define TVM_RELAX_EXPR_FUNCTOR_H_
27
28#include <tvm/ir/node_functor.h>
30#include <tvm/relax/expr.h>
31#include <tvm/relax/type.h>
33#include <tvm/tirx/function.h>
34
35#include <unordered_map>
36#include <utility>
37namespace tvm {
38namespace relax {
39
51template <typename FType>
53
54// functions to be overriden.
55#define EXPR_FUNCTOR_DEFAULT \
56 { \
57 return VisitExprDefault_(op, std::forward<Args>(args)...); \
58 }
59
60#define EXPR_FUNCTOR_DISABLED \
61 final { \
62 TVM_FFI_THROW(TypeError) << "Relax does not support " << op->GetTypeKey() << " expressions"; \
63 throw; \
64 }
65
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)...); \
69 });
70
71#define PY_EXPR_VISITOR_DEFAULT(N, PY_FUNC, DEFAULT_FUNC) \
72 { \
73 if (PY_FUNC != nullptr) \
74 PY_FUNC(N); \
75 else \
76 DEFAULT_FUNC; \
77 }
78
79#define PY_EXPR_MUTATOR_DEFAULT(N, PY_FUNC, DEFAULT_FUNC, RET_TYPE) \
80 { \
81 if (PY_FUNC != nullptr) { \
82 RET_TYPE ret = PY_FUNC(N).cast<RET_TYPE>(); \
83 return ret; \
84 } else { \
85 return DEFAULT_FUNC; \
86 } \
87 }
88
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) \
92 self->PY_FUNC(n); \
93 else \
94 self->VisitExpr_(static_cast<const OP*>(n.get())); \
95 });
96
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>(); \
101 return expr; \
102 } else { \
103 return self->VisitExpr_(static_cast<const OP*>(n.get())); \
104 } \
105 });
106
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())); \
110 });
111
112template <typename R, typename... Args>
113class ExprFunctor<R(const Expr& n, Args...)> {
114 private:
115 using TSelf = ExprFunctor<R(const Expr& n, Args...)>;
116 using FType = tvm::NodeFunctor<R(const ffi::ObjectRef& n, TSelf* self, Args...)>;
117
118 public:
120 using result_type = R;
122 virtual ~ExprFunctor() {}
129 R operator()(const Expr& n, Args... args) { return VisitExpr(n, std::forward<Args>(args)...); }
136 virtual R VisitExpr(const Expr& n, Args... args) {
137 TVM_FFI_ICHECK(n.defined())
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)...);
143 }
144 return VisitExprFallback_(n.get(), std::forward<Args>(args)...);
145 }
146 // Functions that can be overriden by subclass
147 // NOTE: cross dialect calls are invoked through global var
148 // We do not expect inline PrimFunc to appear in relax IR.
149 virtual R VisitExpr_(const ConstantNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
150 virtual R VisitExpr_(const TupleNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
151 virtual R VisitExpr_(const VarNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
156 virtual R VisitExpr_(const FunctionNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
157 virtual R VisitExpr_(const CallNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
160 virtual R VisitExpr_(const prim::AddNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
161 virtual R VisitExpr_(const prim::SubNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
162 virtual R VisitExpr_(const prim::MulNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
163 virtual R VisitExpr_(const prim::DivNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
164 virtual R VisitExpr_(const prim::ModNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
165 virtual R VisitExpr_(const prim::FloorDivNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
166 virtual R VisitExpr_(const prim::FloorModNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
167 virtual R VisitExpr_(const prim::MinNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
168 virtual R VisitExpr_(const prim::MaxNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
169 virtual R VisitExpr_(const prim::EQNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
170 virtual R VisitExpr_(const prim::NENode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
171 virtual R VisitExpr_(const prim::LTNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
172 virtual R VisitExpr_(const prim::LENode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
173 virtual R VisitExpr_(const prim::GTNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
174 virtual R VisitExpr_(const prim::GENode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
175 virtual R VisitExpr_(const prim::AndNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
176 virtual R VisitExpr_(const prim::OrNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
177 virtual R VisitExpr_(const prim::CastNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
178 virtual R VisitExpr_(const prim::NotNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
179 virtual R VisitExpr_(const prim::SelectNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
180 virtual R VisitExpr_(const prim::RampNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
181 virtual R VisitExpr_(const prim::BroadcastNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
182 virtual R VisitExpr_(const prim::ShuffleNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
185 virtual R VisitExpr_(const prim::StringImmNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
187 virtual R VisitExpr_(const IfNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
188 virtual R VisitExpr_(const OpNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
190 virtual R VisitExprFallback_(const ExprNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
193 virtual R VisitExprDefault_(const ffi::Object* op, Args...) {
194 TVM_FFI_THROW(InternalError) << "Do not have a default for " << op->GetTypeKey();
195 throw;
196 }
197
198 private:
199 // initialize the vtable.
200 static FType InitVTable() {
201 FType vtable;
202 // Set dispatch
246 vtable.Finalize();
247 return vtable;
248 }
249};
250
255class ExprVisitor : public ExprFunctor<void(const Expr&)> {
256 public:
261 void VisitExpr(const Expr& expr) override;
262 // specific leaf level visitor functions
263 void VisitExpr_(const ConstantNode* op) override;
264 void VisitExpr_(const TupleNode* op) override;
265 void VisitExpr_(const VarNode* op) override;
266 void VisitExpr_(const DataflowVarNode* op) override;
267 void VisitExpr_(const ShapeExprNode* op) override;
268 void VisitExpr_(const ExternFuncNode* op) override;
269 void VisitExpr_(const GlobalVarNode* op) override;
270 void VisitExpr_(const FunctionNode* op) override;
271 void VisitExpr_(const CallNode* op) override;
272 void VisitExpr_(const TensorLoadNode* op) override;
273 void VisitExpr_(const prim::AddNode* op) override;
274 void VisitExpr_(const prim::SubNode* op) override;
275 void VisitExpr_(const prim::MulNode* op) override;
276 void VisitExpr_(const prim::DivNode* op) override;
277 void VisitExpr_(const prim::ModNode* op) override;
278 void VisitExpr_(const prim::FloorDivNode* op) override;
279 void VisitExpr_(const prim::FloorModNode* op) override;
280 void VisitExpr_(const prim::MinNode* op) override;
281 void VisitExpr_(const prim::MaxNode* op) override;
282 void VisitExpr_(const prim::EQNode* op) override;
283 void VisitExpr_(const prim::NENode* op) override;
284 void VisitExpr_(const prim::LTNode* op) override;
285 void VisitExpr_(const prim::LENode* op) override;
286 void VisitExpr_(const prim::GTNode* op) override;
287 void VisitExpr_(const prim::GENode* op) override;
288 void VisitExpr_(const prim::AndNode* op) override;
289 void VisitExpr_(const prim::OrNode* op) override;
290 void VisitExpr_(const prim::CastNode* op) override;
291 void VisitExpr_(const prim::NotNode* op) override;
292 void VisitExpr_(const prim::SelectNode* op) override;
293 void VisitExpr_(const prim::RampNode* op) override;
294 void VisitExpr_(const prim::BroadcastNode* op) override;
295 void VisitExpr_(const prim::ShuffleNode* op) override;
296 void VisitExpr_(const tvm::IntImmNode* op) override;
297 void VisitExpr_(const tvm::FloatImmNode* op) override;
298 void VisitExpr_(const prim::StringImmNode* op) override;
299 void VisitExpr_(const SeqExprNode* op) override;
300 void VisitExpr_(const IfNode* op) override;
301 void VisitExpr_(const OpNode* op) override;
302 void VisitExpr_(const TupleGetItemNode* op) override;
303 void VisitExprFallback_(const ExprNode* op) override;
304 void VisitExpr_(const StringImmNode* op) override;
305 void VisitExpr_(const DataTypeImmNode* op) override;
306
311 virtual void VisitBinding(const Binding& binding);
312 // specific leaf level visitor functions
314 virtual void VisitBinding_(const MatchCastNode* binding);
315 // second level dispatching based on binding value type.
316 // these dispatching functions get called from first-level dispatch on VarBinding
318 virtual void VisitBinding_(const VarBindingNode* binding, const TupleNode* val);
319 virtual void VisitBinding_(const VarBindingNode* binding, const VarNode* val);
325 virtual void VisitBinding_(const VarBindingNode* binding, const CallNode* val);
326 virtual void VisitBinding_(const VarBindingNode* binding, const SeqExprNode* val);
327 virtual void VisitBinding_(const VarBindingNode* binding, const IfNode* val);
328 virtual void VisitBinding_(const VarBindingNode* binding, const OpNode* val);
330 virtual void VisitBinding_(const VarBindingNode* binding, const ExprNode* val);
337 virtual void VisitBindingBlock(const BindingBlock& block);
338 // specific leaf level visitor functions
339 virtual void VisitBindingBlock_(const BindingBlockNode* block);
340 virtual void VisitBindingBlock_(const DataflowBlockNode* block);
341
347 virtual void VisitVarDef(const Var& var);
348
364 virtual void VisitExprDepTypeField(const Type& ty);
365
366 // specific leaf level visitor functions
367 virtual void VisitVarDef_(const VarNode* var);
368 virtual void VisitVarDef_(const DataflowVarNode* var);
369
370 virtual void VisitSpan(const Span& span);
371 virtual void VisitTypePrimExprField(const PrimExpr& expr);
372
373 private:
374 using TSelf = ExprVisitor;
375 using VisitBindingVTable = tvm::NodeFunctor<void(const ffi::ObjectRef& n, ExprVisitor* self,
376 const VarBindingNode* binding)>;
377 // initialize the vtable.
378 static VisitBindingVTable InitVisitBindingVTable();
388 class DefaultTypeFieldVisitor : public TypeVisitor {
389 public:
390 explicit DefaultTypeFieldVisitor(ExprVisitor* parent);
391
392 // Override defaults in type visitor.
393 void VisitTypeExprField(const Expr& expr) final;
394 void VisitTypeExprField(const PrimExpr& expr) final;
395 void VisitType_(const FuncTypeNode* op) final;
396
397 private:
398 ExprVisitor* parent_;
399 };
400 // This visitor is not visible to child classes and only
401 // used to supported default visiting behavior.
402 DefaultTypeFieldVisitor default_tyfield_visitor_{this};
403};
404
405void PostOrderVisit(const Expr& node, std::function<void(const Expr&)> fvisit);
406
415class ExprMutatorBase : public ExprFunctor<Expr(const Expr&)> {
416 public:
417 Expr VisitExpr(const Expr& expr) override;
418 Expr VisitExpr_(const ConstantNode* op) override;
419 Expr VisitExpr_(const TupleNode* op) override;
420 Expr VisitExpr_(const VarNode* op) override;
421 Expr VisitExpr_(const DataflowVarNode* op) override;
422 Expr VisitExpr_(const ShapeExprNode* op) override;
423 Expr VisitExpr_(const ExternFuncNode* op) override;
424 Expr VisitExpr_(const GlobalVarNode* op) override;
425 Expr VisitExpr_(const FunctionNode* op) override;
426 Expr VisitExpr_(const CallNode* op) override;
427 Expr VisitExpr_(const TensorLoadNode* op) override;
428 Expr VisitExpr_(const prim::AddNode* op) override;
429 Expr VisitExpr_(const prim::SubNode* op) override;
430 Expr VisitExpr_(const prim::MulNode* op) override;
431 Expr VisitExpr_(const prim::DivNode* op) override;
432 Expr VisitExpr_(const prim::ModNode* op) override;
433 Expr VisitExpr_(const prim::FloorDivNode* op) override;
434 Expr VisitExpr_(const prim::FloorModNode* op) override;
435 Expr VisitExpr_(const prim::MinNode* op) override;
436 Expr VisitExpr_(const prim::MaxNode* op) override;
437 Expr VisitExpr_(const prim::EQNode* op) override;
438 Expr VisitExpr_(const prim::NENode* op) override;
439 Expr VisitExpr_(const prim::LTNode* op) override;
440 Expr VisitExpr_(const prim::LENode* op) override;
441 Expr VisitExpr_(const prim::GTNode* op) override;
442 Expr VisitExpr_(const prim::GENode* op) override;
443 Expr VisitExpr_(const prim::AndNode* op) override;
444 Expr VisitExpr_(const prim::OrNode* op) override;
445 Expr VisitExpr_(const prim::CastNode* op) override;
446 Expr VisitExpr_(const prim::NotNode* op) override;
447 Expr VisitExpr_(const prim::SelectNode* op) override;
448 Expr VisitExpr_(const prim::RampNode* op) override;
450 Expr VisitExpr_(const prim::ShuffleNode* op) override;
451 Expr VisitExpr_(const tvm::IntImmNode* op) override;
452 Expr VisitExpr_(const tvm::FloatImmNode* op) override;
454 Expr VisitExpr_(const SeqExprNode* op) override;
455 Expr VisitExpr_(const IfNode* op) override;
456 Expr VisitExpr_(const OpNode* op) override;
457 Expr VisitExpr_(const TupleGetItemNode* op) override;
458 Expr VisitExprFallback_(const ExprNode* op) override;
459 Expr VisitExpr_(const StringImmNode* op) override;
460 Expr VisitExpr_(const DataTypeImmNode* op) override;
461
468
476
493 virtual Type VisitExprDepTypeField(const Type& ty);
494
495 protected:
504 bool VisitAndCheckTypeFieldUnchanged(const ffi::ObjectRef& ty) {
505 if (const TypeNode* ty_node = ty.as<TypeNode>()) {
506 Type type = ffi::GetRef<Type>(ty_node);
507 return type.IsMissing() || this->VisitExprDepTypeField(type).same_as(ty);
508 } else {
509 return true;
510 }
511 }
512
513 private:
521 class DefaultTypeFieldMutator : public TypeMutator {
522 public:
523 explicit DefaultTypeFieldMutator(ExprMutatorBase* parent);
524
525 // Override defaults in type visitor.
526 Expr VisitTypeExprField(const Expr& expr) final;
527 PrimExpr VisitTypeExprField(const PrimExpr& expr) final;
528 Type VisitType_(const FuncTypeNode* op) final;
529
530 private:
531 ExprMutatorBase* parent_;
532 };
533 // This visitor is not visible to child classes and only
534 // used to supported default visiting behavior.
535 DefaultTypeFieldMutator default_tyfield_mutator_{this};
536};
537
546 public:
548
549 ExprMutator(ffi::Optional<IRModule> mod = std::nullopt) { builder_ = BlockBuilder::Create(mod); }
550 Expr VisitExpr(const Expr& expr) override;
551 Expr VisitExpr_(const VarNode* op) override;
552 Expr VisitExpr_(const DataflowVarNode* op) override;
553 Expr VisitExpr_(const FunctionNode* op) override;
554 Expr VisitExpr_(const SeqExprNode* op) override;
555 Expr VisitExpr_(const IfNode* op) override;
556
561 virtual void VisitBinding(const Binding& binding);
562 // specific leaf level visitor functions
564 virtual void VisitBinding_(const MatchCastNode* binding);
565 // second level dispatching based on binding value type.
566 // these dispatching functions get called from first-level dispatch on VarBinding
568 virtual void VisitBinding_(const VarBindingNode* binding, const TupleNode* val);
569 virtual void VisitBinding_(const VarBindingNode* binding, const VarNode* val);
575 virtual void VisitBinding_(const VarBindingNode* binding, const CallNode* val);
576 virtual void VisitBinding_(const VarBindingNode* binding, const SeqExprNode* val);
577 virtual void VisitBinding_(const VarBindingNode* binding, const IfNode* val);
578 virtual void VisitBinding_(const VarBindingNode* binding, const OpNode* val);
580 virtual void VisitBinding_(const VarBindingNode* binding, const ExprNode* val);
588 virtual BindingBlock VisitBindingBlock(const BindingBlock& block) override; // NOLINT(*)
589 // specific leaf level visitor functions
592
599 virtual Var VisitVarDef(const Var& var);
600 // specific leaf level visitor functions
601 virtual Var VisitVarDef_(const VarNode* var);
602 virtual Var VisitVarDef_(const DataflowVarNode* var);
603
604 protected:
617
631 ffi::Optional<ffi::Array<Var>> params = std::nullopt);
632
648
655 ffi::Optional<Expr> LookupBinding(const Var& var);
656
663 template <typename T>
665 return builder_->Normalize(ExprMutator::VisitExpr_(op));
666 }
667
676
679
681 std::unordered_map<Var, Var, ffi::ObjectPtrHash, ffi::ObjectPtrEqual> var_remap_;
682
683 private:
684 using TSelf = ExprMutator;
685 using VisitBindingVTable = tvm::NodeFunctor<void(const ffi::ObjectRef& n, ExprMutator* self,
686 const VarBindingNode* binding)>;
687 // initialize the vtable.
688 static VisitBindingVTable InitVisitBindingVTable();
689};
690
691} // namespace relax
692} // namespace tvm
693#endif // TVM_RELAX_EXPR_FUNCTOR_H_
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
bool IsMissing() const
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
Definition expr.h:473
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 expr.h:290
Definition expr.h:307
Definition expr.h:211
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
Definition expr.h:315
A sub-type of the variable node used to mark dataflow variables from normal visible "function local" ...
Definition expr.h:69
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
Definition expr.h:263
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
TIR Function.
Functors and visitors for Relax type nodes.