From 70adfce16aebeb184d24524cde2667645e5ae090 Mon Sep 17 00:00:00 2001 From: Hongyi Jin Date: Wed, 29 Jul 2026 02:54:46 +0000 Subject: [PATCH] [FIX][TIRx] Remap buffers consistently in ConvertSSA ConvertSSA caches remapped buffers while an SSA-renamed variable is in scope. The cleanup previously popped a cached buffer only when the renamed variable was its data pointer, leaving remaps through elem_offset, shape, strides, or tile layout fields alive after scope exit. Recognize dependencies in every field rewritten by GetRemappedBuffer, and add a regression with reused sibling loop variables and a variable-dependent elem_offset. --- 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 5ced08b335f8..44b2d86e211a 100644 --- a/src/tirx/transform/ir_utils.cc +++ b/src/tirx/transform/ir_utils.cc @@ -29,6 +29,7 @@ #include #include #include +#include #include #include #include @@ -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()) { + 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& 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& 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& 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& 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 59275728c67a..59a9a5c93ba5 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()