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 c067f3967cf8267e7b54fbef43a1acc7e73a11b4
Author: Tianqi Chen <[email protected]>
AuthorDate: Wed Sep 23 04:16:24 2026 +0000

    Preserve lexical scopes and source locations in direct parser lowering
---
 python/tvm/script/parser/__init__.py               |   5 +
 python/tvm/script/parser/entry.py                  |   7 +-
 python/tvm/script/parser/prescan.py                |  82 +++++---
 python/tvm/script/parser/transpile.py              | 211 +++++++++++++++------
 tests/parser/dummy_builder.py                      |   7 +
 tests/parser/test_parser.py                        | 154 +++++++++++++++
 tests/python/relax/test_tvmscript_parser.py        |   4 +-
 tests/python/tirx/test_parser_call_scopes.py       |   2 +-
 tests/python/tirx/test_parser_constexpr.py         |   2 +-
 tests/python/tirx/test_parser_emitted_spans.py     |   2 +-
 tests/python/tirx/test_parser_scope_id_identity.py |   4 +-
 tests/python/tirx/test_parser_scope_id_names.py    |   2 +-
 tests/python/unittest/test_parser_native_frames.py |  86 +++++++++
 13 files changed, 471 insertions(+), 97 deletions(-)

diff --git a/python/tvm/script/parser/__init__.py 
b/python/tvm/script/parser/__init__.py
index 121096e6a7..90e22e6b0a 100644
--- a/python/tvm/script/parser/__init__.py
+++ b/python/tvm/script/parser/__init__.py
@@ -74,6 +74,11 @@ def _initialize():
         ir.ir_module = entry.ir_module
         ir.pyfunc = entry.pyfunc
         for namespace in (tir_namespace, relax_namespace, ir):
