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()