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 aec1416ffb189220abcece64977e25f1647aa614 Author: Tianqi Chen <[email protected]> AuthorDate: Wed Sep 23 08:43:18 2026 +0000 Resolve declaration aliases and module members in their lexical context --- python/tvm/script/parser/prescan.py | 25 +++++++++---- python/tvm/script/parser/transpile.py | 14 ++++---- tests/parser/test_parser.py | 41 ++++++++++++++++++++++ tests/python/unittest/test_parser_native_frames.py | 7 ++-- 4 files changed, 70 insertions(+), 17 deletions(-) diff --git a/python/tvm/script/parser/prescan.py b/python/tvm/script/parser/prescan.py index 063d088b0b..6d7238d9e5 100644 --- a/python/tvm/script/parser/prescan.py +++ b/python/tvm/script/parser/prescan.py @@ -58,6 +58,8 @@ class Binding: annotation: ast.AST | None = None dtype: object = None direct: bool = False + # Original callee root; assignment dispatch checks completed lexical bindings. + declaration_root: ast.Name | None = None @dataclass(frozen=True) @@ -153,11 +155,13 @@ class PrescanCollector(ast.NodeVisitor): def _error(self, node, message): raise SyntaxError(message, (self.filename, node.lineno, node.col_offset + 1, None)) - def _binding(self, name, node, kind="ordinary", annotation=None, dtype=None): + def _binding( + self, name, node, kind="ordinary", annotation=None, dtype=None, *, declaration_root=None + ): if name in self.namespaces: self._error(node, f"Script namespace {name!r} cannot be rebound or shadowed") self.names.add(name) - item = Binding(name, node, kind, annotation, dtype, self.direct) + item = Binding(name, node, kind, annotation, dtype, self.direct, declaration_root) self.bindings[self.scope].append(item) self.sites[node] = item @@ -241,11 +245,14 @@ class PrescanCollector(ast.NodeVisitor): visit_AsyncFunctionDef = visit_FunctionDef - def _target(self, target, value=None, annotation=None, *, binding_declaration=False): + def _target(self, target, value=None, annotation=None, *, binding_declaration=None): constructor = ( resolve_syntax(value.func, self.environment) if isinstance(value, ast.Call) else None ) - binding_declaration |= getattr(constructor, "__tvm_binding_decl__", False) + if getattr(constructor, "__tvm_binding_decl__", False): + binding_declaration = value.func + while isinstance(binding_declaration, ast.Attribute): + binding_declaration = binding_declaration.value if isinstance(target, ast.Name): declaration = getattr(constructor, "__tvm_type_var_decl__", None) if declaration is not None and not value.args and not value.keywords: @@ -266,8 +273,14 @@ class PrescanCollector(ast.NodeVisitor): ) ): self._binding(target.id, target, "mutable", annotation) - elif binding_declaration: - self._binding(target.id, target, "binding_declaration", annotation) + elif binding_declaration is not None: + self._binding( + target.id, + target, + "binding_declaration", + annotation, + declaration_root=binding_declaration, + ) elif isinstance(value, ast.Name) and value.id == self.module_name: self._binding(target.id, target, "module_alias") else: diff --git a/python/tvm/script/parser/transpile.py b/python/tvm/script/parser/transpile.py index 59b8597a08..cfd4733438 100644 --- a/python/tvm/script/parser/transpile.py +++ b/python/tvm/script/parser/transpile.py @@ -246,13 +246,8 @@ class IRBuilderTranspiler(ast.NodeTransformer): # child expressions to transform or locations to invent. if getattr(node, "_tvm_intrinsic", False): return node - # Source: Module.f; Builder: f (the same reserved native GlobalVar). - if ( - isinstance(node.value, ast.Name) - and node.value.id == self.module_name - and node.attr in self.module_functions - ): - return self._at(ast.copy_location(ast.Name(node.attr, node.ctx), node), node) + # Source: Module.f; Builder: Module.f, using the native module's map. + # Retaining the owner keeps a local f from shadowing this GlobalVar. result = self.generic_visit(node) return ( self._at(result, node) @@ -536,6 +531,9 @@ class IRBuilderTranspiler(ast.NodeTransformer): if isinstance(target, ast.Name): site = self.prescan.sites.get(target) if self.prescan else None kind = site.kind if site is not None else "ordinary" + binding_declaration = kind == "binding_declaration" and ( + self._resolve(site.declaration_root) is not None + ) mutable = self.prescan.mutable_names.get(self.current_scope, ()) if self.prescan else () keywords = {"name": ast.Constant(target.id), "name_span": self.span(target)} if ty is not None: @@ -553,7 +551,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 kind != "binding_declaration" and not frame_value: + elif target.id in mutable and not binding_declaration and not frame_value: # Source: x = value; Builder: X.set_mutable_var_(x, value). return [ ast.copy_location( diff --git a/tests/parser/test_parser.py b/tests/parser/test_parser.py index 1ae698747c..56f8c26da9 100644 --- a/tests/parser/test_parser.py +++ b/tests/parser/test_parser.py @@ -402,6 +402,47 @@ def main(): ] [email protected]("scope", ["local", "parameter", "enclosing"]) +def test_shadowed_binding_declaration_alias_uses_ordinary_assignment(language, scope): + # Before: axis_alias = ordinary; cell = X.cell(); cell = axis_alias() + # Expected builder program: X.set_mutable_var_(cell, axis_alias()). + # An ambient declaration policy cannot override a lexical callable binding. + marker, calls = object(), [] + + @protocol.register_binding_decl + def axis_alias(): + raise AssertionError("The shadowed ambient declaration must not run") + + def ordinary(): + calls.append("ordinary") + return marker + + body = "cell = X.cell()\n cell = axis_alias()\n X.record(cell)" + if scope == "local": + source = "@X.script\ndef main():\n axis_alias = ordinary\n " + body + elif scope == "parameter": + # A callable parameter is an opaque value supplied by the fake frame. + def arg(name, annotation, **kwargs): + frame = language.frame() + frame.params.append(ordinary) + frame.function.params.append(ordinary) + return ordinary + + language.X.arg = arg + source = "@X.script\ndef main(axis_alias: X.tensor(())):\n " + body + else: + source = ( + "@X.script\ndef main():\n axis_alias = ordinary\n" + " @X.script\n def inner():\n " + body.replace("\n", "\n ") + ) + language.parse(source, axis_alias=axis_alias, ordinary=ordinary) + declaration = next(event[2] for event in language.events if event[0] == "declare") + stores = [event for event in language.events if event[0] == "set"] + assert calls == ["ordinary"] + assert len(stores) == 1 and stores[0][1] is declaration and stores[0][2] is marker + assert next(event[1] for event in language.events if event[0] == "record") is declaration + + 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/unittest/test_parser_native_frames.py b/tests/python/unittest/test_parser_native_frames.py index 389326e8e2..cc92e890b5 100644 --- a/tests/python/unittest/test_parser_native_frames.py +++ b/tests/python/unittest/test_parser_native_frames.py @@ -191,12 +191,13 @@ def test_module_alias_is_the_native_frame_and_lookup_keeps_reference(dialect): @pytest.mark.parametrize("alias", [False, True]) [email protected]("parameter", ["x", "callee"]) @pytest.mark.parametrize( "dialect, decorator, annotation", [(T, "T.prim_func", "T.int32"), (R, "R.function", 'R.Tensor((4,), "float32")')], ) def test_module_member_calls_use_the_callers_native_dialect( - monkeypatch, dialect, decorator, annotation, alias + monkeypatch, dialect, decorator, annotation, alias, parameter ): # Before: cls = Module; return cls.callee(x) # Expected builder program: @@ -215,9 +216,9 @@ def test_module_member_calls_use_the_callers_native_dialect( @I.ir_module class Module: @{decorator} - def caller(x: {annotation}) -> {annotation}: + def caller({parameter}: {annotation}) -> {annotation}: {setup} - return {owner}.callee(x) + return {owner}.callee({parameter}) @{decorator} def callee(y: {annotation}) -> {annotation}: return y
