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