+            # Materialize declared public entry points after builder bootstrap.
+            # Source-text decorators need the same registered metadata as 
Python
+            # decorator lookup, including entries owned lazily by a dialect.
+            for name in vars(namespace).get("_ENTRY_EXPORTS", ()):
+                getattr(namespace, name)
             namespace.__all__ = sorted(
                 {
                     *vars(namespace).get("__all__", ()),
diff --git a/python/tvm/script/parser/entry.py 
b/python/tvm/script/parser/entry.py
index 4f7212e8b4..ad8afc4dd2 100644
--- a/python/tvm/script/parser/entry.py
+++ b/python/tvm/script/parser/entry.py
@@ -528,7 +528,7 @@ def _prepare_transpiler(
     # Give prescan its registered syntax policy before allocating injected 
names.
     if inspect.isfunction(source):
         tree.body[-1]._tvm_function_info = 
syntax_protocol.function_info(source)
-    prescan = PrescanCollector(metadata).collect(tree)
+    prescan = PrescanCollector(metadata, filename=filename).collect(tree)
     names = dict.fromkeys([*namespace, *prescan.reserved_names], 0)
 
     def fresh(prefix="_t"):
@@ -590,7 +590,6 @@ def _prepare_transpiler(
         infrastructure_name,
         span,
         fresh,
-        name_map=names,
         track_span=track_span,
         definition_scopes_name=definition_scopes_name,
         prescan=prescan,
@@ -666,9 +665,7 @@ def _build(tree, source, environment, definition_scope, 
filename, flags, *, trac
     transformer, namespace = _prepare_transpiler(
         tree, source, environment, definition_scope, filename, 
track_span=track_span
     )
-    original_name = transformer.fresh("_original")
-    namespace[original_name] = source if inspect.isclass(source) else None
-    transformed, result = transformer.program(tree, original_name, namespace)
+    transformed, result = transformer.program(tree)
     runnable = recompose_builder(
         transformed,
         source_fn=source,
diff --git a/python/tvm/script/parser/prescan.py 
b/python/tvm/script/parser/prescan.py
index 1b7bd2987b..aa5004a50f 100644
--- a/python/tvm/script/parser/prescan.py
+++ b/python/tvm/script/parser/prescan.py
@@ -79,14 +79,17 @@ class PrescanContext:
     namespaces: frozenset
     # Explicit source dataflow outputs preserve native export identity on exit.
     with_outputs: object
+    # Only source functions referencing themselves need early standalone refs.
+    recursive_functions: frozenset
 
 
 class PrescanCollector(ast.NodeVisitor):
     """Collect binding syntax once, then discard all traversal accumulators."""
 
-    def __init__(self, environment):
+    def __init__(self, environment, *, filename="<str>"):
         # Fixed lookup inputs last for this scan; never updated by assignments.
         self.environment = environment
+        self.filename = filename
         # Accumulators are frozen by collect(); no native values are stored.
         self.names = set(environment)
         self.bindings = {}
@@ -94,6 +97,10 @@ class PrescanCollector(ast.NodeVisitor):
         self.outputs = {}
         self.namespaces = set()
         self.exports = {}
+        self.recursive = set()
+        # Active source declarations, used only to recognize self references.
+        self.functions = []
+        self.module_name = None
         # Temporary lexical with-stack routes explicit output calls, then 
resets.
         self.regions = []
         # Lexical scope and dialect restore on function/class exit. direct 
marks
@@ -140,11 +147,15 @@ class PrescanCollector(ast.NodeVisitor):
             MappingProxyType(dict(self.outputs)),
             frozenset(self.namespaces),
             MappingProxyType({node: tuple(names) for node, names in 
self.exports.items()}),
+            frozenset(self.recursive),
         )
 
+    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):
         if name in self.namespaces:
-            raise SyntaxError(f"Script namespace {name!r} cannot be rebound or 
shadowed")
+            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)
         self.bindings[self.scope].append(item)
@@ -152,33 +163,39 @@ class PrescanCollector(ast.NodeVisitor):
 
     def visit_Name(self, node):
         if isinstance(node.ctx, ast.Store) and node.id in self.namespaces:
-            raise SyntaxError(f"Script namespace {node.id!r} cannot be rebound 
or shadowed")
+            self._error(node, f"Script namespace {node.id!r} cannot be rebound 
or shadowed")
         self.names.add(node.id)
+        if isinstance(node.ctx, ast.Load):
+            for function in reversed(self.functions):
+                if node.id == function.name:
+                    self.recursive.add(function)
+                    break
 
     def visit_arg(self, node):
         if node.arg in self.namespaces:
-            raise SyntaxError(f"Script namespace {node.arg!r} cannot be 
rebound or shadowed")
+            self._error(node, f"Script namespace {node.arg!r} cannot be 
rebound or shadowed")
         self.names.add(node.arg)
         self.generic_visit(node)
 
     def visit_alias(self, node):
         name = node.asname or node.name.split(".")[0]
         if isinstance(self.scope, ast.FunctionDef) and name in self.namespaces:
-            raise SyntaxError(f"Script namespace {name!r} cannot be rebound or 
shadowed")
+            self._error(node, f"Script namespace {name!r} cannot be rebound or 
shadowed")
         self.names.add(name)
 
     def visit_ClassDef(self, node):
         self.names.add(node.name)
-        old = self.scope
-        self.scope = node
+        old, old_module = self.scope, self.module_name
+        self.scope, self.module_name = node, node.name
         self.bindings[node] = []
         self.generic_visit(node)
-        self.scope = old
+        self.scope, self.module_name = old, old_module
 
     def visit_FunctionDef(self, node):
         self._binding(node.name, node, "function")
         old_scope, old_builder, old_direct = self.scope, self.builder, 
self.direct
         self.scope, self.direct = node, True
+        self.functions.append(node)
         original_kind = getattr(node, "_tvm_function_info", None)
         if original_kind is not None:
             self.builder = original_kind.builder
@@ -203,7 +220,8 @@ class PrescanCollector(ast.NodeVisitor):
                 arg.arg,
                 arg,
                 "mutable_parameter"
-                if protocol.is_mutable_var_decl(constructor, 
syntax="parameter")
+                if getattr(self.builder, "supports_mutable_declarations", True)
+                and protocol.is_mutable_var_decl(constructor, 
syntax="parameter")
                 else "parameter",
                 arg.annotation,
                 (declaration.dtype if declaration is not None else dtype)
@@ -218,6 +236,7 @@ class PrescanCollector(ast.NodeVisitor):
             self.visit(decorator)
         if node.returns:
             self.visit(node.returns)
+        self.functions.pop()
         self.scope, self.builder, self.direct = old_scope, old_builder, 
old_direct
 
     visit_AsyncFunctionDef = visit_FunctionDef
@@ -232,17 +251,24 @@ class PrescanCollector(ast.NodeVisitor):
             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)
-            elif protocol.is_mutable_var_decl(constructor, syntax="call") or (
-                annotation is not None
-                and protocol.is_mutable_var_decl(
-                    resolve_syntax(
-                        annotation.value if isinstance(annotation, 
ast.Subscript) else annotation,
-                        self.environment,
-                    ),
-                    syntax="annotation",
+            elif getattr(self.builder, "supports_mutable_declarations", True) 
and (
+                protocol.is_mutable_var_decl(constructor, syntax="call")
+                or (
+                    annotation is not None
+                    and protocol.is_mutable_var_decl(
+                        resolve_syntax(
+                            annotation.value
+                            if isinstance(annotation, ast.Subscript)
+                            else annotation,
+                            self.environment,
+                        ),
+                        syntax="annotation",
+                    )
                 )
             ):
                 self._binding(target.id, target, "mutable", annotation)
+            elif isinstance(value, ast.Name) and value.id == self.module_name:
+                self._binding(target.id, target, "module_alias")
             else:
                 self._binding(target.id, target, annotation=annotation)
         elif isinstance(target, ast.Tuple | ast.List):
@@ -270,6 +296,7 @@ class PrescanCollector(ast.NodeVisitor):
 
     def visit_If(self, node):
         old_direct, self.direct = self.direct, False
+        self.generic_visit(node)
         marker = (
             isinstance(node.test, ast.Call)
             and isinstance(node.test.func, ast.Attribute)
@@ -287,13 +314,24 @@ class PrescanCollector(ast.NodeVisitor):
                     return last.targets[0].id
                 if isinstance(last, ast.AnnAssign) and isinstance(last.target, 
ast.Name):
                     return last.target.id
-                return None
+                return self.outputs.get(last)
 
             then, otherwise = ending(node.body), ending(node.orelse)
-            if then is None or then != otherwise:
-                raise SyntaxError("IR conditional branches must end with the 
same named output")
-            self.outputs[node] = then
-        self.generic_visit(node)
+            # Effect-only branches have no Python output. Native branch frames
+            # still reject a non-void expression used as an effect-only ending.
+            effects = bool(
+                node.body
+                and node.orelse
+                and isinstance(node.body[-1], ast.Expr)
+                and isinstance(node.orelse[-1], ast.Expr)
+            )
+            if not effects:
+                if then is None or then != otherwise:
+                    location = node.orelse[-1] if node.orelse else node
+                    self._error(
+                        location, "IR conditional branches must end with the 
same named output"
+                    )
+                self.outputs[node] = then
         self.direct = old_direct
 
     def visit_For(self, node):
diff --git a/python/tvm/script/parser/transpile.py 
b/python/tvm/script/parser/transpile.py
index 4974825923..ab18c87811 100644
--- a/python/tvm/script/parser/transpile.py
+++ b/python/tvm/script/parser/transpile.py
@@ -34,7 +34,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
     """A single statement/expression visitor over the entry-owned AST.
 
     ``environment``, ``prescan``, ``filename`` and ``span`` are fixed inputs 
for
-    one translation. ``name_map`` is the unit-wide collision-free name 
allocator;
+    one translation. ``fresh`` allocates names in the entry-owned name map;
     ``bindings`` receives only injected host namespaces and helper functions.
     ``dialect_prefix``, ``current_scope``, ``host_expression`` and annotation
     substitutions are saved/restored at their lexical visitor boundaries. They
@@ -48,10 +48,9 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         builder_name,
         infrastructure_name,
         span,
-        fresh=None,
+        fresh,
         *,
         prescan=None,
-        name_map=None,
         track_span=True,
         definition_scopes_name=None,
         current_scope=None,
@@ -67,8 +66,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         self.track_span = track_span
         self.definition_scopes_name = definition_scopes_name
         # Allocation/injection last for this source unit, including nested 
code.
-        self.name_map = name_map if name_map is not None else 
dict.fromkeys(environment, 0)
-        self.fresh = fresh or self.fresh_unique_name
+        self.fresh = fresh
         self.bindings = bindings if bindings is not None else {}
         # Lexical syntax context, restored on nested function/host/annotation 
exit.
         self.dialect_prefix = builder_name
@@ -81,16 +79,6 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         self.module_name = None
         self.module_functions = frozenset()
 
-    def fresh_unique_name(self, prefix="_t"):
-        """Reserve a generated identifier for the lifetime of this 
translation."""
-        counter = self.name_map.get(prefix, 0)
-        while f"{prefix}{counter}" in self.name_map:
-            counter += 1
-        name = f"{prefix}{counter}"
-        self.name_map[prefix] = counter + 1
-        self.name_map[name] = 0
-        return name
-
     def _inject(self, value, prefix="_host"):
         name = self.fresh(prefix)
         self.bindings[name] = value
@@ -177,6 +165,25 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             self.host_expression = old
 
     def _resolve(self, node):
+        # Fixed namespace meanings coexist with Python lexical value bindings.
+        # A local ``range`` or callable hides the ambient binding for the whole
+        # source function, including reads before its assignment.
+        root = node
+        while isinstance(root, ast.Attribute):
+            root = root.value
+        scope = self.current_scope
+        if isinstance(root, ast.Name) and self.prescan is not None:
+            while isinstance(scope, ast.FunctionDef):
+                if any(item.name == root.id for item in 
self.prescan.bindings.get(scope, ())):
+                    return None
+                scope = next(
+                    (
+                        parent
+                        for parent, items in self.prescan.bindings.items()
+                        if any(item.node is scope and item.kind == "function" 
for item in items)
+                    ),
+                    None,
+                )
         return resolve_syntax(node, self.environment)
 
     def _constexpr_operand(self, node):
@@ -228,7 +235,11 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                 ],
                 node,
             )
-        return node
+        return (
+            self._at(node, node)
+            if isinstance(node.ctx, ast.Load) and not self.host_expression
+            else node
+        )
 
     def visit_Attribute(self, node):
         # Source: Module.f; Builder: f (the same reserved native GlobalVar).
@@ -237,8 +248,13 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             and node.value.id == self.module_name
             and node.attr in self.module_functions
         ):
-            return ast.copy_location(ast.Name(node.attr, node.ctx), node)
-        return self.generic_visit(node)
+            return self._at(ast.copy_location(ast.Name(node.attr, node.ctx), 
node), node)
+        result = self.generic_visit(node)
+        return (
+            self._at(result, node)
+            if isinstance(node.ctx, ast.Load) and not self.host_expression
+            else result
+        )
 
     def visit_Lambda(self, node):
         # Source: lambda n: n + outer; Builder: preserve the lambda's locals
@@ -280,7 +296,51 @@ class IRBuilderTranspiler(ast.NodeTransformer):
     visit_DictComp = visit_ListComp
     visit_GeneratorExp = visit_ListComp
 
-    def visit_Call(self, node):
+    def visit_Constant(self, node):
+        # Source: literal; Builder: I.at_(literal_loc, literal).
+        return node if self.host_expression else self._at(node, node)
+
+    def visit_List(self, node):
+        # Source: [a,b] / (a,b) / {a,b}; Builder: preserve the Python container
+        # and source-location every original child through this same visitor.
+        result = self.generic_visit(node)
+        return result if self.host_expression else self._at(result, node)
+
+    visit_Tuple = visit_List
+    visit_Set = visit_List
+    visit_Dict = visit_List
+
+    def visit_JoinedStr(self, node):
+        # Source: f"value={expr}"; Builder: keep literal fragments, visit expr.
+        # Python requires these fragments to remain Constant/FormattedValue.
+        for child in node.values:
+            if isinstance(child, ast.FormattedValue):
+                child.value = self.visit(child.value)
+                if child.format_spec is not None:
+                    child.format_spec = self.visit_JoinedStr(child.format_spec)
+        return node
+
+    def visit_Subscript(self, node):
+        # Source: buffer[index]; Builder: I.at_(load_loc, buffer[index]).
+        result = self.generic_visit(node)
+        if not self.host_expression and isinstance(node.ctx, ast.Load):
+            return self._at(result, node)
+        return result
+
+    def _module_owner(self, node):
+        """Recognize fixed source module aliases from existing binding 
records."""
+        if not isinstance(node, ast.Name):
+            return False
+        if node.id == self.module_name:
+            return True
+        records = [
+            item
+            for item in self.prescan.bindings.get(self.current_scope, ())
+            if item.name == node.id
+        ]
+        return bool(records) and all(item.kind == "module_alias" for item in 
records)
+
+    def visit_Call(self, node, *, callee=None):
         # Source: X.Tensor(("n",), vdevice="cuda:0")
         # Builder: X.Tensor((X.resolve_type_var_("n"),),
         #                   vdevice=I.resolve_global_info("cuda:0"))
@@ -297,11 +357,18 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             isinstance(node.func, ast.Name) and node.func.id in 
self.module_functions
         ) or (
             isinstance(node.func, ast.Attribute)
-            and isinstance(node.func.value, ast.Name)
-            and node.func.value.id == self.module_name
+            and self._module_owner(node.func.value)
             and node.func.attr in self.module_functions
         )
-        node = self.generic_visit(node)
+        # Generated callees are already builder syntax. Only their original
+        # arguments need the main visitor; do not give injected names fake 
spans.
+        if callee is not None:
+            node.func = callee
+        elif not generated:
+            node.func = self.visit(node.func)
+        node.args = [self.visit(value) for value in node.args]
+        for keyword in node.keywords:
+            keyword.value = self.visit(keyword.value)
         # Source: declared_global(x, y)
         # Builder: X.call_global_var_(declared_global, [x, y])
         if global_call and not self.host_expression:
@@ -773,12 +840,17 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             return self.generic_visit(node)
         if node.orelse:
             self._error(node, "A construction loop does not support an else 
clause")
-        iterable = node.iter
-        if isinstance(iterable, ast.Call) and self._resolve(iterable.func) is 
range:
-            iterable.func = ast.copy_location(
-                ast.Attribute(ast.Name(self.dialect_prefix, ast.Load()), 
"range_", ast.Load()),
-                iterable.func,
+        if isinstance(node.iter, ast.Call) and self._resolve(node.iter.func) 
is range:
+            # Normalize only the known builtin; a lexical range binding follows
+            # the ordinary source-call path. Arguments keep that one call 
scope.
+            iterable = self.visit_Call(
+                node.iter,
+                callee=ast.Attribute(
+                    ast.Name(self.dialect_prefix, ast.Load()), "range_", 
ast.Load()
+                ),
             )
+        else:
+            iterable = self.visit(node.iter)
         if isinstance(node.target, ast.Name):
             names = ast.Constant(node.target.id)
         elif isinstance(node.target, ast.Tuple | ast.List):
@@ -791,7 +863,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             )
         else:
             self._error(node.target, "Loop targets must be names or a flat 
tuple of names")
-        context = self._operation("for_", [self.visit(iterable)], node, 
names=names)
+        context = self._operation("for_", [iterable], node, names=names)
         body = self.transform_statements(node.body)
         return ast.copy_location(
             ast.With([ast.withitem(context, node.target)], body or 
[ast.Pass()]), node
@@ -913,7 +985,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             return protocol.FunctionDecoratorInfo(None, python=True), 
ast.Dict([], [])
         self._error(node, f"Function {node.name!r} has no registered 
construction kind")
 
-    def function_program(self, node, *, local=False):
+    def function_program(self, node, *, local=False, declare=True):
         """Declare one native frame and emit one body function taking that 
frame.
 
         Annotation aliases retain actual definition-local Python values; source
@@ -1045,19 +1117,22 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             )
         for item in facts:
             if item.direct and item.dtype is not None and item.kind in 
("symbol", "parameter"):
-                declaration.append(
-                    ast.copy_location(
-                        ast.Expr(
-                            self._operation(
-                                "resolve_type_var_",
-                                [ast.Constant(item.name)],
-                                item.node,
-                                dtype=ast.Constant(item.dtype),
-                            )
-                        ),
-                        item.node,
-                    )
+                symbol = self._operation(
+                    "resolve_type_var_",
+                    [ast.Constant(item.name)],
+                    item.node,
+                    dtype=ast.Constant(item.dtype),
                 )
+                if item.kind == "parameter":
+                    alias = self.fresh("_symbol")
+                    symbol_aliases[item.name] = alias
+                    declaration.append(
+                        ast.copy_location(
+                            ast.Assign([ast.Name(alias, ast.Store())], 
symbol), item.node
+                        )
+                    )
+                else:
+                    declaration.append(ast.copy_location(ast.Expr(symbol), 
item.node))
         self.annotation_aliases = {**aliases, **symbol_aliases}
         constexpr_aliases = {}
         ordered = sorted(
@@ -1142,7 +1217,9 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         # Builder: with X.function(decl=True) as fn: signature
         with self._host():
             options = self.visit(options)
-        keywords = [ast.keyword(None, options), ast.keyword("decl", 
ast.Constant(True))]
+        keywords = [ast.keyword(None, options)]
+        if declare:
+            keywords.append(ast.keyword("decl", ast.Constant(True)))
         if local:
             keywords.append(ast.keyword("local", ast.Constant(True)))
         if self.track_span:
@@ -1153,21 +1230,21 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             ),
             node,
         )
-        statements.append(
-            ast.copy_location(
-                ast.With([ast.withitem(constructor, ast.Name(frame, 
ast.Store()))], declaration),
-                node,
+        if declare:
+            declaration_scope = ast.With(
+                [ast.withitem(constructor, ast.Name(frame, ast.Store()))], 
declaration
             )
-        )
-        statements.append(
-            ast.copy_location(
-                ast.Assign(
-                    [ast.Name(node.name, ast.Store())],
-                    ast.Attribute(ast.Name(frame, ast.Load()), "reference", 
ast.Load()),
-                ),
-                node,
+            statements.append(ast.copy_location(declaration_scope, node))
+            reference = ast.Attribute(ast.Name(frame, ast.Load()), 
"reference", ast.Load())
+            statements.append(
+                ast.copy_location(ast.Assign([ast.Name(node.name, 
ast.Store())], reference), node)
+            )
+        else:
+            # Source: a standalone nonrecursive function.
+            # Builder: fn = X.function(); build(fn) enters once for 
signature+body.
+            statements.append(
+                ast.copy_location(ast.Assign([ast.Name(frame, ast.Store())], 
constructor), node)
             )
-        )
         # The body re-enters the exact native frame. Runtime parameter storage
         # stays native; the short iterator is consumed once by source 
parameters.
         iterator = self.fresh("_arguments")
@@ -1237,7 +1314,11 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         body.extend(self.transform_statements(node.body))
         self.body_annotation_aliases = old_body_annotations
         resumed = ast.copy_location(
-            ast.With([ast.withitem(ast.Name(frame, ast.Load()))], body), node
+            ast.With(
+                [ast.withitem(ast.Name(frame, ast.Load()))],
+                body if declare else [*declaration, *body],
+            ),
+            node,
         )
         definition = self._definition(body_name, [frame], [resumed], node)
         if not local:
@@ -1251,11 +1332,10 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         self.dialect_prefix, self.current_scope, self.annotation_aliases = old
         return statements, frame, body_name
 
-    def program(self, tree, original_name, bindings):
+    def program(self, tree):
         """Emit direct native module construction, declarations, then 
bodies."""
         from .entry import make_opaque_function
 
-        self.bindings = bindings
         root, prefix = tree.body[-1], tree.body[:-1]
         if not isinstance(root, ast.ClassDef | ast.FunctionDef):
             self._error(root, "Source must contain one function or module 
class")
@@ -1266,9 +1346,10 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             self._error(root, "Duplicate function declaration")
         self.module_name = root.name if is_module else None
         self.module_functions = frozenset(item.name for item in functions)
+        declare = is_module or root in self.prescan.recursive_functions
         builder, result = self.fresh("_builder"), self.fresh("_result")
         body = []
-        for function in functions:
+        for function in functions if declare else ():
             body.append(
                 ast.copy_location(
                     ast.Assign(
@@ -1371,7 +1452,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                 )
                 python_functions.append((function.name, host))
                 continue
-            declaration, frame, callback = self.function_program(function)
+            declaration, frame, callback = self.function_program(function, 
declare=declare)
             body.extend(declaration)
             definitions.append(
                 ast.copy_location(
@@ -1390,7 +1471,13 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         #   result = builder.get()
         module = ast.copy_location(
             ast.With(
-                [ast.withitem(self._call(self.infrastructure_name, 
"ir_module", [], root))], body
+                [
+                    ast.withitem(
+                        self._call(self.infrastructure_name, "ir_module", [], 
root),
+                        ast.Name(root.name, ast.Store()) if is_module else 
None,
+                    )
+                ],
+                body,
             ),
             root,
         )
@@ -1402,7 +1489,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                         ast.Name(builder, ast.Store()),
                     )
                 ],
-                [module],
+                [module] if declare else body,
             ),
             root,
         )
diff --git a/tests/parser/dummy_builder.py b/tests/parser/dummy_builder.py
index 7e4f25cb58..87823e10ae 100644
--- a/tests/parser/dummy_builder.py
+++ b/tests/parser/dummy_builder.py
@@ -113,6 +113,11 @@ class Frame:
     def __getitem__(self, name):
         return self.language.functions[name]
 
+    def __getattr__(self, name):
+        if self.kind == "module" and name in self.language.references:
+            return self.language.references[name]
+        raise AttributeError(name)
+
 
 class Language:
     """One recording builder namespace X and shared infrastructure namespace 
I."""
@@ -131,12 +136,14 @@ class Language:
             with_at_group_=self.with_at_group,
             resolve_global_info=self.resolve_global_info,
             reserve_function=self.reserve_function,
+            module_member_=lambda name, value: value,
             require_defined=self.require_defined,
             annotation_value_=lambda name, value: value,
             MISSING=self.missing,
             constexpr=protocol.constexpr,
         )
         self.X = SimpleNamespace(
+            supports_mutable_declarations=True,
             function=lambda **kwargs: Frame(self, "function", **kwargs),
             func_name=self.func_name,
             arg=self.arg,
diff --git a/tests/parser/test_parser.py b/tests/parser/test_parser.py
index e66cebf1fc..b3d1a671dd 100644
--- a/tests/parser/test_parser.py
+++ b/tests/parser/test_parser.py
@@ -532,3 +532,157 @@ def main():
 """,
             marker=protocol.constexpr,
         )
+
+
[email protected]("recursive", [False, True])
+def 
test_source_function_uses_declaration_only_when_reference_is_needed(language, 
recursive):
+    # Before: @X.script def main(x: ...): X.record(x)  (or main(x))
+    # Expected builder program:
+    # ordinary: with X.function(): x = X.arg(...); X.emit_(X.record(x))
+    # recursive: with X.function(decl=True) as fn: X.arg(...)
+    #            with fn: X.emit_(X.call_global_var_(fn.reference, 
[fn.params[0]]))
+    statement = "main(x)" if recursive else "X.record(x)"
+    result = language.parse(f"@X.script\ndef main(x: X.tensor((4,))):\n    
{statement}\n")
+    entries = [event for event in language.events if event[:2] == ("enter", 
"function")]
+    assert [event[2] for event in entries] == ([True, False] if recursive else 
[False])
+    assert len([event for event in language.events if event[0] == "arg"]) == 1
+    if recursive:
+        assert entries[0][3] is entries[1][3]
+        call = result.body[0][1]
+        assert call.op == "call" and call.args[0] is 
language.references["main"]
+        assert call.args[1] is result.params[0]
+    else:
+        assert result.body[0][1] is result.params[0]
+
+
[email protected]("alias", [False, True])
+def test_module_alias_keeps_frame_identity_and_caller_dialect(language, alias):
+    # Before: cls = Module; cls.callee(x)
+    # Expected builder program:
+    # with I.ir_module() as Module: ...
+    # cls = X.bind_(Module, name="cls"); 
X.emit_(X.call_global_var_(cls.callee, [x]))
+    setup, owner = ("cls = Module", "cls") if alias else ("pass", "Module")
+    result = language.parse(f"""
[email protected]_module
+class Module:
+    @X.script
+    def caller(x: X.tensor((4,))):
+        {setup}
+        {owner}.callee(x)
+    @X.script
+    def callee(y: X.tensor((4,))):
+        X.record(y)
+""")
+    call = result["caller"].body[0][1]
+    assert call.op == "call"
+    assert call.args == (language.references["callee"], 
result["caller"].params[0])
+    modules = [event[3] for event in language.events if event[:2] == ("enter", 
"module")]
+    assert len(modules) == 1
+    if alias:
+        bound = next(event[2] for event in language.events if event[:2] == 
("bind", "cls"))
+        assert bound is modules[0]
+        assert bound.callee is language.references["callee"]
+
+
[email protected](
+    "body, line, column, message",
+    [
+        ("    X = 1\n", 3, 5, "namespace"),
+        ("    for X in range(2):\n        pass\n", 3, 9, "namespace"),
+        (
+            "    if condition:\n        y = left\n    else:\n        z = 
right\n",
+            6,
+            9,
+            "same named output",
+        ),
+    ],
+)
+def test_syntax_diagnostics_point_to_the_offending_binding(language, body, 
line, column, message):
+    # Before: X = 1; or if condition: y = left; else: z = right
+    # Expected builder program: reject the offending namespace/output binding
+    # at its original filename, line and column, before running the builder.
+    language.X.__tvm_value_if__ = True
+    with pytest.raises(Exception, match=message) as error:
+        language.parse("@X.script\ndef main():\n" + body)
+    cause = error.value.__cause__
+    assert isinstance(cause, SyntaxError)
+    assert (cause.filename, cause.lineno, cause.offset) == ("dummy.py", line, 
column)
+    assert f"dummy.py:{line}: SyntaxError:" in str(error.value)
+    assert f" {line} | {body.splitlines()[line - 3]}" in str(error.value)
+    assert not language.events
+
+
+def test_void_branch_statements_need_no_synthetic_named_output(language):
+    # Before: if condition: X.record(None); else: X.record(None)
+    # Expected builder program:
+    # with X.If(condition):
+    #     with X.Then(): X.emit_(X.record(None))
+    #     with X.Else(): X.emit_(X.record(None))
+    language.X.__tvm_value_if__ = True
+    result = language.parse(
+        """
[email protected]
+def main():
+    if condition:
+        X.record(None)
+    else:
+        X.record(None)
+    X.record(9)
+""",
+        condition=Value("condition"),
+    )
+    assert [value for kind, value in result.body] == [None, None, 9]
+    assert not any(event[0] == "bind" for event in language.events)
+
+
+def test_lexical_range_binding_calls_the_custom_iterator_once(language):
+    # Before: range = custom_range; for i in range(4): X.record(i)
+    # Expected builder program: range = X.bind_(custom_range, name="range")
+    # with X.for_(range(4), names="i") as i: X.emit_(X.record(i))
+    calls = []
+
+    def custom_range(extent):
+        calls.append(extent)
+        return language.X.grid(2)
+
+    result = language.parse(
+        "@X.script\ndef main():\n    range = custom_range\n"
+        "    for i in range(4):\n        X.record(i)\n",
+        custom_range=custom_range,
+    )
+    assert calls == [4]
+    variable = result.body[0][1]
+    assert variable.op == "loop" and variable.args == (2,) and variable.name 
== "i"
+
+
[email protected]("expression", ["value", "holder.item", "items[0]", 
"7"])
+def test_non_call_expression_reads_keep_their_source_range(language, 
expression):
+    # Before: value; holder.item; items[0]; 7
+    # Expected builder program: X.emit_(I.at_(expression_loc, expression)).
+    # Native expression identity survives; host constants stay ordinary values.
+    value, reads, located = Value("read"), [], []
+
+    class Holder:
+        @property
+        def item(self):
+            reads.append("attribute")
+            return value
+
+        def __getitem__(self, index):
+            reads.append(index)
+            return value
+
+    def at(location, result):
+        located.append(((str(location[0].name), *location[1:]), result))
+        return language.at(location, result)
+
+    language.I.at_ = at
+    holder = Holder()
+    result = language.parse(
+        f"@X.script\ndef main():\n    {expression}\n", value=value, 
holder=holder, items=holder
+    )
+    expected = 7 if expression == "7" else value
+    location = ("dummy.py", 3, 3, 5, 5 + len(expression))
+    assert result.body[0][1] is expected
+    assert [(loc, result) for loc, result in located if loc == location] == 
[(location, expected)]
+    assert reads == {"holder.item": ["attribute"], "items[0]": 
[0]}.get(expression, [])
diff --git a/tests/python/relax/test_tvmscript_parser.py 
b/tests/python/relax/test_tvmscript_parser.py
index 9cda2610d3..875d2619e6 100644
--- a/tests/python/relax/test_tvmscript_parser.py
+++ b/tests/python/relax/test_tvmscript_parser.py
@@ -1239,8 +1239,8 @@ def test_if_branch_with_match_cast():
     @R.function
     def func(A: R.Tensor([16, 16]), is_bfloat16: R.Prim("bool")):
         if is_bfloat16:
-            A = R.match_cast(A, R.Tensor([16, 16], "bfloat16"))
-            B = A.astype("float16")
+            matched = R.match_cast(A, R.Tensor([16, 16], "bfloat16"))
+            B = matched.astype("float16")
         else:
             B = R.match_cast(A, R.Tensor([16, 16], "float16"))
         return B
diff --git a/tests/python/tirx/test_parser_call_scopes.py 
b/tests/python/tirx/test_parser_call_scopes.py
index 3510c021b9..f111a3af6c 100644
--- a/tests/python/tirx/test_parser_call_scopes.py
+++ b/tests/python/tirx/test_parser_call_scopes.py
@@ -20,9 +20,9 @@ import pytest
 
 from tvm import ir, tirx
 from tvm.script import parser
+from tvm.script import tirx as T
 from tvm.script.ir_builder import IRBuilder
 from tvm.script.ir_builder import ir as I
-from tvm.script import tirx as T
 
 
 def loc(line, name="calls.py"):
diff --git a/tests/python/tirx/test_parser_constexpr.py 
b/tests/python/tirx/test_parser_constexpr.py
index 4e8822b408..2d0ffc958e 100644
--- a/tests/python/tirx/test_parser_constexpr.py
+++ b/tests/python/tirx/test_parser_constexpr.py
@@ -20,8 +20,8 @@ import pytest
 
 from tvm import ir, tirx
 from tvm.script import parser
-from tvm.script.ir_builder import ir as I
 from tvm.script import tirx as T
+from tvm.script.ir_builder import ir as I
 
 
 @pytest.mark.parametrize("marker", ["I.constexpr", "T.constexpr"])
diff --git a/tests/python/tirx/test_parser_emitted_spans.py 
b/tests/python/tirx/test_parser_emitted_spans.py
index 1d07c9366b..9996a7256c 100644
--- a/tests/python/tirx/test_parser_emitted_spans.py
+++ b/tests/python/tirx/test_parser_emitted_spans.py
@@ -19,10 +19,10 @@
 import pytest
 
 from tvm import ir, tirx
+from tvm.script import tirx as T
 from tvm.script.ir_builder import IRBuilder
 from tvm.script.ir_builder import ir as I
 from tvm.script.ir_builder.base import BypassEmit
-from tvm.script import tirx as T
 from tvm.tirx.script.builder import ir as native
 
 
diff --git a/tests/python/tirx/test_parser_scope_id_identity.py 
b/tests/python/tirx/test_parser_scope_id_identity.py
index bf02752e3b..d4815f18df 100644
--- a/tests/python/tirx/test_parser_scope_id_identity.py
+++ b/tests/python/tirx/test_parser_scope_id_identity.py
@@ -20,10 +20,10 @@ import pytest
 
 import tvm
 from tvm.script import parser
+from tvm.script import tirx as T
 from tvm.script.ir_builder import IRBuilder
-from tvm.script.ir_builder.base import BypassBind
 from tvm.script.ir_builder import ir as I
-from tvm.script import tirx as T
+from tvm.script.ir_builder.base import BypassBind
 
 
 @pytest.mark.parametrize(
diff --git a/tests/python/tirx/test_parser_scope_id_names.py 
b/tests/python/tirx/test_parser_scope_id_names.py
index 327ea41a8f..027f7c60cb 100644
--- a/tests/python/tirx/test_parser_scope_id_names.py
+++ b/tests/python/tirx/test_parser_scope_id_names.py
@@ -20,9 +20,9 @@ import pytest
 
 import tvm
 from tvm.script import parser
+from tvm.script import tirx as T
 from tvm.script.ir_builder import IRBuilder
 from tvm.script.ir_builder.base import BypassBind
-from tvm.script import tirx as T
 
 
 @pytest.mark.parametrize(
diff --git a/tests/python/unittest/test_parser_native_frames.py 
b/tests/python/unittest/test_parser_native_frames.py
index ae38929e00..389326e8e2 100644
--- a/tests/python/unittest/test_parser_native_frames.py
+++ b/tests/python/unittest/test_parser_native_frames.py
@@ -162,3 +162,89 @@ def 
test_body_attribute_overrides_default_after_declaration():
             T.func_attr({"global_symbol": "custom"})
             T.evaluate(0)
     assert builder.get()["main"].attrs["global_symbol"] == "custom"
+
+
[email protected]("dialect", [T, R])
+def test_module_alias_is_the_native_frame_and_lookup_keeps_reference(dialect):
+    # Before: cls = Module; cls.callee
+    # Expected builder program: cls = X.bind_(Module, name="cls")
+    # The binding returns the same module frame and its plain native GlobalVar.
+    with IRBuilder() as builder, I.ir_module() as module:
+        with dialect.function(decl=True) as function:
+            dialect.func_name("callee")
+            annotation = T.int32() if dialect is T else R.Tensor((4,), 
"float32")
+            argument = dialect.arg("x", annotation)
+        depth = len(builder.frames)
+        alias = dialect.bind_(module, name="cls")
+        assert alias is module
+        assert len(builder.frames) == depth
+        assert alias.callee.same_as(function.reference)
+        assert isinstance(alias.callee, ir.GlobalVar)
+        with pytest.raises(AttributeError, match="missing"):
+            alias.missing
+        with function:
+            if dialect is T:
+                T.evaluate(argument)
+            else:
+                R.func_ret_value(argument)
+    assert builder.get().get_global_var("callee").same_as(function.reference)
+
+
[email protected]("alias", [False, True])
[email protected](
+    "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
+):
+    # Before: cls = Module; return cls.callee(x)
+    # Expected builder program:
+    # cls = X.bind_(Module, name="cls")
+    # X.return_(X.call_global_var_(cls.callee, [x]))
+    calls = []
+    original = dialect.call_global_var_
+
+    def call(reference, arguments):
+        calls.append((reference, arguments))
+        return original(reference, arguments)
+
+    monkeypatch.setattr(dialect, "call_global_var_", call)
+    setup, owner = ("cls = Module", "cls") if alias else ("pass", "Module")
+    module = parser.parse(f"""
[email protected]_module
+class Module:
+    @{decorator}
+    def caller(x: {annotation}) -> {annotation}:
+        {setup}
+        return {owner}.callee(x)
+    @{decorator}
+    def callee(y: {annotation}) -> {annotation}:
+        return y
+""")
+    assert len(calls) == 1
+    reference, arguments = calls[0]
+    assert reference.same_as(module.get_global_var("callee"))
+    assert len(arguments) == 1
+    assert arguments[0].same_as(module["caller"].params[0])
+
+
+def test_lexical_range_binding_preserves_custom_native_loop_extent():
+    # Before: range = custom_range; for i in range(4): X.evaluate(i)
+    # Expected builder program: range = X.bind_(custom_range, name="range")
+    # with X.for_(range(4), names="i") as i: X.emit_(X.evaluate(i))
+    calls = []
+
+    def custom_range(extent):
+        calls.append(extent)
+        return T.serial(2)
+
+    function = parser.parse(
+        "@T.prim_func\ndef main():\n    range = custom_range\n"
+        "    for i in range(4):\n        T.evaluate(i)\n",
+        extra_vars={"custom_range": custom_range},
+    )
+    assert calls == [4]
+    assert isinstance(function.body, tirx.For)
+    assert int(function.body.extent) == 2
+    assert function.body.body.value.same_as(function.body.loop_var)

Reply via email to