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 db03466a8bbabf96ba56361d2136ebdb1d2b48f0
Author: Tianqi Chen <[email protected]>
AuthorDate: Wed Sep 23 07:31:29 2026 +0000

    Preserve explicit native axis binding declarations
---
 python/tvm/script/parser/prescan.py                | 18 ++---
 python/tvm/script/parser/protocol.py               | 12 ++++
 python/tvm/script/parser/transpile.py              |  2 +-
 python/tvm/tirx/script/builder/protocol.py         |  7 +-
 tests/parser/test_parser.py                        | 82 ++++++++++++++++++++++
 .../python/tvmscript/test_tvmscript_parser_tir.py  | 28 ++++++++
 6 files changed, 139 insertions(+), 10 deletions(-)

diff --git a/python/tvm/script/parser/prescan.py 
b/python/tvm/script/parser/prescan.py
index aa5004a50f..063d088b0b 100644
--- a/python/tvm/script/parser/prescan.py
+++ b/python/tvm/script/parser/prescan.py
@@ -241,13 +241,12 @@ class PrescanCollector(ast.NodeVisitor):
 
     visit_AsyncFunctionDef = visit_FunctionDef
 
-    def _target(self, target, value=None, annotation=None):
+    def _target(self, target, value=None, annotation=None, *, 
binding_declaration=False):
+        constructor = (
+            resolve_syntax(value.func, self.environment) if isinstance(value, 
ast.Call) else None
+        )
+        binding_declaration |= getattr(constructor, "__tvm_binding_decl__", 
False)
         if isinstance(target, ast.Name):
-            constructor = (
-                resolve_syntax(value.func, self.environment)
-                if isinstance(value, ast.Call)
-                else None
-            )
             declaration = getattr(constructor, "__tvm_type_var_decl__", None)
             if declaration is not None and not value.args and not 
value.keywords:
                 self._binding(target.id, target, "symbol", value, 
declaration.dtype)
@@ -267,6 +266,8 @@ class PrescanCollector(ast.NodeVisitor):
                 )
             ):
                 self._binding(target.id, target, "mutable", annotation)
+            elif binding_declaration:
+                self._binding(target.id, target, "binding_declaration", 
annotation)
             elif isinstance(value, ast.Name) and value.id == self.module_name:
                 self._binding(target.id, target, "module_alias")
             else:
@@ -278,9 +279,10 @@ class PrescanCollector(ast.NodeVisitor):
                 else [None] * len(target.elts)
             )
             for child, rhs in zip(target.elts, values):
-                self._target(child, rhs)
+                # One declaration call may return several already-owned values.
+                self._target(child, rhs, 
binding_declaration=binding_declaration)
         elif isinstance(target, ast.Starred):
-            self._target(target.value)
+            self._target(target.value, binding_declaration=binding_declaration)
         self.visit(target)
 
     def visit_Assign(self, node):
diff --git a/python/tvm/script/parser/protocol.py 
b/python/tvm/script/parser/protocol.py
index 1a5760b0e2..6213967141 100644
--- a/python/tvm/script/parser/protocol.py
+++ b/python/tvm/script/parser/protocol.py
@@ -328,6 +328,18 @@ def register_type_var_decl(constructor, *, 
value_parameter="expr", dtype=None):
     return constructor
 
 
