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;
 

Reply via email to