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_TIR_EXPR_FUNCTOR_H_
26#define TVM_TIR_EXPR_FUNCTOR_H_
27
28#include <tvm/ir/node_functor.h>
29#include <tvm/ir/prim/expr.h>
31
32#include <utility>
33
34namespace tvm {
35namespace tirx {
36
75template <typename FType>
77
78// functions to be overriden.
79#define EXPR_FUNCTOR_DEFAULT \
80 { \
81 return VisitExprDefault_(op, std::forward<Args>(args)...); \
82 }
83
84#define IR_EXPR_FUNCTOR_DISPATCH(OP) \
85 vtable.template set_dispatch<OP>([](const ffi::ObjectRef& n, TSelf* self, Args... args) { \
86 return self->VisitExpr_(static_cast<const OP*>(n.get()), std::forward<Args>(args)...); \
87 });
88
89template <typename R, typename... Args>
90class ExprFunctor<R(const Expr& n, Args...)> {
91 private:
92 using TSelf = ExprFunctor<R(const Expr& n, Args...)>;
93 using FType = NodeFunctor<R(const ffi::ObjectRef& n, TSelf* self, Args...)>;
94
95 public:
97 using result_type = R;
99 virtual ~ExprFunctor() {}
106 R operator()(const Expr& n, Args... args) { return VisitExpr(n, std::forward<Args>(args)...); }
113 virtual R VisitExpr(const Expr& n, Args... args) {
114 static FType vtable = InitVTable();
115 return vtable(n, this, std::forward<Args>(args)...);
116 }
117 // Functions that can be overriden by subclass
118 virtual R VisitExpr_(const VarNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
122 virtual R VisitExpr_(const TupleNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
125 virtual R VisitExpr_(const CallNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
135 virtual R VisitExpr_(const prim::EQNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
136 virtual R VisitExpr_(const prim::NENode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
137 virtual R VisitExpr_(const prim::LTNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
138 virtual R VisitExpr_(const prim::LENode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
139 virtual R VisitExpr_(const prim::GTNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
140 virtual R VisitExpr_(const prim::GENode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
142 virtual R VisitExpr_(const prim::OrNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
149 virtual R VisitExpr_(const IntImmNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
150 virtual R VisitExpr_(const FloatImmNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
152 virtual R VisitExprDefault_(const ffi::Object* op, Args...) {
153 TVM_FFI_THROW(InternalError) << "Do not have a default for " << op->GetTypeKey();
155 }
156
157 private:
158 // initialize the vtable.
159 static FType InitVTable() {
160 FType vtable;
161 // Set dispatch
196 vtable.Finalize();
197 return vtable;
198 }
199};
200
201#undef IR_EXPR_FUNCTOR_DISPATCH
202#undef EXPR_FUNCTOR_DEFAULT
203
207class TVM_DLL ExprVisitor : public ExprFunctor<void(const Expr&)> {
208 public:
209 using ExprFunctor::operator();
210
211 protected:
212 using ExprFunctor::VisitExpr;
213 // list of functions to override.
214 void VisitExpr_(const VarNode* op) override;
215 void VisitExpr_(const TensorLoadNode* op) override;
216 void VisitExpr_(const OpaqueExprNode* op) override;
217 void VisitExpr_(const BufferRegionNode* op) override;
218 void VisitExpr_(const TupleNode* op) override;
219 void VisitExpr_(const TupleGetItemNode* op) override;
220 void VisitExpr_(const prim::LetNode* op) override;
221 void VisitExpr_(const CallNode* op) override;
222 void VisitExpr_(const prim::AddNode* op) override;
223 void VisitExpr_(const prim::SubNode* op) override;
224 void VisitExpr_(const prim::MulNode* op) override;
225 void VisitExpr_(const prim::DivNode* op) override;
226 void VisitExpr_(const prim::ModNode* op) override;
227 void VisitExpr_(const prim::FloorDivNode* op) override;
228 void VisitExpr_(const prim::FloorModNode* op) override;
229 void VisitExpr_(const prim::MinNode* op) override;
230 void VisitExpr_(const prim::MaxNode* op) override;
231 void VisitExpr_(const prim::EQNode* op) override;
232 void VisitExpr_(const prim::NENode* op) override;
233 void VisitExpr_(const prim::LTNode* op) override;
234 void VisitExpr_(const prim::LENode* op) override;
235 void VisitExpr_(const prim::GTNode* op) override;
236 void VisitExpr_(const prim::GENode* op) override;
237 void VisitExpr_(const prim::AndNode* op) override;
238 void VisitExpr_(const prim::OrNode* op) override;
239 void VisitExpr_(const prim::CastNode* op) override;
240 void VisitExpr_(const prim::NotNode* op) override;
241 void VisitExpr_(const prim::SelectNode* op) override;
242 void VisitExpr_(const prim::RampNode* op) override;
243 void VisitExpr_(const prim::BroadcastNode* op) override;
244 void VisitExpr_(const prim::ShuffleNode* op) override;
245 void VisitExpr_(const IntImmNode* op) override;
246 void VisitExpr_(const FloatImmNode* op) override;
247 void VisitExpr_(const prim::StringImmNode* op) override;
248};
249
253class TVM_DLL ExprMutator : protected ExprFunctor<Expr(const Expr&)> {
254 public:
255 using ExprFunctor::operator();
256
257 protected:
258 using ExprFunctor::VisitExpr;
260 PrimExpr VisitPrimExpr(const PrimExpr& expr) { return VisitExpr(expr).as_or_throw<PrimExpr>(); }
261 // list of functions to override.
262 Expr VisitExpr_(const VarNode* op) override;
263 Expr VisitExpr_(const TensorLoadNode* op) override;
264 Expr VisitExpr_(const OpaqueExprNode* op) override;
265 Expr VisitExpr_(const BufferRegionNode* op) override;
266 Expr VisitExpr_(const TupleNode* op) override;
267 Expr VisitExpr_(const TupleGetItemNode* op) override;
268 Expr VisitExpr_(const prim::LetNode* op) override;
269 Expr VisitExpr_(const CallNode* op) override;
270 Expr VisitExpr_(const prim::AddNode* op) override;
271 Expr VisitExpr_(const prim::SubNode* op) override;
272 Expr VisitExpr_(const prim::MulNode* op) override;
273 Expr VisitExpr_(const prim::DivNode* op) override;
274 Expr VisitExpr_(const prim::ModNode* op) override;
275 Expr VisitExpr_(const prim::FloorDivNode* op) override;
276 Expr VisitExpr_(const prim::FloorModNode* op) override;
277 Expr VisitExpr_(const prim::MinNode* op) override;
278 Expr VisitExpr_(const prim::MaxNode* op) override;
279 Expr VisitExpr_(const prim::EQNode* op) override;
280 Expr VisitExpr_(const prim::NENode* op) override;
281 Expr VisitExpr_(const prim::LTNode* op) override;
282 Expr VisitExpr_(const prim::LENode* op) override;
283 Expr VisitExpr_(const prim::GTNode* op) override;
284 Expr VisitExpr_(const prim::GENode* op) override;
285 Expr VisitExpr_(const prim::AndNode* op) override;
286 Expr VisitExpr_(const prim::OrNode* op) override;
287 Expr VisitExpr_(const prim::CastNode* op) override;
288 Expr VisitExpr_(const prim::NotNode* op) override;
289 Expr VisitExpr_(const prim::SelectNode* op) override;
290 Expr VisitExpr_(const prim::RampNode* op) override;
292 Expr VisitExpr_(const prim::ShuffleNode* op) override;
293 Expr VisitExpr_(const IntImmNode* op) override;
294 Expr VisitExpr_(const FloatImmNode* op) override;
296};
297
298} // namespace tirx
299} // namespace tvm
300#endif // TVM_TIR_EXPR_FUNCTOR_H_
Call corresponds to callable invocation.
Definition expr.h:440
Managed reference to ExprNode.
Definition base_expr.h:335
Constant floating point literals in the program.
Definition expr.h:550
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
Base node for opaque construction-time expressions.
Definition base_expr.h:352
Typed reference/view over any Expr whose ExprNode::ty is PrimType.
Definition base_expr.h:401
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
A local variable in the IR.
Definition expr.h:355
RAII wrapper function to enter and exit a context object similar to python's with syntax.
Definition with_context.h:59
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
Representing a region of multi-dimensional buffer access.
Definition buffer_region.h:49
virtual R VisitExpr_(const IntImmNode *op, Args... args)
Definition expr_functor.h:149
virtual R VisitExpr_(const prim::ShuffleNode *op, Args... args)
Definition expr_functor.h:148
virtual R VisitExprDefault_(const ffi::Object *op, Args...)
Definition expr_functor.h:152
virtual R VisitExpr_(const prim::RampNode *op, Args... args)
Definition expr_functor.h:146
virtual R VisitExpr_(const OpaqueExprNode *op, Args... args)
Definition expr_functor.h:120
virtual R VisitExpr_(const prim::FloorModNode *op, Args... args)
Definition expr_functor.h:132
virtual R VisitExpr_(const TupleGetItemNode *op, Args... args)
Definition expr_functor.h:123
virtual R VisitExpr_(const prim::CastNode *op, Args... args)
Definition expr_functor.h:143
virtual R VisitExpr_(const prim::MinNode *op, Args... args)
Definition expr_functor.h:133
virtual R VisitExpr_(const prim::BroadcastNode *op, Args... args)
Definition expr_functor.h:147
virtual R VisitExpr_(const prim::EQNode *op, Args... args)
Definition expr_functor.h:135
virtual R VisitExpr_(const prim::MulNode *op, Args... args)
Definition expr_functor.h:128
R operator()(const Expr &n, Args... args)
Same as call.
Definition expr_functor.h:106
virtual R VisitExpr_(const prim::StringImmNode *op, Args... args)
Definition expr_functor.h:151
virtual R VisitExpr_(const prim::OrNode *op, Args... args)
Definition expr_functor.h:142
virtual R VisitExpr_(const prim::GTNode *op, Args... args)
Definition expr_functor.h:139
virtual R VisitExpr_(const TupleNode *op, Args... args)
Definition expr_functor.h:122
virtual R VisitExpr_(const prim::NotNode *op, Args... args)
Definition expr_functor.h:144
virtual R VisitExpr_(const prim::LetNode *op, Args... args)
Definition expr_functor.h:124
virtual R VisitExpr_(const prim::SelectNode *op, Args... args)
Definition expr_functor.h:145
virtual R VisitExpr_(const TensorLoadNode *op, Args... args)
Definition expr_functor.h:119
virtual R VisitExpr_(const VarNode *op, Args... args)
Definition expr_functor.h:118
virtual R VisitExpr_(const prim::LTNode *op, Args... args)
Definition expr_functor.h:137
virtual R VisitExpr_(const prim::NENode *op, Args... args)
Definition expr_functor.h:136
virtual R VisitExpr_(const prim::ModNode *op, Args... args)
Definition expr_functor.h:130
virtual R VisitExpr_(const prim::AndNode *op, Args... args)
Definition expr_functor.h:141
virtual R VisitExpr(const Expr &n, Args... args)
The functor call.
Definition expr_functor.h:113
virtual R VisitExpr_(const prim::MaxNode *op, Args... args)
Definition expr_functor.h:134
virtual R VisitExpr_(const prim::GENode *op, Args... args)
Definition expr_functor.h:140
virtual R VisitExpr_(const BufferRegionNode *op, Args... args)
Definition expr_functor.h:121
virtual R VisitExpr_(const FloatImmNode *op, Args... args)
Definition expr_functor.h:150
virtual R VisitExpr_(const prim::FloorDivNode *op, Args... args)
Definition expr_functor.h:131
virtual R VisitExpr_(const prim::SubNode *op, Args... args)
Definition expr_functor.h:127
virtual R VisitExpr_(const prim::LENode *op, Args... args)
Definition expr_functor.h:138
virtual R VisitExpr_(const prim::DivNode *op, Args... args)
Definition expr_functor.h:129
virtual R VisitExpr_(const prim::AddNode *op, Args... args)
Definition expr_functor.h:126
virtual ~ExprFunctor()
virtual destructor
Definition expr_functor.h:99
virtual R VisitExpr_(const CallNode *op, Args... args)
Definition expr_functor.h:125
A dynamical functor that dispatches on in the first Expr argument. You can use this as a more powerfu...
Definition expr_functor.h:76
ExprMutator that mutates expressions.
Definition expr_functor.h:253
Expr VisitExpr_(const prim::LTNode *op) override
Expr VisitExpr_(const prim::CastNode *op) override
Expr VisitExpr_(const prim::ShuffleNode *op) override
Expr VisitExpr_(const IntImmNode *op) override
Expr VisitExpr_(const prim::EQNode *op) override
Expr VisitExpr_(const TupleNode *op) override
Expr VisitExpr_(const prim::AddNode *op) override
Expr VisitExpr_(const prim::AndNode *op) override
Expr VisitExpr_(const prim::RampNode *op) override
Expr VisitExpr_(const TensorLoadNode *op) override
Expr VisitExpr_(const prim::NENode *op) override
Expr VisitExpr_(const prim::DivNode *op) override
Expr VisitExpr_(const prim::FloorModNode *op) override
Expr VisitExpr_(const VarNode *op) override
Expr VisitExpr_(const prim::LENode *op) override
Expr VisitExpr_(const prim::SelectNode *op) override
Expr VisitExpr_(const prim::StringImmNode *op) override
Expr VisitExpr_(const prim::MinNode *op) override
Expr VisitExpr_(const prim::BroadcastNode *op) override
Expr VisitExpr_(const prim::GTNode *op) override
Expr VisitExpr_(const prim::OrNode *op) override
Expr VisitExpr_(const prim::GENode *op) override
Expr VisitExpr_(const prim::SubNode *op) override
Expr VisitExpr_(const BufferRegionNode *op) override
Expr VisitExpr_(const prim::LetNode *op) override
Expr VisitExpr_(const OpaqueExprNode *op) override
Expr VisitExpr_(const prim::ModNode *op) override
Expr VisitExpr_(const prim::FloorDivNode *op) override
Expr VisitExpr_(const prim::NotNode *op) override
Expr VisitExpr_(const TupleGetItemNode *op) override
Expr VisitExpr_(const CallNode *op) override
Expr VisitExpr_(const prim::MaxNode *op) override
Expr VisitExpr_(const prim::MulNode *op) override
Expr VisitExpr_(const FloatImmNode *op) override
PrimExpr VisitPrimExpr(const PrimExpr &expr)
Visit a primitive expression and verify that it remains primitive.
Definition expr_functor.h:260
ExprVisitor.
Definition expr_functor.h:207
void VisitExpr_(const prim::FloorDivNode *op) override
void VisitExpr_(const prim::NotNode *op) override
void VisitExpr_(const prim::MaxNode *op) override
void VisitExpr_(const prim::SelectNode *op) override
void VisitExpr_(const prim::FloorModNode *op) override
void VisitExpr_(const prim::LTNode *op) override
void VisitExpr_(const prim::SubNode *op) override
void VisitExpr_(const prim::RampNode *op) override
void VisitExpr_(const prim::GTNode *op) override
void VisitExpr_(const prim::ModNode *op) override
void VisitExpr_(const TupleNode *op) override
void VisitExpr_(const FloatImmNode *op) override
void VisitExpr_(const prim::MulNode *op) override
void VisitExpr_(const prim::DivNode *op) override
void VisitExpr_(const BufferRegionNode *op) override
void VisitExpr_(const prim::LENode *op) override
void VisitExpr_(const prim::NENode *op) override
void VisitExpr_(const prim::CastNode *op) override
void VisitExpr_(const prim::StringImmNode *op) override
void VisitExpr_(const prim::GENode *op) override
void VisitExpr_(const CallNode *op) override
void VisitExpr_(const TensorLoadNode *op) override
void VisitExpr_(const prim::ShuffleNode *op) override
void VisitExpr_(const prim::AddNode *op) override
void VisitExpr_(const IntImmNode *op) override
void VisitExpr_(const TupleGetItemNode *op) override
void VisitExpr_(const prim::BroadcastNode *op) override
void VisitExpr_(const prim::OrNode *op) override
void VisitExpr_(const OpaqueExprNode *op) override
void VisitExpr_(const prim::LetNode *op) override
void VisitExpr_(const prim::AndNode *op) override
void VisitExpr_(const prim::MinNode *op) override
void VisitExpr_(const VarNode *op) override
void VisitExpr_(const prim::EQNode *op) override
TIR expressions.
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
a <= b
Definition expr.h:367
#define IR_EXPR_FUNCTOR_DISPATCH(OP)
Definition expr_functor.h:84