This is an automated email from the ASF dual-hosted git repository.

spectrometerHBH 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 0fbc04baeb [FIX][TIRx] Remap buffers consistently in ConvertSSA 
(#20069)
0fbc04baeb is described below

commit 0fbc04baeb285086802c0792ad4f65ddb17bc40d
Author: Hongyi Jin <[email protected]>
AuthorDate: Wed Jul 29 14:03:07 2026 -0400

    [FIX][TIRx] Remap buffers consistently in ConvertSSA (#20069)
    
    ## Motivation and context
    
    ConvertSSA caches remapped buffers while an SSA-renamed variable is in
    scope. Cleanup previously removed a cached buffer only when the renamed
    variable was its data pointer, so remaps through other buffer fields
    could survive after scope exit.
    
    ## Changes
    
    - Track dependencies in every buffer field rewritten by
    `GetRemappedBuffer`, including shape, strides, element offset, and tile
    layout fields.
    - Invalidate cached remaps consistently when an SSA-renamed variable
    leaves scope.
    - Add a regression using reused sibling loop variables and a
    variable-dependent element offset.
    
    ## Testing
    
    - `python -m pytest
    tests/python/tirx-transform/test_tir_transform_convert_ssa.py`
    - Changed-files pre-commit checks
---
 src/tirx/transform/ir_utils.cc                     | 36 +++++++++++++++++++---
 .../test_tir_transform_convert_ssa.py              | 24 +++++++++++++++
 2 files changed, 56 insertions(+), 4 deletions(-)

diff --git a/src/tirx/transform/ir_utils.cc b/src/tirx/transform/ir_utils.cc
index 5ced08b335..44b2d86e21 100644
--- a/src/tirx/transform/ir_utils.cc
+++ b/src/tirx/transform/ir_utils.cc
@@ -29,6 +29,7 @@
 #include <tvm/ffi/reflection/registry.h>
 #include <tvm/ir/scope_stack.h>
 #include <tvm/s_tir/stmt.h>
+#include <tvm/tirx/analysis.h>
 #include <tvm/tirx/layout.h>
 #include <tvm/tirx/stmt_functor.h>
 #include <tvm/tirx/transform.h>
@@ -526,6 +527,33 @@ class IRConvertSSA final : public StmtExprMutator {
     Var new_var;
   };
 
+  /*! \brief Check whether a buffer uses a variable in any remapped field. */
+  static bool BufferDependsOnVar(const Buffer& buffer, const VarNode* var) {
+    if (buffer->data.get() == var) return true;
+
+    auto uses_var = [var](const PrimExpr& expr) {
+      return expr.defined() && UsesVar(expr, [var](const VarNode* node) { 
return node == var; });
+    };
+    if (uses_var(buffer->elem_offset)) return true;
+    for (const PrimExpr& dim : buffer->shape) {
+      if (uses_var(dim)) return true;
+    }
+    for (const PrimExpr& stride : buffer->strides) {
+      if (uses_var(stride)) return true;
+    }
+    if (buffer->layout.has_value()) {
+      if (const auto* tile_layout = 
buffer->layout.value().as<TileLayoutNode>()) {
+        for (const Iter& iter : tile_layout->shard) {
+          if (uses_var(iter->extent) || uses_var(iter->stride)) return true;
+        }
+        for (const Iter& iter : tile_layout->replica) {
+          if (uses_var(iter->extent) || uses_var(iter->stride)) return true;
+        }
+      }
+    }
+    return false;
+  }
+
   /*! \brief Create a new variable with the same name and type as the 
original. */
   static Var MakeNewVar(const Var& old_var) { return Var(old_var->name, 
old_var->ty); }
 
@@ -542,7 +570,7 @@ class IRConvertSSA final : public StmtExprMutator {
     var_remap_[old_var.get()].pop_back();
     for (auto& kv : buf_remap_) {
       std::vector<Buffer>& buffers = kv.second;
-      if (buffers.size() && (buffers.back()->data.get() == new_var.get())) {
+      if (buffers.size() && BufferDependsOnVar(buffers.back(), new_var.get())) 
{
         buffers.pop_back();
       }
     }
@@ -561,7 +589,7 @@ class IRConvertSSA final : public StmtExprMutator {
       var_remap_[remap.old_var.get()].pop_back();
       for (auto& kv : buf_remap_) {
         std::vector<Buffer>& buffers = kv.second;
-        if (buffers.size() && (buffers.back()->data.get() == 
remap.new_var.get())) {
+        if (buffers.size() && BufferDependsOnVar(buffers.back(), 
remap.new_var.get())) {
           buffers.pop_back();
         }
       }
@@ -598,7 +626,7 @@ class IRConvertSSA final : public StmtExprMutator {
         parent->var_remap_[remap.old_var.get()].pop_back();
         for (auto& kv : parent->buf_remap_) {
           std::vector<Buffer>& buffers = kv.second;
-          if (buffers.size() && (buffers.back()->data.get() == 
remap.new_var.get())) {
+          if (buffers.size() && BufferDependsOnVar(buffers.back(), 
remap.new_var.get())) {
             buffers.pop_back();
           }
         }
@@ -622,7 +650,7 @@ class IRConvertSSA final : public StmtExprMutator {
             parent->var_remap_[remap.old_var.get()].pop_back();
             for (auto& kv : parent->buf_remap_) {
               std::vector<Buffer>& buffers = kv.second;
-              if (buffers.size() && (buffers.back()->data.get() == 
remap.new_var.get())) {
+              if (buffers.size() && BufferDependsOnVar(buffers.back(), 
remap.new_var.get())) {
                 buffers.pop_back();
               }
             }
diff --git a/tests/python/tirx-transform/test_tir_transform_convert_ssa.py 
b/tests/python/tirx-transform/test_tir_transform_convert_ssa.py
index 59275728c6..59a9a5c93b 100644
--- a/tests/python/tirx-transform/test_tir_transform_convert_ssa.py
+++ b/tests/python/tirx-transform/test_tir_transform_convert_ssa.py
@@ -535,5 +535,29 @@ def test_shared_shape_var_in_buffer_map_and_alloc_buffer():
     tvm.ir.assert_structural_equal(after["main"], before)
 
 
+def test_reused_loop_var_in_decl_buffer_elem_offset():
+    """Remap a buffer whose elem_offset depends on an SSA-renamed loop var."""
+    loop_var = tirx.Var("loop_var", "int32")
+    buffer = tirx.decl_buffer(
+        (128,),
+        "float32",
+        "buffer",
+        elem_offset=loop_var * 128,
+        scope="shared.dyn",
+    )
+    loop = tirx.For(
+        loop_var,
+        0,
+        128,
+        tirx.ForKind.SERIAL,
+        tirx.DeclBuffer(buffer, tirx.Evaluate(tirx.BufferLoad(buffer, [0]))),
+    )
+    func = tirx.PrimFunc([buffer.data], tirx.SeqStmt([loop, loop, loop]))
+
+    after = tvm.tirx.transform.ConvertSSA()(tvm.IRModule.from_expr(func))
+
+    tvm.tirx.analysis.verify_well_formed(after["main"], assert_mode=True)
+
+
 if __name__ == "__main__":
     tvm.testing.main()

Reply via email to