diff --git a/src/tirx/ir/specialize.cc b/src/tirx/ir/specialize.cc index 3890829dc14a..d1a687af5cdf 100644 --- a/src/tirx/ir/specialize.cc +++ b/src/tirx/ir/specialize.cc @@ -174,12 +174,13 @@ class PrimFuncSpecializer : public StmtExprMutator { BufferVar VisitBufferUse(const BufferVar& buffer) final { return GetNewBuffer(buffer); } Expr VisitExpr_(const VarNode* op) final { + Var var = ffi::GetRef(op); if (constrained_buffer_params_.count(op)) { - return ffi::GetRef(op); + return var; } - auto it = var_map_.find(ffi::GetRef(op)); + auto it = var_map_.find(var); if (it == var_map_.end()) { - return ffi::GetRef(op); + return StmtExprMutator::VisitExpr_(op); } else { return it->second; } diff --git a/tests/python/tirx-base/test_tir_specialize.py b/tests/python/tirx-base/test_tir_specialize.py index f47a6dc591ae..4dffc8dc11e9 100644 --- a/tests/python/tirx-base/test_tir_specialize.py +++ b/tests/python/tirx-base/test_tir_specialize.py @@ -266,6 +266,24 @@ def expected(A_data: T.handle("float32")): tvm.ir.assert_structural_equal(expected, after) +def test_specialize_preserves_decl_buffer_alias(): + @T.prim_func(private=True, s_tir=True) + def before(A_handle: T.handle, n: T.int32): + A = T.match_buffer(A_handle, (n,), "int32") + A_flat = T.decl_buffer((n,), "int32", data=A.data) + A_flat[n - 1] = 42 + + @T.prim_func(private=True, s_tir=True) + def expected(A_handle: T.handle): + A = T.match_buffer(A_handle, (8,), "int32") + A_flat = T.decl_buffer((8,), "int32", data=A.data) + A_flat[7] = 42 + + after = before.specialize({before.params[1]: 8}) + + tvm.ir.assert_structural_equal(expected, after) + + def test_specialize_buffer_var_to_var(): """A buffer var may be remapped by specialization