This is an automated email from the ASF dual-hosted git repository. tqchen pushed a commit to branch script/canonical-parser-df in repository https://gitbox.apache.org/repos/asf/tvm.git
commit 3c7079b615e24c0026776a04cb319af5820264dc Author: Tianqi Chen <[email protected]> AuthorDate: Wed Sep 23 08:48:06 2026 +0000 Cover binding names and allocation identity under explicit storage declarations --- .../tvmscript/test_tvmscript_error_report.py | 22 +++++++++++----- .../tvmscript/test_tvmscript_printer_annotation.py | 9 +++---- .../tvmscript/test_tvmscript_syntax_sugar.py | 29 ++++++++++++++-------- 3 files changed, 39 insertions(+), 21 deletions(-) diff --git a/tests/python/tvmscript/test_tvmscript_error_report.py b/tests/python/tvmscript/test_tvmscript_error_report.py index 1d2ef09ba3..aec053fbe7 100644 --- a/tests/python/tvmscript/test_tvmscript_error_report.py +++ b/tests/python/tvmscript/test_tvmscript_error_report.py @@ -14,7 +14,7 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -# ruff: noqa: E741, F401, F821, F841, RUF005 +# ruff: noqa: E741, F821, F841, RUF005 import inspect import re @@ -224,12 +224,22 @@ def test_invalid_match_buffer_region(): check_error(invalid_match_buffer_region, 5) -def test_duplicate_buffer(): - def duplicate_buffer() -> None: +def test_buffer_rebinding_preserves_distinct_allocations(): + @T.prim_func(s_tir=True) + def rebound_buffer() -> None: A = T.sblock_alloc_buffer((128, 128), "float32") - A = T.sblock_alloc_buffer((128, 128), "float32") # error - - check_error(duplicate_buffer, 3) + A = T.sblock_alloc_buffer((128, 128), "float32") + A[0, 0] = A[0, 1] + T.float32(1) + + # Python rebinding selects the second buffer and retains both native allocations. + block = rebound_buffer.body.block + assert len(block.alloc_buffers) == 2 + first, second = block.alloc_buffers + assert not first.same_as(second) + store = block.body + assert isinstance(store, tirx.BufferStore) + assert store.buffer.same_as(second) + assert store.value.a.source.same_as(second) def test_duplicate_block_signature(): diff --git a/tests/python/tvmscript/test_tvmscript_printer_annotation.py b/tests/python/tvmscript/test_tvmscript_printer_annotation.py index 7442bd7afc..02fb421b15 100644 --- a/tests/python/tvmscript/test_tvmscript_printer_annotation.py +++ b/tests/python/tvmscript/test_tvmscript_printer_annotation.py @@ -93,13 +93,12 @@ def main(): def test_disable_concise_scoping_when_scope_annotated(): @T.prim_func(s_tir=True) def _func(): - x = 1 - y = x + 1 + x: T.int32 = 1 + y: T.int32 = x + 1 T.evaluate(y - 1) - # In fork, each bare `x = expr` lowers to AllocBuffer + BufferStore (local_scalar); - # the printer fuses each pair into a single `y: T.int32 = x + 1` line. Annotate the - # AllocBuffer that originates this fused line. + # Explicit scalar declarations lower to AllocBuffer + BufferStore (local_scalar). + # The printer fuses each pair into one line; annotate the allocation for y. result = _func.with_attr("global_symbol", "main").script( obj_to_annotate={ _func.body.seq[2]: "annotation 1", diff --git a/tests/python/tvmscript/test_tvmscript_syntax_sugar.py b/tests/python/tvmscript/test_tvmscript_syntax_sugar.py index 57354f72db..006e415012 100644 --- a/tests/python/tvmscript/test_tvmscript_syntax_sugar.py +++ b/tests/python/tvmscript/test_tvmscript_syntax_sugar.py @@ -455,18 +455,27 @@ def test_preserve_parameter_name(): assert param_name == "i" -def test_preserve_variable_name(): [email protected]("mutable", [False, True]) +def test_preserve_variable_name(mutable): """Use variable name when generating tirx::Bind / AllocBuffer""" - @T.prim_func(s_tir=True) - def func(): - for i in T.serial(16): - j = i // 4 - T.evaluate(j) - - # In fork, bare `j = i // 4` lowers to AllocBuffer (local_scalar) in the for-body - # SeqStmt; the variable name lives on the underlying buffer. - var_name = func.body.body.seq[0].buffer.name + # Bare bindings name the immutable Var; explicit declarations name scalar storage. + annotation = ": T.int32" if mutable else "" + func = from_source( + f"""@T.prim_func(s_tir=True) +def func(): + for i in T.serial(16): + j{annotation} = i // 4 + T.evaluate(j) +""" + ) + binding = func.body.body.seq[0] + if mutable: + assert isinstance(binding, tvm.tirx.AllocBuffer) + var_name = binding.buffer.name + else: + assert isinstance(binding, tvm.tirx.Bind) + var_name = binding.var.name assert var_name == "j"
