This is an automated email from the ASF dual-hosted git repository.
tqchen pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/main by this push:
new e30f85965b [REFACTOR][TIR] Inline StructuralWalk at variable-use
checks (#20306)
e30f85965b is described below
commit e30f85965bec1eb57de0183214d42b59b16b0a48
Author: Tianqi Chen <[email protected]>
AuthorDate: Thu Sep 10 13:30:31 2026 -0400
[REFACTOR][TIR] Inline StructuralWalk at variable-use checks (#20306)
Remove the public `UsesVar` helper and inline preorder
`ffi::StructuralWalk` at each variable-use check.
Each matching callback returns the encountered `Var` in `VisitInterrupt`
and uses the walk result as the signal, preserving reflected traversal
and immediate early exit. Adjacent `SideEffect` behavior remains
unchanged.
Validation: LLVM 23 CPU build; full C++ suite (133/133); clang-format.
---
include/tvm/tirx/analysis.h | 16 -----
src/arith/detect_linear_equation.cc | 12 +++-
src/arith/int_set.cc | 10 ++-
src/arith/ir_mutator_with_analyzer.h | 7 +-
src/arith/iter_affine_map.cc | 28 ++++++--
src/relax/analysis/tir_op_pattern_kind.cc | 30 ++++----
src/relax/transform/rewrite_dataflow_reshape.cc | 8 ++-
.../postproc/rewrite_reduction_block.cc | 7 +-
src/s_tir/schedule/analysis/analysis.cc | 14 ++--
src/s_tir/schedule/analysis/reducer.cc | 9 ++-
src/s_tir/schedule/primitive/blockize_tensorize.cc | 28 +++++---
src/s_tir/schedule/primitive/for_kind.cc | 6 +-
.../schedule/primitive/loop_transformation.cc | 32 ++++-----
src/s_tir/schedule/primitive/reduction.cc | 28 ++++++--
src/s_tir/transform/compact_buffer_region.cc | 7 +-
src/s_tir/transform/hoist_expression.cc | 32 +++++----
src/s_tir/transform/loop_partition.cc | 16 ++++-
src/s_tir/transform/thread_storage_sync.cc | 13 +++-
src/tirx/analysis/var_touch.cc | 80 ----------------------
src/tirx/script/printer/for_loop.cc | 11 ++-
src/tirx/transform/ir_utils.cc | 7 +-
src/tirx/transform/lower_warp_memory.cc | 8 ++-
22 files changed, 221 insertions(+), 188 deletions(-)
diff --git a/include/tvm/tirx/analysis.h b/include/tvm/tirx/analysis.h
index 6323c03f96..eb90ca5226 100644
--- a/include/tvm/tirx/analysis.h
+++ b/include/tvm/tirx/analysis.h
@@ -106,22 +106,6 @@ TVM_DLL ffi::Array<Var> UndefinedVars(const PrimExpr&
expr, const ffi::Array<Var
*/
TVM_DLL CallEffectKind SideEffect(const PrimExpr& expr);
-/*!
- * \brief Whether the given Stmt uses any var in the given variable set.
- * \param stmt The Stmt to be checked.
- * \param vset_contains The check function to see if a var is in the variable
set.
- * \return Whether `stmt` uses any var in the given variable set.
- */
-TVM_DLL bool UsesVar(const Stmt& stmt, std::function<bool(const VarNode*)>
vset_contains);
-
-/*!
- * \brief Whether the given PrimExpr uses any var in the given variable set.
- * \param expr The PrimExpr to be checked.
- * \param vset_contains The check function to see if var is in the variable
set.
- * \return Whether `expr` uses any var in the given variable set.
- */
-TVM_DLL bool UsesVar(const PrimExpr& expr, std::function<bool(const VarNode*)>
vset_contains);
-
/*!
* \brief Verifies whether the IR stmt or Expr is in SSA form.
* That is: each Var is defined and assigned once(in Let/For)
diff --git a/src/arith/detect_linear_equation.cc
b/src/arith/detect_linear_equation.cc
index f00bb889e6..a2a8ffdd5d 100644
--- a/src/arith/detect_linear_equation.cc
+++ b/src/arith/detect_linear_equation.cc
@@ -110,7 +110,11 @@ class LinearEqDetector : public
ExprFunctor<LinearEqEntry(const Expr&, const Pri
}
LinearEqEntry VisitExprDefault_(const ffi::Object* op, const PrimExpr& e)
final {
if (fail_) return LinearEqEntry();
- if (UsesVar(e, [this](const VarNode* var) { return var == var_.get(); })) {
+ auto walkfn = [this](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return var.get() == var_.get() ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(e, walkfn).has_value())
{
fail_ = true;
return LinearEqEntry();
} else {
@@ -157,11 +161,15 @@ ffi::Array<PrimExpr> DetectLinearEquation(const PrimExpr&
e, const ffi::Array<Pr
std::unordered_set<const VarNode*> vset;
auto vset_contains = [&](const VarNode* node) { return vset.count(node) !=
0; };
+ auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return vset_contains(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
for (size_t i = vars.size(); i > 1; --i) {
vset.insert(vars[i - 1].get());
// The previous coeff contains the variable
- if (UsesVar(coeff[i - 2], vset_contains)) {
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(coeff[i - 2],
walkfn).has_value()) {
return ffi::Array<PrimExpr>();
}
}
diff --git a/src/arith/int_set.cc b/src/arith/int_set.cc
index 7c63eb908f..9063c0b620 100644
--- a/src/arith/int_set.cc
+++ b/src/arith/int_set.cc
@@ -24,6 +24,7 @@
#include <tvm/arith/int_set.h>
#include <tvm/arith/iter_affine_map.h>
#include <tvm/ffi/cast.h>
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/function.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/prim/expr.h>
@@ -599,10 +600,13 @@ class IntervalSetEvaluator : public
ExprFunctor<IntervalSet(const Expr&)> {
}
// If the indices do not contain any variables to be relaxed, return the
TensorLoad itself.
// Otherwise return `IntervalSet::everything()` since we have no knowledge
on the buffer data.
+ auto walkfn = [dom_map = &this->dom_map_](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return dom_map->find(var) != dom_map->end()
+ ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
for (const PrimExpr& index : op->indices) {
- if (UsesVar(index, [dom_map = &this->dom_map_](const VarNode* var) {
- return dom_map->find(ffi::GetRef<Var>(var)) != dom_map->end();
- })) {
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(index,
walkfn).has_value()) {
return IntervalSet::Everything();
}
}
diff --git a/src/arith/ir_mutator_with_analyzer.h
b/src/arith/ir_mutator_with_analyzer.h
index da331ada86..d05e8e9a5f 100644
--- a/src/arith/ir_mutator_with_analyzer.h
+++ b/src/arith/ir_mutator_with_analyzer.h
@@ -26,6 +26,7 @@
#include <tvm/arith/analyzer.h>
#include <tvm/ffi/cast.h>
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ir/scope_stack.h>
#include <tvm/ir/with_context.h>
#include <tvm/tirx/analysis.h>
@@ -108,8 +109,12 @@ class IRMutatorWithAnalyzer : public tirx::StmtExprMutator
{
auto f_use_itervar = [&iter_var_nodes](const tirx::VarNode* v) {
return iter_var_nodes.count(v);
};
+ auto walkfn = [&](const tirx::Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return f_use_itervar(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
// simple heuristics for detecting predicate
- if (tirx::UsesVar(condition, f_use_itervar)) {
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(condition,
walkfn).has_value()) {
iter_predicates_.push_back(condition);
callback();
iter_predicates_.pop_back();
diff --git a/src/arith/iter_affine_map.cc b/src/arith/iter_affine_map.cc
index 3b8daaa7f4..534b826d61 100644
--- a/src/arith/iter_affine_map.cc
+++ b/src/arith/iter_affine_map.cc
@@ -23,6 +23,7 @@
#include <tvm/arith/analyzer.h>
#include <tvm/arith/iter_affine_map.h>
#include <tvm/ffi/cast.h>
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/prim/expr.h>
#include <tvm/tirx/analysis.h>
@@ -1348,12 +1349,20 @@ bool MatchBoundConstraints(PrimExpr pred,
ffi::Map<PrimVar, Range>* input_iters,
auto f_use_itervar = [&input_iter_nodes](const VarNode* v) {
return input_iter_nodes.count(v);
};
+ auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return f_use_itervar(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ bool lhs_uses_itervar =
+ ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(lhs_expr,
walkfn).has_value();
+ bool rhs_uses_itervar =
+ ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(rhs_expr,
walkfn).has_value();
bool bound_at_left;
- if (UsesVar(lhs_expr, f_use_itervar) || UsesVar(rhs_expr, f_use_itervar)) {
+ if (lhs_uses_itervar || rhs_uses_itervar) {
// At least it uses one input iter
- if (is_const_int(lhs_expr) || !UsesVar(lhs_expr, f_use_itervar)) {
+ if (is_const_int(lhs_expr) || !lhs_uses_itervar) {
bound_at_left = true;
- } else if (is_const_int(rhs_expr) || !UsesVar(rhs_expr, f_use_itervar)) {
+ } else if (is_const_int(rhs_expr) || !rhs_uses_itervar) {
bound_at_left = false;
} else {
bound_at_left = false; // accumulate bound to rhs
@@ -1361,14 +1370,14 @@ bool MatchBoundConstraints(PrimExpr pred,
ffi::Map<PrimVar, Range>* input_iters,
lhs_expr = 0;
rhs_expr = 0;
std::function<void(const PrimExpr&, bool)> f_extract =
- [&lhs_expr, &rhs_expr, f_use_itervar, &f_extract](const PrimExpr&
part, bool sign) {
+ [&lhs_expr, &rhs_expr, &walkfn, &f_extract](const PrimExpr& part,
bool sign) {
if (const prim::AddNode* add = part.as<prim::AddNode>()) {
f_extract(add->a, sign);
f_extract(add->b, sign);
} else if (const prim::SubNode* sub = part.as<prim::SubNode>()) {
f_extract(sub->a, sign);
f_extract(sub->b, !sign);
- } else if (UsesVar(part, f_use_itervar)) {
+ } else if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(part,
walkfn).has_value()) {
lhs_expr = sign ? lhs_expr + part : lhs_expr - part;
} else {
rhs_expr = sign ? rhs_expr - part : rhs_expr + part;
@@ -1429,8 +1438,15 @@ bool IterRangeSanityCheck(const ffi::Map<PrimVar,
Range>& iter_ranges) {
std::unordered_set<Var> iters;
for (const auto& it : iter_ranges) iters.insert(it.first);
auto f = [&](const VarNode* var) { return
iters.count(ffi::GetRef<Var>(var)); };
+ auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return f(var.get()) ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
for (const auto& it : iter_ranges) {
- if (UsesVar(it.second->min, f) || UsesVar(it.second->extent, f)) return
false;
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(it.second->min,
walkfn).has_value() ||
+ ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(it.second->extent,
walkfn).has_value()) {
+ return false;
+ }
}
return true;
}
diff --git a/src/relax/analysis/tir_op_pattern_kind.cc
b/src/relax/analysis/tir_op_pattern_kind.cc
index de8a8d06d9..646f7d4cca 100644
--- a/src/relax/analysis/tir_op_pattern_kind.cc
+++ b/src/relax/analysis/tir_op_pattern_kind.cc
@@ -236,10 +236,13 @@ class PatternKindAnalyzer : public StmtExprVisitor {
return false;
}
}
+ auto walkfn = [&vars](const tirx::Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return !vars.count(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
for (const PrimExpr& load_index : load->indices) {
// return false if there are vars used in load indices but not in store
indices.
- if (tirx::UsesVar(load_index,
- [&vars](const tirx::VarNode* var) { return
!vars.count(var); })) {
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(load_index,
walkfn).has_value()) {
return false;
}
}
@@ -318,17 +321,20 @@ class PatternKindAnalyzer : public StmtExprVisitor {
*/
static bool IsPureReducePattern(ffi::Array<tirx::Var> reduce_loops,
ffi::Array<PrimExpr> indices) {
+ auto walkfn = [&](const tirx::Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return std::any_of(reduce_loops.begin(), reduce_loops.end(),
+ [&](const tirx::Var& loop) { return
loop.same_as(var); })
+ ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
for (const PrimExpr& e : indices) {
- int id = -1;
- if (UsesVar(e, [&](const tirx::VarNode* var) {
- for (size_t i = 0; i < reduce_loops.size(); ++i) {
- if (reduce_loops[i].get() == var) {
- id = i;
- return true;
- }
- }
- return false;
- })) {
+ auto result = ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(e, walkfn);
+ if (result.has_value()) {
+ tirx::Var var = result.value()->value.cast<tirx::Var>();
+ int id =
+ std::distance(reduce_loops.begin(),
+ std::find_if(reduce_loops.begin(),
reduce_loops.end(),
+ [&](const tirx::Var& loop) { return
loop.same_as(var); }));
if (!reduce_loops[id].same_as(e)) {
return false;
}
diff --git a/src/relax/transform/rewrite_dataflow_reshape.cc
b/src/relax/transform/rewrite_dataflow_reshape.cc
index 464be7d774..e99f10219d 100644
--- a/src/relax/transform/rewrite_dataflow_reshape.cc
+++ b/src/relax/transform/rewrite_dataflow_reshape.cc
@@ -22,6 +22,7 @@
*/
#include <tvm/arith/analyzer.h>
#include <tvm/ffi/cast.h>
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/relax/analysis.h>
#include <tvm/relax/expr_functor.h>
@@ -41,8 +42,11 @@ std::vector<size_t> GetUsedTensorArgIndices(const
tirx::PrimFunc& fn, size_t num
for (size_t i = 0; i < num_args; ++i) {
if (auto buffer = fn->params[i].as<tirx::BufferVar>()) {
auto buffer_var = buffer.value().var();
- if (tirx::UsesVar(fn->body,
- [=](const tirx::VarNode* var) { return var ==
buffer_var.get(); })) {
+ auto walkfn = [=](const tirx::Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return var.get() == buffer_var.get() ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(fn->body,
walkfn).has_value()) {
indices.push_back(i);
}
}
diff --git a/src/s_tir/meta_schedule/postproc/rewrite_reduction_block.cc
b/src/s_tir/meta_schedule/postproc/rewrite_reduction_block.cc
index 92ab70201d..9a7a4acf6e 100644
--- a/src/s_tir/meta_schedule/postproc/rewrite_reduction_block.cc
+++ b/src/s_tir/meta_schedule/postproc/rewrite_reduction_block.cc
@@ -16,6 +16,7 @@
* specific language governing permissions and limitations
* under the License.
*/
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/s_tir/stmt.h>
@@ -67,6 +68,10 @@ struct ReductionBlockFinder : private StmtVisitor {
return true;
}
auto f_find = [this](const VarNode* var) -> bool { return
thread_bound_loop_vars_.count(var); };
+ auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return f_find(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
const SBlockNode* block = realize->block.get();
TVM_FFI_ICHECK_EQ(block->iter_vars.size(), realize->iter_values.size());
int n = block->iter_vars.size();
@@ -74,7 +79,7 @@ struct ReductionBlockFinder : private StmtVisitor {
IterVar iter_var = block->iter_vars[i];
PrimExpr binding = realize->iter_values[i];
if (iter_var->iter_type == tirx::kCommReduce) {
- if (UsesVar(binding, f_find)) {
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(binding,
walkfn).has_value()) {
return false;
}
}
diff --git a/src/s_tir/schedule/analysis/analysis.cc
b/src/s_tir/schedule/analysis/analysis.cc
index 342632b6ce..8e715c83fd 100644
--- a/src/s_tir/schedule/analysis/analysis.cc
+++ b/src/s_tir/schedule/analysis/analysis.cc
@@ -1812,6 +1812,14 @@ ffi::Optional<TensorizeInfo>
GetTensorizeLoopMapping(const s_tir::ScheduleState&
// C[i, j] += A[i, k] * B[k, j]
int next_block_ind = block_loops.size() - 1;
+ auto desc_walkfn = [&desc_loop_vars](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return desc_loop_vars.count(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ auto block_walkfn = [&block_loop_vars](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return block_loop_vars.count(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
for (int i_desc = n_desc_vars - 1; i_desc >= 0; --i_desc) {
// Step 3.1. Find the corresponding loop of the i_desc-th block var of desc
const PrimExpr& desc_bind = desc_block->iter_values[i_desc];
@@ -1820,8 +1828,7 @@ ffi::Optional<TensorizeInfo>
GetTensorizeLoopMapping(const s_tir::ScheduleState&
for (int i = 0, n = desc_loops.size(); i < n; ++i) {
// Check if desc_bind = loops[i]->loop_var +
stuff-irrelevant-of-loop-vars
PrimExpr residual = analyzer->Simplify(desc_bind -
desc_loops[i]->loop_var);
- if (!UsesVar(residual,
- [&desc_loop_vars](const VarNode* var) { return
desc_loop_vars.count(var); })) {
+ if (!ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(residual,
desc_walkfn).has_value()) {
desc_loop = desc_loops[i];
iter_type_desc = iter_types_desc[i];
break;
@@ -1855,8 +1862,7 @@ ffi::Optional<TensorizeInfo>
GetTensorizeLoopMapping(const s_tir::ScheduleState&
if (ret->loop_map.find(block_loop_sref) != ret->loop_map.end()) continue;
PrimExpr residual = analyzer->Simplify(block_bind -
block_loops[i]->loop_var);
- if (UsesVar(residual,
- [&block_loop_vars](const VarNode* var) { return
block_loop_vars.count(var); })) {
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(residual,
block_walkfn).has_value()) {
continue;
}
// padding is allowed only when the block has trivial bindings
diff --git a/src/s_tir/schedule/analysis/reducer.cc
b/src/s_tir/schedule/analysis/reducer.cc
index 477b9b38fb..bff5bec91e 100644
--- a/src/s_tir/schedule/analysis/reducer.cc
+++ b/src/s_tir/schedule/analysis/reducer.cc
@@ -554,10 +554,13 @@ bool ReductionIterNotIndexOutputBuffer(const SBlock&
block) {
buffer_allocated.insert(buffer.get());
}
+ auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return reduction_block_iters.count(var.get())
+ ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
auto f_uses_reduction_block_var = [&](const PrimExpr& expr) -> bool {
- return UsesVar(expr, [&](const VarNode* var) { //
- return reduction_block_iters.count(var);
- });
+ return ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(expr,
walkfn).has_value();
};
std::unordered_map<const VarNode*, const VarNode*> match_buffer_sources;
diff --git a/src/s_tir/schedule/primitive/blockize_tensorize.cc
b/src/s_tir/schedule/primitive/blockize_tensorize.cc
index ca5053c921..3dea068849 100644
--- a/src/s_tir/schedule/primitive/blockize_tensorize.cc
+++ b/src/s_tir/schedule/primitive/blockize_tensorize.cc
@@ -18,6 +18,7 @@
*/
#include <tvm/ffi/cast.h>
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/runtime/logging.h>
#include <functional>
@@ -32,11 +33,6 @@ namespace s_tir {
using namespace tvm::prim;
using namespace tvm::tirx;
-template <class T>
-bool UsesVar(const T& x, const Var& var) {
- return tirx::UsesVar(x, [tgt = var.get()](const VarNode* v) { return v ==
tgt; });
-}
-
Range RangeFromExtent(const PrimExpr& extent) {
return Range::FromMinExtent(IntImm(extent.ty(), 0), extent);
}
@@ -110,9 +106,11 @@ ffi::Array<ffi::Array<arith::IterMark>>
TrivialSubspaceDivision(
var_set.insert(var.get());
}
return [var_set = std::move(var_set)](const PrimExpr& expr) -> bool {
- return tirx::UsesVar(expr, [&var_set](const VarNode* var) {
- return var_set.count(var); //
- });
+ auto walkfn = [&var_set](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return var_set.count(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ return ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(expr,
walkfn).has_value();
};
};
auto use_outer_loop_vars = make_uses_var(outer_iters);
@@ -335,8 +333,13 @@ Stmt GenerateOuterInit(const Stmt& block_init, const
SBlockRealize& inner_realiz
for (int i = 0; i < n; ++i) {
const IterVar& old_iter_var = inner_block->iter_vars[i];
const PrimExpr& iter_value = inner_realize->iter_values[i];
+ auto walkfn = [target =
+ old_iter_var->var.get()](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return var.get() == target ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
if (old_iter_var->iter_type == IterVarType::kDataPar &&
- UsesVar(block_init, old_iter_var->var)) {
+ ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(block_init,
walkfn).has_value()) {
ffi::ObjectPtr<IterVarNode> new_iter_var =
ffi::make_object<IterVarNode>(*old_iter_var.get());
new_iter_var->var = new_iter_var->var.CopyWithSuffix("_init");
subst_map.Set(old_iter_var->var, new_iter_var->var);
@@ -358,8 +361,13 @@ Stmt GenerateOuterInit(const Stmt& block_init, const
SBlockRealize& inner_realiz
// Step 3. Create the loop nest on top of the block
for (const ForNode* loop : loops) {
bool is_init_loop = false;
+ auto walkfn = [target =
+ loop->loop_var.get()](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return var.get() == target ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
for (const PrimExpr& init_binding : iter_values) {
- if (UsesVar(init_binding, loop->loop_var)) {
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(init_binding,
walkfn).has_value()) {
is_init_loop = true;
break;
}
diff --git a/src/s_tir/schedule/primitive/for_kind.cc
b/src/s_tir/schedule/primitive/for_kind.cc
index 833877af21..58e19d8496 100644
--- a/src/s_tir/schedule/primitive/for_kind.cc
+++ b/src/s_tir/schedule/primitive/for_kind.cc
@@ -98,7 +98,11 @@ void CheckLoopParallelizableInBlock(const ScheduleState&
self, ForKind for_kind,
const IterVar& iter_var = block->iter_vars[i];
const PrimExpr& binding = block_realize->iter_values[i];
- if (!UsesVar(binding, [v = loop_var.get()](const VarNode* var) { return
var == v; })) {
+ auto walkfn = [v = loop_var.get()](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return var.get() == v ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ if (!ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(binding,
walkfn).has_value()) {
continue;
}
// Only two cases are allowed:
diff --git a/src/s_tir/schedule/primitive/loop_transformation.cc
b/src/s_tir/schedule/primitive/loop_transformation.cc
index e7ee773ac5..6c354769df 100644
--- a/src/s_tir/schedule/primitive/loop_transformation.cc
+++ b/src/s_tir/schedule/primitive/loop_transformation.cc
@@ -17,6 +17,7 @@
* under the License.
*/
#include <tvm/ffi/cast.h>
+#include <tvm/ffi/extra/structural_visit.h>
#include "../utils.h"
@@ -907,15 +908,13 @@ StmtSRef Fuse(ScheduleState self, const
ffi::Array<StmtSRef>& loop_srefs,
outer_loop_sref = sref;
outer_loop = loop;
CheckLoopStartsWithZero(self, sref, analyzer.get());
- const VarNode* used_var = nullptr;
- auto f_contain = [&outer_loop_vars, &used_var](const VarNode* var) {
- if (outer_loop_vars.count(var)) {
- used_var = var;
- return true;
- }
- return false;
+ auto walkfn = [&outer_loop_vars](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return outer_loop_vars.count(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
};
- if (UsesVar(loop->extent, f_contain)) {
+ auto result = ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(loop->extent,
walkfn);
+ if (result.has_value()) {
+ Var used_var = result.value()->value.cast<Var>();
throw DependentLoopError(self->mod, ffi::GetRef<For>(loop),
used_var->name,
DependentLoopError::PrimitiveKind::kFuse);
}
@@ -1105,15 +1104,16 @@ For ConstructNewLoopChain(const ScheduleState& self,
std::vector<const StmtSRefN
} else {
n->body = loop_sref->StmtAs<ForNode>()->body;
}
- const VarNode* used_var = nullptr;
- auto f_contain = [&inner_vars, &used_var](const VarNode* var) {
- if (inner_vars.count(var)) {
- used_var = var;
- return true;
- }
- return false;
+ auto walkfn = [&inner_vars](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return inner_vars.count(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
};
- if (UsesVar(copy->min, f_contain) || UsesVar(copy->extent, f_contain)) {
+ auto result = ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(copy->min,
walkfn);
+ if (!result.has_value()) {
+ result = ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(copy->extent,
walkfn);
+ }
+ if (result.has_value()) {
+ Var used_var = result.value()->value.cast<Var>();
throw DependentLoopError(self->mod, ffi::GetRef<For>(copy),
used_var->name,
DependentLoopError::PrimitiveKind::kReorder);
}
diff --git a/src/s_tir/schedule/primitive/reduction.cc
b/src/s_tir/schedule/primitive/reduction.cc
index beb05afbd5..a96face461 100644
--- a/src/s_tir/schedule/primitive/reduction.cc
+++ b/src/s_tir/schedule/primitive/reduction.cc
@@ -17,6 +17,7 @@
* under the License.
*/
#include <tvm/ffi/cast.h>
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/te/operation.h>
@@ -128,7 +129,11 @@ class LoopHeightError : public ScheduleError {
}
// loop_var of a higher loop shouldn't contain loop var
const Var& loop_var = higher_loop->StmtAs<ForNode>()->loop_var;
- if (UsesVar(binding, [v = loop_var.get()](const VarNode* var) { return
var == v; })) {
+ auto walkfn = [v = loop_var.get()](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return var.get() == v ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(binding,
walkfn).has_value()) {
const ForNode* loop = TVM_SREF_TO_FOR(loop_sref);
throw LoopHeightError(mod, ffi::GetRef<For>(loop),
ffi::GetRef<SBlock>(block));
}
@@ -169,7 +174,13 @@ PrimExpr RewriteInitPredicate(PrimExpr pred,
auto uses_discarded_loop = [&discarded_loops](const VarNode* var) {
return discarded_loops.count(var);
};
- return UsesVar(pred, uses_discarded_loop) ? IntImm::Bool(true) : pred;
+ auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return uses_discarded_loop(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ return ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(pred,
walkfn).has_value()
+ ? IntImm::Bool(true)
+ : pred;
}
StmtSRef DecomposeReduction(ScheduleState self, const StmtSRef& block_sref,
@@ -244,8 +255,12 @@ StmtSRef DecomposeReduction(ScheduleState self, const
StmtSRef& block_sref,
for (int i = static_cast<int>(loops.size()) - 1; i >= 0; --i) {
const VarNode* loop_var = loops[i]->StmtAs<ForNode>()->loop_var.get();
bool discarded = true;
+ auto walkfn = [v = loop_var](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return var.get() == v ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
for (const PrimExpr& expr : init_realize->iter_values) {
- if (!UsesVar(expr, [v = loop_var](const VarNode* var) { return var == v;
})) {
+ if (!ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(expr,
walkfn).has_value()) {
continue;
}
// The loop is related to init block bindings;
@@ -896,9 +911,12 @@ class RFactorBlockCreator : public BaseBlockCreator {
void CreateNormalIters(int idx) final {
IterVar old_iter = old_block_realize_->block->iter_vars[idx];
PrimExpr old_binding = old_block_realize_->iter_values[idx];
+ auto walkfn = [v = rf_loop_->loop_var.get()](const Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return var.get() == v ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
if (old_iter->iter_type == IterVarType::kDataPar ||
- !UsesVar(old_binding,
- [v = rf_loop_->loop_var.get()](const VarNode* var) { return
var == v; })) {
+ !ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(old_binding,
walkfn).has_value()) {
// The old block iter is either a data parallel block iter, or a
reduction block iter that
// doesn't touch the rfactor loop. In this case reuse the old reduction
block iter and its
// corresponding binding.
diff --git a/src/s_tir/transform/compact_buffer_region.cc
b/src/s_tir/transform/compact_buffer_region.cc
index 2ccbd8e905..7255317fa3 100644
--- a/src/s_tir/transform/compact_buffer_region.cc
+++ b/src/s_tir/transform/compact_buffer_region.cc
@@ -25,6 +25,7 @@
#include <tvm/arith/int_set.h>
#include <tvm/arith/int_solver.h>
#include <tvm/ffi/cast.h>
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/s_tir/stmt.h>
#include <tvm/s_tir/transform.h>
@@ -472,7 +473,11 @@ class BufferAccessRegionCollector : public StmtExprVisitor
{
return std::any_of(ancestor_iters_.begin(), ancestor_iters_.end(),
[v](const IterVar& n) { return n->var.get() == v;
});
};
- if (UsesVar(extent, is_loop_var)) {
+ auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return is_loop_var(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(extent,
walkfn).has_value()) {
// try estimate a constant upperbound on region's extent
int64_t upperbound = dom_analyzer_->const_int_bound(extent)->max_value;
if (upperbound != arith::ConstIntBound::kPosInf) {
diff --git a/src/s_tir/transform/hoist_expression.cc
b/src/s_tir/transform/hoist_expression.cc
index 6aa8082a98..40a4b47daf 100644
--- a/src/s_tir/transform/hoist_expression.cc
+++ b/src/s_tir/transform/hoist_expression.cc
@@ -22,6 +22,7 @@
*/
#include <tvm/arith/analyzer.h>
#include <tvm/ffi/cast.h>
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/function.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/prim/expr.h>
@@ -217,9 +218,14 @@ class HoistInfoCollector : public StmtExprVisitor {
if (auto info = FindHoistDestination(cond)) {
if (!info->reached_sequential_node) {
// Record whether this conditional uses any block variables.
- bool uses_block_var = active_block_vars.size() && UsesVar(cond,
[&](const VarNode* var) {
- return active_block_vars.count(var);
- });
+ auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return active_block_vars.count(var.get())
+ ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ bool uses_block_var =
+ active_block_vars.size() &&
+ ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(cond,
walkfn).has_value();
std::unordered_set<const VarNode*> let_bindings_used;
@@ -389,18 +395,16 @@ class HoistInfoCollector : public StmtExprVisitor {
for (auto it = active_loops.rbegin(); it != active_loops.rend(); it++) {
Var loop_var = it->loop_var;
- bool uses_loop_var = UsesVar(expr, [&](const VarNode* var) -> bool {
- if (var == loop_var.get()) {
- return true;
+ auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ bool matches = var.get() == loop_var.get();
+ if (!matches) {
+ auto it = let_var_to_loop_vars.find(var.get());
+ matches = it != let_var_to_loop_vars.end() &&
it->second.count(loop_var.get());
}
-
- auto it = let_var_to_loop_vars.find(var);
- if (it == let_var_to_loop_vars.end()) {
- return false;
- }
-
- return it->second.count(loop_var.get());
- });
+ return matches ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ bool uses_loop_var =
ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(expr, walkfn).has_value();
bool is_disabled_hoist_across_block_var =
!config->FlagSet(HoistedConditionals::kUsingBlockVar) &&
it->IsBlockVariable();
diff --git a/src/s_tir/transform/loop_partition.cc
b/src/s_tir/transform/loop_partition.cc
index 56d6b13038..e04a48687e 100644
--- a/src/s_tir/transform/loop_partition.cc
+++ b/src/s_tir/transform/loop_partition.cc
@@ -23,6 +23,7 @@
#include <tvm/arith/analyzer.h>
#include <tvm/arith/bound.h>
#include <tvm/ffi/cast.h>
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/function.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/prim/builtin.h>
@@ -248,7 +249,14 @@ class PartitionFinder : public StmtExprVisitor {
void VisitStmt_(const ForNode* op) final {
auto f_vset_contains = [this](const VarNode* var) { return
out_vars_.count(var); };
- if (UsesVar(op->min, f_vset_contains) || UsesVar(op->extent,
f_vset_contains)) return;
+ auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return f_vset_contains(var.get()) ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(op->min,
walkfn).has_value() ||
+ ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(op->extent,
walkfn).has_value()) {
+ return;
+ }
const VarNode* var = op->loop_var.get();
hint_map_.insert({var, IntSet::Interval(op->min, op->min + op->extent -
1)});
@@ -299,7 +307,11 @@ class PartitionFinder : public StmtExprVisitor {
// For cond, find out the interval, if exists, in which we can prove that
cond is
// true. Also find the interval, if exists, in which we can prove that
cond is
// false.
- if (UsesVar(cond, [this](const VarNode* var) { return var ==
current_var_.get(); })) {
+ auto walkfn = [this](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return var.get() == current_var_.get() ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ if (ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(cond,
walkfn).has_value()) {
IntSet interval =
DeduceBound(current_var_.as_or_throw<PrimExpr>(), cond, hint_map_,
relax_map_);
if (!interval.IsNothing()) {
diff --git a/src/s_tir/transform/thread_storage_sync.cc
b/src/s_tir/transform/thread_storage_sync.cc
index d284accd42..200ef98ae9 100644
--- a/src/s_tir/transform/thread_storage_sync.cc
+++ b/src/s_tir/transform/thread_storage_sync.cc
@@ -20,6 +20,7 @@
/*!
* \file thread_storage_sync.cc
*/
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/function.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/prim/builtin.h>
@@ -237,9 +238,15 @@ class ThreadSyncPlanner : public StorageAccessVisitor {
auto f_uses_thread_index = [=](const tvm::tirx::VarNode* parameter) {
return parameter == thread_index_var;
};
- depends_on_thread_index = depends_on_thread_index &&
- UsesVar(curr_index, f_uses_thread_index) &&
- UsesVar(prev_index, f_uses_thread_index);
+ auto walkfn = [&](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return f_uses_thread_index(var.get())
+ ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ depends_on_thread_index =
+ depends_on_thread_index &&
+ ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(curr_index,
walkfn).has_value() &&
+ ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(prev_index,
walkfn).has_value();
}
} else {
has_same_index = false;
diff --git a/src/tirx/analysis/var_touch.cc b/src/tirx/analysis/var_touch.cc
deleted file mode 100644
index d3441b3351..0000000000
--- a/src/tirx/analysis/var_touch.cc
+++ /dev/null
@@ -1,80 +0,0 @@
-/*
- * Licensed to the Apache Software Foundation (ASF) under one
- * or more contributor license agreements. See the NOTICE file
- * distributed with this work for additional information
- * regarding copyright ownership. The ASF licenses this file
- * to you under the Apache License, Version 2.0 (the
- * "License"); you may not use this file except in compliance
- * with the License. You may obtain a copy of the License at
- *
- * http://www.apache.org/licenses/LICENSE-2.0
- *
- * Unless required by applicable law or agreed to in writing,
- * software distributed under the License is distributed on an
- * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
- * KIND, either express or implied. See the License for the
- * specific language governing permissions and limitations
- * under the License.
- */
-
-/*!
- * \file var_touch.cc
- * \brief Implementation of simple passes
- */
-#include <tvm/tirx/analysis.h>
-#include <tvm/tirx/stmt_functor.h>
-
-namespace tvm {
-namespace tirx {
-
-class VarTouchVisitor : public StmtExprVisitor {
- public:
- explicit VarTouchVisitor(std::function<bool(const VarNode*)> var_set)
- : var_set_(std::move(var_set)) {}
-
- void VisitStmt(const Stmt& stmt) final {
- if (use_var_) return;
- StmtExprVisitor::VisitStmt(stmt);
- }
-
- void VisitExpr(const Expr& e) final {
- if (use_var_) return;
- StmtExprVisitor::VisitExpr(e);
- }
-
- void VisitExpr_(const VarNode* op) final { Handle(op); }
-
- void VisitStmt_(const BufferStoreNode* op) final {
- Handle(op->buffer.get());
- StmtVisitor::VisitStmt_(op);
- }
-
- void VisitExpr_(const TensorLoadNode* op) final {
- Handle(op->source.as_or_throw<tvm::tirx::BufferVar>().get());
- ExprVisitor::VisitExpr_(op);
- }
-
- void Handle(const VarNode* var) {
- if (var_set_(var)) use_var_ = true;
- }
-
- bool use_var_{false};
-
- private:
- std::function<bool(const VarNode*)> var_set_;
-};
-
-bool UsesVar(const Stmt& stmt, std::function<bool(const VarNode*)> var_set) {
- VarTouchVisitor visitor(std::move(var_set));
- visitor(stmt);
- return visitor.use_var_;
-}
-
-bool UsesVar(const PrimExpr& expr, std::function<bool(const VarNode*)>
var_set) {
- VarTouchVisitor visitor(std::move(var_set));
- visitor(expr);
- return visitor.use_var_;
-}
-
-} // namespace tirx
-} // namespace tvm
diff --git a/src/tirx/script/printer/for_loop.cc
b/src/tirx/script/printer/for_loop.cc
index 363a2d466f..4d0f2016cd 100644
--- a/src/tirx/script/printer/for_loop.cc
+++ b/src/tirx/script/printer/for_loop.cc
@@ -16,6 +16,8 @@
* specific language governing permissions and limitations
* under the License.
*/
+#include <tvm/ffi/extra/structural_visit.h>
+
#include "./utils.h"
namespace tvm {
@@ -28,9 +30,12 @@ TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable)
std::vector<const tirx::ForNode*> grid;
std::unordered_set<const tirx::VarNode*> grid_loop_vars;
auto f_var_dep = [&grid_loop_vars](const PrimExpr& e) -> bool {
- return tirx::UsesVar(e, [&grid_loop_vars](const tirx::VarNode* v) ->
bool { //
- return grid_loop_vars.count(v);
- });
+ auto walkfn = [&grid_loop_vars](const tirx::Var& var) ->
ffi::Expected<ffi::WalkResult> {
+ return grid_loop_vars.count(var.get())
+ ? ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+ return ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(e,
walkfn).has_value();
};
if (d->cfg->syntax_sugar) {
for (const tirx::ForNode* l = loop.get(); l != nullptr; l =
l->body.as<tirx::ForNode>()) {
diff --git a/src/tirx/transform/ir_utils.cc b/src/tirx/transform/ir_utils.cc
index 6e7e44b780..7c5c717caa 100644
--- a/src/tirx/transform/ir_utils.cc
+++ b/src/tirx/transform/ir_utils.cc
@@ -569,7 +569,12 @@ class IRConvertSSA final : public StmtExprMutator {
if (buffer.get() == var) return true;
auto uses_var = [var](const PrimExpr& expr) {
- return expr.defined() && UsesVar(expr, [var](const VarNode* node) {
return node == var; });
+ auto walkfn = [var](const Var& candidate) ->
ffi::Expected<ffi::WalkResult> {
+ return candidate.get() == var ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(candidate))
+ : ffi::WalkResult::Advance();
+ };
+ return expr.defined() &&
+ ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(expr,
walkfn).has_value();
};
if (uses_var(buffer->elem_offset)) return true;
for (const PrimExpr& dim : buffer->shape) {
diff --git a/src/tirx/transform/lower_warp_memory.cc
b/src/tirx/transform/lower_warp_memory.cc
index ed2afbfc2a..4f268b40e0 100644
--- a/src/tirx/transform/lower_warp_memory.cc
+++ b/src/tirx/transform/lower_warp_memory.cc
@@ -28,6 +28,7 @@
#include <tvm/arith/analyzer.h>
#include <tvm/arith/pattern.h>
#include <tvm/ffi/cast.h>
+#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/function.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/op.h>
@@ -384,8 +385,11 @@ class WarpAccessRewriter : protected StmtExprMutator {
auto [local_index, group] = SplitIndexByGroup(op->indices[0]);
// invariance: local index must do not contain warp id
- TVM_FFI_ICHECK(
- !UsesVar(local_index, [this](const VarNode* var) { return var ==
warp_index_.get(); }))
+ auto walkfn = [this](const Var& var) -> ffi::Expected<ffi::WalkResult> {
+ return var.get() == warp_index_.get() ?
ffi::WalkResult::Interrupt(ffi::VisitInterrupt(var))
+ : ffi::WalkResult::Advance();
+ };
+
TVM_FFI_ICHECK(!ffi::StructuralWalk<ffi::WalkOrder::kPreOrder>(local_index,
walkfn).has_value())
<< "LowerWarpMemory failed to rewrite load to shuffle for index " <<
op->indices[0]
<< " local_index=" << local_index;