26#ifndef TVM_TIRX_STMT_FUNCTOR_H_
27#define TVM_TIRX_STMT_FUNCTOR_H_
36#include <unordered_map>
46template <
typename FType>
49#define STMT_FUNCTOR_DEFAULT \
51 return VisitStmtDefault_(op, std::forward<Args>(args)...); \
54#define IR_STMT_FUNCTOR_DISPATCH(OP) \
55 vtable.template set_dispatch<OP>([](const ffi::ObjectRef& n, TSelf* self, Args... args) { \
56 return self->VisitStmt_(static_cast<const OP*>(n.get()), std::forward<Args>(args)...); \
59template <
typename R,
typename... Args>
84 static FType vtable = InitVTable();
85 return vtable(
n,
this, std::forward<Args>(args)...);
113 static FType InitVTable() {
138#undef IR_STMT_FUNCTOR_DISPATCH
139#undef STMT_FUNCTOR_DEFAULT
146 using StmtFunctor::operator();
149 using StmtFunctor::VisitStmt;
207 allow_copy_on_write_ =
true;
208 return VisitStmt(stmt);
224 bool allow_copy_on_write_{
false};
234 template <
typename TNode>
236 static_assert(std::is_base_of<StmtNode, TNode>::value,
237 "StmtMutator:: CopyOnWrite requires us to track uniqueness of all parent "
238 "nodes during the recursion. Because the child classes do not necessarily "
239 "check the Array, Expr and other structures during the visit, it is only safe to "
240 "call this function with StmtNodes for now. "
241 "Please create a new node directly in other cases.");
242 if (allow_copy_on_write_) {
244 return ffi::GetObjectPtr<TNode>(
const_cast<TNode*
>(node));
248 return ffi::make_object<TNode>(*node);
258 if (allow_copy_on_write_ && !stmt.unique()) {
259 allow_copy_on_write_ =
false;
260 Stmt ret = StmtFunctor::VisitStmt(stmt);
261 allow_copy_on_write_ =
true;
264 return StmtFunctor::VisitStmt(stmt);
336 using StmtVisitor::operator();
337 using ExprVisitor::operator();
340 using ExprVisitor::VisitExpr;
341 using ExprVisitor::VisitExpr_;
342 using StmtVisitor::VisitStmt;
354 using StmtMutator::operator();
355 using ExprMutator::operator();
358 using ExprMutator::VisitExpr;
359 using ExprMutator::VisitExpr_;
360 using ExprMutator::VisitPrimExpr;
361 using StmtMutator::VisitStmt;
385 ffi::Optional<ffi::Array<ffi::String>>
only_enable = std::nullopt);
394 std::function<
void(
const ffi::ObjectRef&)>
fvisit);
423 std::function<ffi::Optional<Expr>(
const Var& var)>
vmap) {
435 std::function<ffi::Optional<Expr>(
const Var& var)>
vmap) {
450template <
typename Obj>
452 auto func = [&
vmap](
const Var& var) -> ffi::Optional<Expr> {
return vmap.Get(var); };
465template <
typename Obj,
typename Replacement>
467 auto func = [&
vmap](
const Var& var) -> ffi::Optional<Expr> {
483template <
typename Obj,
typename Replacement>
485 auto func = [&
vmap](
const Var& var) -> ffi::Optional<Expr> {
503template <
typename Obj,
typename Replacement,
typename Hasher,
typename EqualityChecker>
505 const std::unordered_map<Var, Replacement, Hasher, EqualityChecker>&
vmap) {
506 auto func = [&
vmap](
const Var& var) -> ffi::Optional<Expr> {
524template <
typename Obj,
typename Replacement>
526 std::unordered_map<const VarNode*, Expr>
vmap;
527 for (
const auto& [iter_var, expr] :
iter_vmap) {
528 vmap[iter_var->var.get()] =
Expr(expr);
531 auto func = [&
vmap](
const Var& var) -> ffi::Optional<Expr> {
552 Stmt stmt, std::function<ffi::Optional<PrimExpr>(
const Var&)>
vmap);
565 PrimExpr expr, std::function<ffi::Optional<PrimExpr>(
const Var&)>
vmap);
575 const std::function<
bool(
const ffi::ObjectRef&)>&
fvisit);
587template <
typename Node,
typename = std::enable_if_t<std::is_base_of_v<StmtNode, Node>>>
591 void VisitStmt(
const Stmt& stmt)
final {
595 StmtVisitor::VisitStmt(stmt);
Managed reference to ExprNode.
Definition base_expr.h:335
A dynamically dispatched functor on the type of the first argument.
Definition node_functor.h:62
Typed reference/view over any Expr whose ExprNode::ty is PrimType.
Definition base_expr.h:401
Range container
Definition expr.h:610
static Range FromMinExtent(PrimExpr min, PrimExpr extent, Span span=Span())
construct a new range with min and extent The corresponding constructor is removed,...
Load a value from an indexed expression source.
Definition expr.h:109
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
Allocate a buffer and declare it in scope.
Definition stmt.h:262
Assert condition, if an error occurs, return the error message.
Definition stmt.h:162
Define certain auxiliary attribute for the body to be a symbolic value. This provide auxiliary inform...
Definition stmt.h:118
Bind a variable to a value in the enclosing scope.
Definition stmt.h:79
A Break in control flow.
Definition stmt.h:724
Representing a region of multi-dimensional buffer access.
Definition buffer_region.h:49
Store value to the high dimension buffer.
Definition stmt.h:204
Checked zero-state view over an ordinary VarNode with BufferType.
Definition buffer.h:179
A Continue in control flow.
Definition stmt.h:749
Declare a buffer that can be used in the body.
Definition stmt.h:237
Evaluates an expression. This is mostly used for putting a Call node into Stmt.
Definition stmt.h:340
ExprMutator that mutates expressions.
Definition expr_functor.h:253
ExprVisitor.
Definition expr_functor.h:207
A for loop, with possible type annotations.
Definition stmt.h:590
IfThenElse statement.
Definition stmt.h:520
A return from the current function.
Definition stmt.h:696
A block is a basic schedule unit in TIR.
Definition stmt.h:833
A block realization node represents execution of the block at the binding values.
Definition stmt.h:903
Standalone statement that declares a scope-id binding (e.g. cta_id, warp_id, lane_id)....
Definition stmt.h:946
The container of seq statement. Represent a sequence of statements.
Definition stmt.h:315
Mutator that recursively mutates stmts and exprs on them.
Definition stmt_functor.h:352
Expr VisitExpr(const Expr &e) override
Visitor to Exprs, can be overriden to do recursive changes to Exprs.
Definition stmt_functor.h:363
Expr VisitExpr_(const BufferRegionNode *op) override
Expr VisitExpr_(const VarNode *op) override
Expr VisitExpr_(const TensorLoadNode *op) override
Visitor that recursively visit stmts and exprs on them.
Definition stmt_functor.h:334
void VisitExpr_(const BufferRegionNode *op) override
void VisitExpr_(const TensorLoadNode *op) override
void VisitExpr(const Expr &e) override
Visitor to Exprs, can be overriden to do recursive changes to Exprs.
Definition stmt_functor.h:344
Definition stmt_functor.h:60
virtual R VisitStmt_(const ScopeIdDefStmtNode *op, Args... args)
Definition stmt_functor.h:104
virtual R VisitStmt_(const ContinueNode *op, Args... args)
Definition stmt_functor.h:95
virtual R VisitStmt_(const BufferStoreNode *op, Args... args)
Definition stmt_functor.h:98
virtual ~StmtFunctor()
virtual destructor
Definition stmt_functor.h:69
virtual R VisitStmtDefault_(const ffi::Object *op, Args...)
Definition stmt_functor.h:106
R operator()(const Stmt &n, Args... args)
Same as call.
Definition stmt_functor.h:76
virtual R VisitStmt_(const IfThenElseNode *op, Args... args)
Definition stmt_functor.h:90
virtual R VisitStmt_(const SBlockRealizeNode *op, Args... args)
Definition stmt_functor.h:103
virtual R VisitStmt_(const ForNode *op, Args... args)
Definition stmt_functor.h:91
virtual R VisitStmt_(const SeqStmtNode *op, Args... args)
Definition stmt_functor.h:100
virtual R VisitStmt_(const SBlockNode *op, Args... args)
Definition stmt_functor.h:102
virtual R VisitStmt_(const tirx::TilePrimitiveCallNode *op, Args... args)
Definition stmt_functor.h:105
virtual R VisitStmt(const Stmt &n, Args... args)
The functor call.
Definition stmt_functor.h:83
virtual R VisitStmt_(const AttrStmtNode *op, Args... args)
Definition stmt_functor.h:89
virtual R VisitStmt_(const EvaluateNode *op, Args... args)
Definition stmt_functor.h:101
virtual R VisitStmt_(const BreakNode *op, Args... args)
Definition stmt_functor.h:94
virtual R VisitStmt_(const ReturnNode *op, Args... args)
Definition stmt_functor.h:93
virtual R VisitStmt_(const DeclBufferNode *op, Args... args)
Definition stmt_functor.h:97
virtual R VisitStmt_(const AssertStmtNode *op, Args... args)
Definition stmt_functor.h:99
virtual R VisitStmt_(const AllocBufferNode *op, Args... args)
Definition stmt_functor.h:96
virtual R VisitStmt_(const WhileNode *op, Args... args)
Definition stmt_functor.h:92
virtual R VisitStmt_(const BindNode *op, Args... args)
Definition stmt_functor.h:88
Same as ExprFunctor except it is applied on statements.
Definition stmt_functor.h:47
StmtMutator that mutates the statements.
Definition stmt_functor.h:196
ffi::ObjectPtr< TNode > CopyOnWrite(const TNode *node)
Perform copy on write on node.
Definition stmt_functor.h:235
virtual Expr VisitExpr(const Expr &e)
Visitor to Exprs, can be overriden to do recursive changes to Exprs.
Definition stmt_functor.h:274
Stmt operator()(Stmt stmt)
Mutate stmt.
Definition stmt_functor.h:206
Stmt VisitStmt_(const ReturnNode *op) override
Stmt VisitStmt_(const tirx::TilePrimitiveCallNode *op) override
Stmt VisitStmt_(const EvaluateNode *op) override
Stmt VisitStmt_(const SBlockNode *op) override
virtual BufferVar VisitBufferUse(const BufferVar &buffer)
Visit buffer at use site (BufferStore, BufferLoad, SBlock reads/writes). By default,...
Stmt VisitStmt_(const WhileNode *op) override
Stmt VisitStmt_(const SeqStmtNode *op) override
virtual BufferVar VisitBufferDef(const BufferVar &buffer, bool alloc_data)
Visit buffer at definition site. Visits shape/strides/elem_offset via VisitExpr. If any field changes...
Stmt VisitStmt_(const ContinueNode *op) override
Stmt VisitStmt_(const SBlockRealizeNode *op) override
Stmt VisitStmt_(const IfThenElseNode *op) override
Stmt VisitStmt_(const ScopeIdDefStmtNode *op) override
Stmt VisitStmt_(const AttrStmtNode *op) override
Stmt VisitStmt(const Stmt &stmt) override
Internal mutator that everyone calls.
Definition stmt_functor.h:257
Stmt VisitStmt_(const BindNode *op) override
ffi::Map< BufferVar, BufferVar > buffer_remap_
Map from old buffer to new buffer, populated by VisitBufferDef.
Definition stmt_functor.h:213
Stmt VisitStmt_(const ForNode *op) override
Stmt VisitSeqStmt_(const SeqStmtNode *op, bool flatten_before_visit, std::function< Stmt(const Stmt &)> fmutate=nullptr)
Alternative advance method for SeqStmtNode.
PrimExpr VisitPrimExpr(const PrimExpr &e)
Mutate a primitive expression and verify that it remains primitive.
Definition stmt_functor.h:276
Stmt VisitStmt_(const BreakNode *op) override
Stmt VisitStmt_(const AllocBufferNode *op) override
Stmt VisitStmt_(const AssertStmtNode *op) override
Stmt VisitStmt_(const BufferStoreNode *op) override
Stmt VisitStmt_(const DeclBufferNode *op) override
StmtVisitor.
Definition stmt_functor.h:144
virtual void VisitBufferUse(const BufferVar &buffer)
Visit buffer at use site (BufferStore, BufferLoad, SBlock reads/writes). By default,...
void VisitStmt_(const ForNode *op) override
void VisitStmt_(const BindNode *op) override
void VisitStmt_(const SBlockNode *op) override
void VisitStmt_(const AttrStmtNode *op) override
void VisitStmt_(const BufferStoreNode *op) override
void VisitStmt_(const ReturnNode *op) override
void VisitStmt_(const SBlockRealizeNode *op) override
void VisitStmt_(const SeqStmtNode *op) override
void VisitStmt_(const AssertStmtNode *op) override
void VisitStmt_(const tirx::TilePrimitiveCallNode *op) override
void VisitStmt_(const ContinueNode *op) override
void VisitStmt_(const WhileNode *op) override
void VisitStmt_(const BreakNode *op) override
void VisitStmt_(const EvaluateNode *op) override
virtual void VisitExpr(const Expr &e)
Visitor to Exprs, can be overriden to do recursive changes to Exprs.
Definition stmt_functor.h:157
void VisitStmt_(const ScopeIdDefStmtNode *op) override
void VisitStmt_(const IfThenElseNode *op) override
virtual void VisitBufferDef(const BufferVar &buffer, bool alloc_data)
Visit buffer at definition site (AllocBuffer, DeclBuffer, SBlock alloc_buffers). Visits buffer shape,...
void VisitStmt_(const AllocBufferNode *op) override
void VisitStmt_(const DeclBufferNode *op) override
Container of all statements.
Definition stmt.h:67
TIRX TilePrimitiveCall stmt.
Definition tile_primitive.h:188
A While loop.
Definition stmt.h:665
void PostOrderVisit(const ffi::ObjectRef &node, std::function< void(const ffi::ObjectRef &)> fvisit)
Recursively visit a statement or expression in post DFS order, applying fvisit. Each node is guarante...
IndexMap Substitute(const IndexMap &index_map, std::function< ffi::Optional< PrimExpr >(const Var &var)> f_subst)
Substitute variables in an index map.
bool ContainsNode(const Stmt &stmt)
Check if the statement contains the specified node type.
Definition stmt_functor.h:588
Stmt IRTransform(Stmt stmt, const ffi::Function &preorder, const ffi::Function &postorder, ffi::Optional< ffi::Array< ffi::String > > only_enable=std::nullopt)
recursively visit the ir nodes in post DFS order, and transform it
void PreOrderVisit(const ffi::ObjectRef &stmt_or_expr, const std::function< bool(const ffi::ObjectRef &)> &fvisit)
Recursively visit a statement or expression in pre DFS order, applying fvisit. If fvisit returns fals...
Stmt SubstituteWithDataTypeLegalization(Stmt stmt, std::function< ffi::Optional< PrimExpr >(const Var &)> vmap)
Substitute the var specified by vmap and legalize data types after substitution.
An object that builds and maintains block scope and StmtSref mapping for Dependence analysis.
Definition analyzer.h:40
Defines the Functor data structures.
#define IR_STMT_FUNCTOR_DISPATCH(OP)
Definition stmt_functor.h:54
#define STMT_FUNCTOR_DEFAULT
Definition stmt_functor.h:49
TIRX tile primitive statements, operators, and reified lambda expressions.
Functors for tirx expressions.