+def register_binding_decl(constructor):
+    """Mark a call that explicitly introduces an ordinary source binding.
+
+    Its scalar or unpacked targets bind the returned values even when an outer
+    declaration uses the same name for mutable storage. Builders still own the
+    returned values and their identity; this flag only selects assignment 
syntax.
+    Callable aliases share the metadata, without a separate registry.
+    """
+    constructor.__tvm_binding_decl__ = True
+    return constructor
+
+
 def register_mutable_var_decl(constructor, *, syntax="call"):
     """Register mutable storage in call, annotation or parameter position.
 
diff --git a/python/tvm/script/parser/transpile.py 
b/python/tvm/script/parser/transpile.py
index e3b0aaec4c..59b8597a08 100644
--- a/python/tvm/script/parser/transpile.py
+++ b/python/tvm/script/parser/transpile.py
@@ -553,7 +553,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             elif kind == "mutable" and not frame_value:
                 # Source: x = X.local_scalar(...); Builder: x = 
X.decl_mutable_var_(...).
                 value = self._operation("decl_mutable_var_", [value], 
statement, **keywords)
-            elif target.id in mutable and not frame_value:
+            elif target.id in mutable and kind != "binding_declaration" and 
not frame_value:
                 # Source: x = value; Builder: X.set_mutable_var_(x, value).
                 return [
                     ast.copy_location(
diff --git a/python/tvm/tirx/script/builder/protocol.py 
b/python/tvm/tirx/script/builder/protocol.py
index aa9b3beecf..ad97ee7422 100644
--- a/python/tvm/tirx/script/builder/protocol.py
+++ b/python/tvm/tirx/script/builder/protocol.py
@@ -226,7 +226,12 @@ def set_mutable_var_(target, value, *, span=None):
 
 
 def _register_declarations():
-    from tvm.script.parser.protocol import register_mutable_var_decl
+    from tvm.script.parser.protocol import register_binding_decl, 
register_mutable_var_decl
+
+    # Native axes declare fresh bindings, including unpacked remap results.
+    # Their names can shadow outer storage handles without emitting stores.
+    for name in ("spatial", "reduce", "scan", "opaque", "remap"):
+        register_binding_decl(getattr(_builder.axis, name))
 
     for constructor in vars(_native).values():
         if isinstance(constructor, _native.DtypeConstructor):
diff --git a/tests/parser/test_parser.py b/tests/parser/test_parser.py
index b3d1a671dd..1ae698747c 100644
--- a/tests/parser/test_parser.py
+++ b/tests/parser/test_parser.py
@@ -320,6 +320,88 @@ def main():
     assert mutation[2] is marker
 
 
[email protected]("callee", ["X.axes", "axis_alias"])
[email protected]("target", ["cell", "i, cell", "[i, cell]", "i, (cell, 
*tail)", "i, *cell"])
+def 
test_binding_declarations_override_mutable_targets_and_unpack_once(language, 
callee, target):
+    # Before: cell = X.cell(); i, cell = X.axes()
+    # Expected builder program: cell = X.decl_mutable_var_(X.cell(), 
name="cell")
+    # values = X.axes(); i, cell = X.unpack(values); cell = X.bind_(cell, 
name="cell")
+    # The declaration preserves returned identity rather than storing into the 
old cell.
+    marker, other = object(), object()
+    returned = {
+        "cell": marker,
+        "i, cell": (other, marker),
+        "[i, cell]": [other, marker],
+        "i, (cell, *tail)": (other, (marker, other)),
+        "i, *cell": (other, marker),
+    }[target]
+    calls = []
+
+    @protocol.register_binding_decl
+    def axes():
+        calls.append("axes")
+        return returned
+
+    language.X.axes = axes
+    language.X.unpack = lambda value: value
+    language.parse(
+        f"""
[email protected]
+def main():
+    cell = X.cell()
+    {target} = {callee}()
+    X.record(cell)
+""",
+        axis_alias=axes,
+    )
+    assert calls == ["axes"]
+    assert not any(event[0] == "set" for event in language.events)
+    result = next(event[1] for event in language.events if event[0] == 
"record")
+    if target == "i, *cell":
+        assert len(result) == 1 and result[0] is marker
+    else:
+        assert result is marker
+
+
+def test_ordinary_tuple_and_outer_branch_assignments_still_store(language):
+    # Before: a = X.cell(); b = X.cell(); a, b = values(); if cond: a = first
+    # Expected builder program: declare a/b; unpack values once; set a/b;
+    # with X.Then(): def branch(): X.set_mutable_var_(a, first); branch()
+    first, second = object(), object()
+    calls = []
+
+    def values():
+        calls.append("values")
+        return first, second
+
+    language.X.unpack = lambda value: value
+    language.parse(
+        """
[email protected]
+def main():
+    a = X.cell()
+    b = X.cell()
+    a, b = values()
+    if X.value():
+        a = first
+    else:
+        a = second
+""",
+        values=values,
+        first=first,
+        second=second,
+    )
+    declarations = {event[1]: event[2] for event in language.events if 
event[0] == "declare"}
+    stores = [(event[1], event[2]) for event in language.events if event[0] == 
"set"]
+    assert calls == ["values"]
+    assert stores == [
+        (declarations["a"], first),
+        (declarations["b"], second),
+        (declarations["a"], first),
+        (declarations["a"], second),
+    ]
+
+
 def 
test_quoted_symbols_share_identity_without_introducing_python_bindings(language):
     # Before: def main(x: X.tensor(("n", "n"))): X.record(n)
     # Expected builder program: X.tensor((X.resolve_type_var_("n"), 
X.resolve_type_var_("n")))
diff --git a/tests/python/tvmscript/test_tvmscript_parser_tir.py 
b/tests/python/tvmscript/test_tvmscript_parser_tir.py
index e7ba984464..8b2bbd77f6 100644
--- a/tests/python/tvmscript/test_tvmscript_parser_tir.py
+++ b/tests/python/tvmscript/test_tvmscript_parser_tir.py
@@ -778,6 +778,34 @@ def test_alloc_inside_block():
     tvm.ir.assert_structural_equal(func, expected)
 
 
[email protected](
+    "axes",
+    [
+        'i, k = T.axis.remap("SS", [li, lk])',
+        'i, k = axis_alias("SS", [li, lk])',
+        "i = T.axis.spatial(8, li)\n            k = T.axis.S(8, lk)",
+    ],
+)
+def test_block_axes_shadow_outer_buffer_without_stores(axes):
+    func = tvm.script.from_source(
+        f"""@T.prim_func(s_tir=True)
+def main(buffer: T.handle, output: T.Buffer((8, 8), "float32")):
+    k = T.match_buffer(buffer, (8,), "float32")
+    for li, lk in T.grid(8, 8):
+        with T.sblock("output"):
+            {axes}
+            output[i, k] = T.cast(i + k, "float32")
+""",
+        extra_vars={"axis_alias": T.axis.remap},
+    )
+    block = func.body.block.body.body.body.block
+    store = block.body
+    assert isinstance(store, tirx.BufferStore)
+    assert store.indices[0].same_as(block.iter_vars[0].var)
+    assert store.indices[1].same_as(block.iter_vars[1].var)
+    assert [axis.var.name for axis in block.iter_vars] == ["i", "k"]
+
+
 def test_tir_macro_block_name_suffix():
     @T.inline
     def operation(A, idx):

Reply via email to