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

Reply via email to