Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion docs/arch/codegen.rst
Original file line number Diff line number Diff line change
Expand Up @@ -144,7 +144,7 @@ backend (x86, ARM, NVPTX, AMDGPU, etc.).
├── CodeGenNVPTX ← NVIDIA PTX via LLVM (target.build.nvptx)
└── CodeGenAMDGPU ← AMD GPU via LLVM (target.build.rocm)

``CodeGenLLVM`` inherits from both ``ExprFunctor<llvm::Value*(const PrimExpr&)>`` and
``CodeGenLLVM`` inherits from both ``ExprFunctor<llvm::Value*(const Expr&)>`` and
``StmtFunctor<void(const Stmt&)>``. Each TIR node type has a corresponding visitor:

- **Expressions** (``VisitExpr_``) convert TIR expressions to LLVM ``Value``\ s:
Expand Down
80 changes: 40 additions & 40 deletions include/tvm/tirx/expr_functor.h
Original file line number Diff line number Diff line change
Expand Up @@ -86,9 +86,9 @@ class ExprFunctor;
});

template <typename R, typename... Args>
class ExprFunctor<R(const PrimExpr& n, Args...)> {
class ExprFunctor<R(const Expr& n, Args...)> {
private:
using TSelf = ExprFunctor<R(const PrimExpr& n, Args...)>;
using TSelf = ExprFunctor<R(const Expr& n, Args...)>;
using FType = NodeFunctor<R(const ffi::ObjectRef& n, TSelf* self, Args...)>;

public:
Expand All @@ -102,16 +102,14 @@ class ExprFunctor<R(const PrimExpr& n, Args...)> {
* \param args Additional arguments.
* \return The result of the call
*/
R operator()(const PrimExpr& n, Args... args) {
return VisitExpr(n, std::forward<Args>(args)...);
}
R operator()(const Expr& n, Args... args) { return VisitExpr(n, std::forward<Args>(args)...); }
/*!
* \brief The functor call.
* \param n The expression node.
* \param args Additional arguments.
* \return The result of the call
*/
virtual R VisitExpr(const PrimExpr& n, Args... args) {
virtual R VisitExpr(const Expr& n, Args... args) {
static FType vtable = InitVTable();
return vtable(n, this, std::forward<Args>(args)...);
}
Expand Down Expand Up @@ -201,7 +199,7 @@ class ExprFunctor<R(const PrimExpr& n, Args...)> {
/*!
* \brief ExprVisitor
*/
class TVM_DLL ExprVisitor : public ExprFunctor<void(const PrimExpr&)> {
class TVM_DLL ExprVisitor : public ExprFunctor<void(const Expr&)> {
public:
using ExprFunctor::operator();

Expand Down Expand Up @@ -245,45 +243,47 @@ class TVM_DLL ExprVisitor : public ExprFunctor<void(const PrimExpr&)> {
/*!
* \brief ExprMutator that mutates expressions.
*/
class TVM_DLL ExprMutator : protected ExprFunctor<PrimExpr(const PrimExpr&)> {
class TVM_DLL ExprMutator : protected ExprFunctor<Expr(const Expr&)> {
public:
using ExprFunctor::operator();

protected:
using ExprFunctor::VisitExpr;
/*! \brief Visit a primitive expression and verify that it remains primitive. */
PrimExpr VisitPrimExpr(const PrimExpr& expr) { return VisitExpr(expr).as_or_throw<PrimExpr>(); }
// list of functions to override.
PrimExpr VisitExpr_(const VarNode* op) override;
PrimExpr VisitExpr_(const BufferLoadNode* op) override;
PrimExpr VisitExpr_(const ProducerLoadNode* op) override;
PrimExpr VisitExpr_(const LetNode* op) override;
PrimExpr VisitExpr_(const CallNode* op) override;
PrimExpr VisitExpr_(const AddNode* op) override;
PrimExpr VisitExpr_(const SubNode* op) override;
PrimExpr VisitExpr_(const MulNode* op) override;
PrimExpr VisitExpr_(const DivNode* op) override;
PrimExpr VisitExpr_(const ModNode* op) override;
PrimExpr VisitExpr_(const FloorDivNode* op) override;
PrimExpr VisitExpr_(const FloorModNode* op) override;
PrimExpr VisitExpr_(const MinNode* op) override;
PrimExpr VisitExpr_(const MaxNode* op) override;
PrimExpr VisitExpr_(const EQNode* op) override;
PrimExpr VisitExpr_(const NENode* op) override;
PrimExpr VisitExpr_(const LTNode* op) override;
PrimExpr VisitExpr_(const LENode* op) override;
PrimExpr VisitExpr_(const GTNode* op) override;
PrimExpr VisitExpr_(const GENode* op) override;
PrimExpr VisitExpr_(const AndNode* op) override;
PrimExpr VisitExpr_(const OrNode* op) override;
PrimExpr VisitExpr_(const ReduceNode* op) override;
PrimExpr VisitExpr_(const CastNode* op) override;
PrimExpr VisitExpr_(const NotNode* op) override;
PrimExpr VisitExpr_(const SelectNode* op) override;
PrimExpr VisitExpr_(const RampNode* op) override;
PrimExpr VisitExpr_(const BroadcastNode* op) override;
PrimExpr VisitExpr_(const ShuffleNode* op) override;
PrimExpr VisitExpr_(const IntImmNode* op) override;
PrimExpr VisitExpr_(const FloatImmNode* op) override;
PrimExpr VisitExpr_(const StringImmNode* op) override;
Expr VisitExpr_(const VarNode* op) override;
Expr VisitExpr_(const BufferLoadNode* op) override;
Expr VisitExpr_(const ProducerLoadNode* op) override;
Expr VisitExpr_(const LetNode* op) override;
Expr VisitExpr_(const CallNode* op) override;
Expr VisitExpr_(const AddNode* op) override;
Expr VisitExpr_(const SubNode* op) override;
Expr VisitExpr_(const MulNode* op) override;
Expr VisitExpr_(const DivNode* op) override;
Expr VisitExpr_(const ModNode* op) override;
Expr VisitExpr_(const FloorDivNode* op) override;
Expr VisitExpr_(const FloorModNode* op) override;
Expr VisitExpr_(const MinNode* op) override;
Expr VisitExpr_(const MaxNode* op) override;
Expr VisitExpr_(const EQNode* op) override;
Expr VisitExpr_(const NENode* op) override;
Expr VisitExpr_(const LTNode* op) override;
Expr VisitExpr_(const LENode* op) override;
Expr VisitExpr_(const GTNode* op) override;
Expr VisitExpr_(const GENode* op) override;
Expr VisitExpr_(const AndNode* op) override;
Expr VisitExpr_(const OrNode* op) override;
Expr VisitExpr_(const ReduceNode* op) override;
Expr VisitExpr_(const CastNode* op) override;
Expr VisitExpr_(const NotNode* op) override;
Expr VisitExpr_(const SelectNode* op) override;
Expr VisitExpr_(const RampNode* op) override;
Expr VisitExpr_(const BroadcastNode* op) override;
Expr VisitExpr_(const ShuffleNode* op) override;
Expr VisitExpr_(const IntImmNode* op) override;
Expr VisitExpr_(const FloatImmNode* op) override;
Expr VisitExpr_(const StringImmNode* op) override;
};

} // namespace tirx
Expand Down
21 changes: 12 additions & 9 deletions include/tvm/tirx/stmt_functor.h
Original file line number Diff line number Diff line change
Expand Up @@ -152,7 +152,7 @@ class TVM_DLL StmtVisitor : protected StmtFunctor<void(const Stmt&)> {
* or have a class sub-class both StmtVisitor and ExprVisitor
* and redirect Visit to ExprMutator::VisitExpr(Expr)
*/
virtual void VisitExpr(const PrimExpr& e) {}
virtual void VisitExpr(const Expr& e) {}
/*!
* \brief Visit buffer at definition site (AllocBuffer, DeclBuffer, SBlock alloc_buffers).
* Visits buffer shape, strides, elem_offset via VisitExpr.
Expand Down Expand Up @@ -268,7 +268,9 @@ class TVM_DLL StmtMutator : protected StmtFunctor<Stmt(const Stmt&)> {
* or have a class sub-class both StmtMutator and ExprMutator
* and redirect Mutate to ExprMutator::Mutate(Expr)
*/
virtual PrimExpr VisitExpr(const PrimExpr& e) { return e; }
virtual Expr VisitExpr(const Expr& e) { return e; }
/*! \brief Mutate a primitive expression and verify that it remains primitive. */
PrimExpr VisitPrimExpr(const PrimExpr& e) { return VisitExpr(e).as_or_throw<PrimExpr>(); }
/*!
* \brief Visit buffer at definition site. Visits shape/strides/elem_offset via VisitExpr.
* If any field changes, creates a new buffer and records it in buffer_remap_.
Expand Down Expand Up @@ -335,7 +337,7 @@ class TVM_DLL StmtExprVisitor : public ExprVisitor, public StmtVisitor {
using ExprVisitor::VisitExpr_;
using StmtVisitor::VisitStmt;

void VisitExpr(const PrimExpr& e) override { return ExprVisitor::VisitExpr(e); }
void VisitExpr(const Expr& e) override { return ExprVisitor::VisitExpr(e); }
void VisitExpr_(const BufferLoadNode* op) override;
};

Expand All @@ -350,10 +352,11 @@ class TVM_DLL StmtExprMutator : public ExprMutator, public StmtMutator {
protected:
using ExprMutator::VisitExpr;
using ExprMutator::VisitExpr_;
using ExprMutator::VisitPrimExpr;
using StmtMutator::VisitStmt;

PrimExpr VisitExpr(const PrimExpr& e) override { return ExprMutator::VisitExpr(e); }
PrimExpr VisitExpr_(const BufferLoadNode* op) override;
Expr VisitExpr(const Expr& e) override { return ExprMutator::VisitExpr(e); }
Expr VisitExpr_(const BufferLoadNode* op) override;
};

/*!
Expand All @@ -375,9 +378,9 @@ TVM_DLL Stmt IRTransform(Stmt stmt, const ffi::Function& preorder, const ffi::Fu
ffi::Optional<ffi::Array<ffi::String>> only_enable = std::nullopt);

/*!
* \brief Recursively visit the ir in post DFS order node, apply fvisit
* \brief Recursively visit a statement or expression in post DFS order, applying fvisit.
* Each node is guaranteed to be visited only once.
* \param node The ir to be visited.
* \param node The statement or expression to be visited.
* \param fvisit The visitor function to be applied.
*/
TVM_DLL void PostOrderVisit(const ffi::ObjectRef& node,
Expand Down Expand Up @@ -560,9 +563,9 @@ TVM_DLL PrimExpr SubstituteWithDataTypeLegalization(
PrimExpr expr, std::function<ffi::Optional<PrimExpr>(const Var&)> vmap);

/*!
* \brief Recursively visit the IR in pre DFS order node, apply fvisit.
* \brief Recursively visit a statement or expression in pre DFS order, applying fvisit.
* If fvisit returns false, it won't visit the children of the node.
* \param stmt_or_expr The ir to be visited.
* \param stmt_or_expr The statement or expression to be visited.
* \param fvisit The visitor function to be applied. If fvisit returns false, it won't visit the
* children of the node
*/
Expand Down
18 changes: 12 additions & 6 deletions python/tvm/tirx/stmt_functor.py
Original file line number Diff line number Diff line change
Expand Up @@ -994,28 +994,34 @@ def ir_transform(stmt, preorder, postorder, only_enable=None):
return _ffi_api.IRTransform(stmt, preorder, postorder, only_enable) # type: ignore


def post_order_visit(stmt, fvisit):
"""Recursively visit the ir in post DFS order node, apply fvisit
def post_order_visit(node, fvisit):
"""Recursively visit a statement or expression in post DFS order, applying fvisit.
Each node is guaranteed to be visited only once.

Parameters
----------
node : tvm.tirx.Stmt or tvm.ir.Expr
The statement or expression to visit.

fvisit: function
The visitor function.
"""
return _ffi_api.PostOrderVisit(stmt, fvisit) # type: ignore
return _ffi_api.PostOrderVisit(node, fvisit) # type: ignore


def pre_order_visit(stmt, fvisit):
"""Recursive pre-order visit on stmt AST, applying fvisit on each node.
def pre_order_visit(node, fvisit):
"""Recursively visit a statement or expression in pre-order, applying fvisit.
If fvisit returns False, it won't visit the children of the node.

Parameters
----------
node : tvm.tirx.Stmt or tvm.ir.Expr
The statement or expression to visit.

fvisit: function of the signature Object -> bool
The visitor function.
"""
return _ffi_api.PreOrderVisit(stmt, fvisit) # type: ignore
return _ffi_api.PreOrderVisit(node, fvisit) # type: ignore


def substitute(node, vmap):
Expand Down
8 changes: 4 additions & 4 deletions src/arith/bound_deducer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ class VariablePathFinder : public ExprVisitor {
public:
explicit VariablePathFinder(PrimExpr target) : target_(target) {}

void VisitExpr(const PrimExpr& node) final {
void VisitExpr(const Expr& node) final {
if (visited_.count(node.get()) != 0) return;
visited_.insert(node.get());

Expand Down Expand Up @@ -72,7 +72,7 @@ std::vector<const ffi::Object*> GetPath(PrimExpr target, PrimExpr expr) {
enum CompareOp { kGreater, kLess, kEqual };

// a visitor to deduce the bound of a variable from a expression
class BoundDeducer : public ExprFunctor<void(const PrimExpr&)> {
class BoundDeducer : public ExprFunctor<void(const Expr&)> {
public:
friend class BoundDeduceInputChecker;
friend class Converter;
Expand All @@ -83,7 +83,7 @@ class BoundDeducer : public ExprFunctor<void(const PrimExpr&)> {

void Deduce();

void VisitExpr(const PrimExpr& e) final {
void VisitExpr(const Expr& e) final {
if (!success_) return;
if (iter_ < path_.size() && e.get() == path_[iter_++]) {
ExprFunctor::VisitExpr(e);
Expand Down Expand Up @@ -239,7 +239,7 @@ class BoundDeduceInputChecker : public ExprVisitor {
return target_count == 1;
}

void VisitExpr(const PrimExpr& e) final {
void VisitExpr(const Expr& e) final {
if (e.same_as(deducer_->target_)) ++target_count;
ExprVisitor::VisitExpr(e);
}
Expand Down
Loading
Loading