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 71f6e9a72b46cd7f33214d22e8fbfff90db3a9bc
Author: Tianqi Chen <[email protected]>
AuthorDate: Wed Sep 23 13:51:36 2026 +0000

    [REFACTOR][Script] Separate function emission and restore lexical contexts
    
    Express signature, frame and body construction through typed AST helpers. 
Share assignment and lazy specialization emission, preserve stable constexpr 
ordering, and restore lexical aliases on every exit.
---
 python/tvm/script/parser/transpile.py              | 595 +++++++++++----------
 .../python/script/test_parser_alias_restoration.py |  85 +++
 2 files changed, 404 insertions(+), 276 deletions(-)

diff --git a/python/tvm/script/parser/transpile.py 
b/python/tvm/script/parser/transpile.py
index 30a708ba43..1234bd9a8b 100644
--- a/python/tvm/script/parser/transpile.py
+++ b/python/tvm/script/parser/transpile.py
@@ -31,7 +31,7 @@ from typing import Any, NoReturn, TypeVar
 
 from . import protocol
 from .call_args_policy import handle_call_args_policy, parse_annotation
-from .prescan import PrescanContext, resolve_syntax
+from .prescan import Binding, PrescanContext, resolve_syntax
 
 _Node = TypeVar("_Node", bound=ast.AST)
 
@@ -119,6 +119,11 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             return value
         return self._call(self.infrastructure_name, "at_", [self.span(node), 
value], node)
 
+    @staticmethod
+    def _assign(name: str, value: ast.expr, node: ast.AST) -> ast.Assign:
+        """Assign an injected or source name while retaining its source 
range."""
+        return ast.copy_location(ast.Assign([ast.Name(name, ast.Store())], 
value), node)
+
     @staticmethod
     def _lambda(names: list[str], value: ast.expr) -> ast.Lambda:
         return ast.Lambda(
@@ -174,6 +179,20 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         finally:
             self.host_expression = old
 
+    @contextmanager
+    def _aliases(
+        self, mapping: dict[str, str], *, body: dict[str, str] | None = None
+    ) -> Iterator[None]:
+        """Restore lexical annotation substitutions even when a visitor 
fails."""
+        old = self.annotation_aliases, self.body_annotation_aliases
+        self.annotation_aliases = mapping
+        if body is not None:
+            self.body_annotation_aliases = body
+        try:
+            yield
+        finally:
+            self.annotation_aliases, self.body_annotation_aliases = old
+
     def _resolve(self, node: ast.AST | None) -> object:
         # Fixed namespace meanings coexist with Python lexical value bindings.
         # A local ``range`` or callable hides the ambient binding for the whole
@@ -272,15 +291,17 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         node.args.kw_defaults = [
             self.visit(value) if value is not None else None for value in 
node.args.kw_defaults
         ]
-        old = self.annotation_aliases
         arguments = [*node.args.posonlyargs, *node.args.args, 
*node.args.kwonlyargs]
         arguments += [arg for arg in (node.args.vararg, node.args.kwarg) if 
arg]
         local_names = {argument.arg for argument in arguments}
-        self.annotation_aliases = {
-            name: alias for name, alias in old.items() if name not in 
local_names
-        }
-        node.body = self.visit(node.body)
-        self.annotation_aliases = old
+        with self._aliases(
+            {
+                name: alias
+                for name, alias in self.annotation_aliases.items()
+                if name not in local_names
+            }
+        ):
+            node.body = self.visit(node.body)
         return node
 
     def visit_ListComp(
@@ -288,19 +309,17 @@ class IRBuilderTranspiler(ast.NodeTransformer):
     ) -> ast.ListComp | ast.SetComp | ast.DictComp | ast.GeneratorExp:
         # Source: [f(n) for n in values]; Builder: preserve Python 
comprehension
         # scope while translating its call/operand expressions exactly once.
-        old = self.annotation_aliases
-        self.annotation_aliases = dict(old)
-        for generator in node.generators:
-            generator.iter = self.visit(generator.iter)
-            for target in ast.walk(generator.target):
-                if isinstance(target, ast.Name):
-                    self.annotation_aliases.pop(target.id, None)
-            generator.ifs = [self.visit(value) for value in generator.ifs]
-        if isinstance(node, ast.DictComp):
-            node.key, node.value = self.visit(node.key), self.visit(node.value)
-        else:
-            node.elt = self.visit(node.elt)
-        self.annotation_aliases = old
+        with self._aliases(dict(self.annotation_aliases)):
+            for generator in node.generators:
+                generator.iter = self.visit(generator.iter)
+                for target in ast.walk(generator.target):
+                    if isinstance(target, ast.Name):
+                        self.annotation_aliases.pop(target.id, None)
+                generator.ifs = [self.visit(value) for value in generator.ifs]
+            if isinstance(node, ast.DictComp):
+                node.key, node.value = self.visit(node.key), 
self.visit(node.value)
+            else:
+                node.elt = self.visit(node.elt)
         return node
 
     visit_SetComp = visit_ListComp
@@ -644,11 +663,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             value = ast.Constant(None) if site and site.kind == "symbol" else 
self.visit(node.value)
             return self._bind(target, value, node)
         temporary = self.fresh("_value")
-        result: list[ast.stmt] = [
-            ast.copy_location(
-                ast.Assign([ast.Name(temporary, ast.Store())], 
self.visit(node.value)), node
-            )
-        ]
+        result: list[ast.stmt] = [self._assign(temporary, 
self.visit(node.value), node)]
         for target in node.targets:
             result.extend(self._bind(target, ast.Name(temporary, ast.Load()), 
node))
         return result
@@ -666,10 +681,8 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                 ast.Name(self.infrastructure_name, ast.Load()), "MISSING", 
ast.Load()
             )
         )
-        old_aliases = self.annotation_aliases
-        self.annotation_aliases = self.body_annotation_aliases
-        annotation = self.visit(parse_annotation(node.annotation, 
self.filename))
-        self.annotation_aliases = old_aliases
+        with self._aliases(self.body_annotation_aliases):
+            annotation = self.visit(parse_annotation(node.annotation, 
self.filename))
         return self._bind(node.target, value, node, ty=annotation)
 
     def visit_AugAssign(self, node: ast.AugAssign) -> ast.AugAssign | 
list[ast.stmt]:
@@ -689,31 +702,19 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         if not isinstance(node.target, ast.Subscript | ast.Attribute):
             self._error(node.target, "An augmented assignment requires a name, 
attribute, or index")
         base = self.fresh("_base")
-        statements: list[ast.stmt] = [
-            ast.copy_location(
-                ast.Assign([ast.Name(base, ast.Store())], 
self.visit(node.target.value)), node
-            )
-        ]
+        statements: list[ast.stmt] = [self._assign(base, 
self.visit(node.target.value), node)]
         if isinstance(node.target, ast.Attribute):
             key = ast.Constant(node.target.attr)
             load = ast.Attribute(ast.Name(base, ast.Load()), node.target.attr, 
ast.Load())
             operation = "setattr"
         else:
             index = self.fresh("_index")
-            statements.append(
-                ast.copy_location(
-                    ast.Assign([ast.Name(index, ast.Store())], 
self._index(node.target.slice)), node
-                )
-            )
+            statements.append(self._assign(index, 
self._index(node.target.slice), node))
             key = ast.Name(index, ast.Load())
             load = ast.Subscript(ast.Name(base, ast.Load()), key, ast.Load())
             operation = "setitem"
         old = self.fresh("_old")
-        statements.append(
-            ast.copy_location(
-                ast.Assign([ast.Name(old, ast.Store())], self._at(load, 
node.target)), node
-            )
-        )
+        statements.append(self._assign(old, self._at(load, node.target), node))
         value = self._at(
             ast.BinOp(ast.Name(old, ast.Load()), node.op, 
self.visit(node.value)), node
         )
@@ -847,12 +848,8 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         output = self.prescan.conditional_outputs.get(node) if self.prescan 
else None
         if output is not None:
             result.append(
-                ast.copy_location(
-                    ast.Assign(
-                        [ast.Name(output, ast.Store())],
-                        ast.Attribute(ast.Name(frame, ast.Load()), "var", 
ast.Load()),
-                    ),
-                    node,
+                self._assign(
+                    output, ast.Attribute(ast.Name(frame, ast.Load()), "var", 
ast.Load()), node
                 )
             )
         return result
@@ -937,7 +934,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         # The native frame owns conversion to ordinary output variables.
         frame = self.fresh("_dataflow")
         statements: list[ast.stmt] = [
-            ast.copy_location(ast.Assign([ast.Name(frame, ast.Store())], 
context), node),
+            self._assign(frame, context, node),
             ast.copy_location(
                 ast.With([ast.withitem(ast.Name(frame, ast.Load()), target)], 
translated_body), node
             ),
@@ -948,9 +945,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                 ast.Constant(index),
                 ast.Load(),
             )
-            statements.append(
-                ast.copy_location(ast.Assign([ast.Name(name, ast.Store())], 
value), node)
-            )
+            statements.append(self._assign(name, value, node))
         return statements
 
     def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.FunctionDef | 
list[ast.stmt]:
@@ -1015,32 +1010,12 @@ 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: ast.FunctionDef, *, local: bool = False, declare: bool = 
True
-    ) -> tuple[list[ast.stmt], str, ast.With]:
-        """Declare a native frame and emit a lexical body helper inside its 
scope.
-
-        Annotation aliases retain actual definition-local Python values; source
-        parameters are read from the enclosing frame while its zero-argument
-        helper runs. No factory, callback record, copied parameter map or 
symbol
-        owner is generated.
-        """
-        from . import jit_support
-
-        kind, options = self.function_metadata(node)
-        builder = self.fresh("_X")
-        self.bindings[builder] = kind.builder
-        frame, body_name = self.fresh("_fn"), self.fresh("_build")
-        old = (self.dialect_prefix, self.current_scope, 
self.annotation_aliases)
-        self.dialect_prefix, self.current_scope = builder, node
-        parameters = [*node.args.posonlyargs, *node.args.args, 
*node.args.kwonlyargs]
-        if node.args.vararg or node.args.kwarg:
-            self._error(node, "IR signatures require ordinary named 
parameters")
-        facts = self.prescan.bindings.get(node, ()) if self.prescan else ()
-        type_parameters = list(getattr(node, "type_params", ()))
-        declared_names = {item.name for item in type_parameters}
+    def _function_annotations(
+        self, node: ast.FunctionDef, parameters: list[ast.arg], facts: 
tuple[Binding, ...]
+    ) -> tuple[list[ast.expr | None], ast.expr | None, dict[str, str]]:
+        """Find definition-scope names needed by signatures and body 
annotations."""
+        declared_names = {item.name for item in getattr(node, "type_params", 
())}
         local_names = {item.name for item in facts}
-        # Definition-context aliases are needed only for real annotation names.
         # Quoted expression names are created later by argument normalization.
         annotations = [
             parse_annotation(parameter.annotation, self.filename) if 
parameter.annotation else None
@@ -1071,7 +1046,14 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         aliases = {
             name: self.fresh("_annotation") for name in 
sorted(annotation_names - declared_names)
         }
-        captures = self.fresh("_definition")
+        return annotations, returns, aliases
+
+    def _definition_bindings(
+        self, node: ast.FunctionDef, aliases: dict[str, str], *, captures: 
str, local: bool
+    ) -> list[ast.stmt]:
+        """Capture lexical annotation values without entering a construction 
frame."""
+        # Inject builtin objects under fresh names: a source binding named 
globals,
+        # locals, iter or next must not replace these generated operations.
         captures_expr = ast.Dict(
             [None, None, None],
             [
@@ -1087,11 +1069,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                 ast.Call(self._inject(locals), [], []),
             ],
         )
-        statements: list[ast.stmt] = [
-            ast.copy_location(ast.Assign([ast.Name(captures, ast.Store())], 
captures_expr), node)
-        ]
-        fallback: ast.expr
-        value: ast.expr
+        statements: list[ast.stmt] = [self._assign(captures, captures_expr, 
node)]
         for name, alias in aliases.items():
             fallback = (
                 self._inject(getattr(builtins, name))
@@ -1101,10 +1079,15 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                 )
             )
             value = self._call(captures, "get", [ast.Constant(name), 
fallback], node)
-            statements.append(
-                ast.copy_location(ast.Assign([ast.Name(alias, ast.Store())], 
value), node)
-            )
-        special, absent = self.fresh("_specialization"), self.fresh("_absent")
+            statements.append(self._assign(alias, value, node))
+        return statements
+
+    def _specialization_bindings(
+        self, node: ast.FunctionDef, *, special: str, absent: str, local: bool
+    ) -> list[ast.stmt]:
+        """Read root JIT inputs; nested functions keep ordinary runtime 
parameters."""
+        from . import jit_support
+
         special_expr: ast.expr
         absent_expr: ast.expr
         if local:
@@ -1116,22 +1099,21 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             absent_expr = ast.Call(
                 self._inject(jit_support.absent_parameter_names), 
[ast.Constant(node.name)], []
             )
-        statements.extend(
-            [
-                ast.copy_location(ast.Assign([ast.Name(special, ast.Store())], 
special_expr), node),
-                ast.copy_location(ast.Assign([ast.Name(absent, ast.Store())], 
absent_expr), node),
-            ]
-        )
-        declaration: list[ast.stmt] = [
-            ast.copy_location(
-                ast.Expr(self._call(builder, "func_name", 
[ast.Constant(node.name)], node)), node
-            )
+        return [
+            self._assign(special, special_expr, node),
+            self._assign(absent, absent_expr, node),
         ]
+
+    def _symbol_declarations(
+        self, node: ast.FunctionDef, facts: tuple[Binding, ...]
+    ) -> tuple[list[ast.stmt], dict[str, str]]:
+        """Predeclare symbol types and bind explicit signature type 
parameters."""
+        declaration: list[ast.stmt] = []
+        symbol_aliases: dict[str, str] = {}
         # Source: def f[n](...); Builder: n = X.resolve_type_var_("n").
-        # Explicit dtype declarations are predeclared before quoted shapes, but
-        # only explicit type parameters introduce signature Python bindings.
-        symbol_aliases = {}
-        for parameter in type_parameters:
+        # Explicit dtype declarations precede quoted shapes, but only explicit
+        # type parameters introduce signature Python bindings.
+        for parameter in getattr(node, "type_params", ()):
             if not isinstance(parameter, getattr(ast, "TypeVar", ())):
                 self._error(parameter, "Only scalar type parameters are 
supported")
             bound = getattr(parameter, "bound", None)
@@ -1142,13 +1124,9 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             alias = self.fresh("_symbol")
             symbol_aliases[parameter.name] = alias
             declaration.append(
-                ast.copy_location(
-                    ast.Assign(
-                        [ast.Name(alias, ast.Store())],
-                        self._operation(
-                            "resolve_type_var_", 
[ast.Constant(parameter.name)], parameter
-                        ),
-                    ),
+                self._assign(
+                    alias,
+                    self._operation("resolve_type_var_", 
[ast.Constant(parameter.name)], parameter),
                     parameter,
                 )
             )
@@ -1160,23 +1138,60 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                     item.node,
                     dtype=ast.Constant(item.dtype),
                 )
-                # The native map knows later scalar dtypes, but a later Python
-                # parameter name is not in annotation scope until its own arg.
+                # A later Python parameter name does not enter annotation scope
+                # until its own arg, even though the native map knows its 
dtype.
                 declaration.append(ast.copy_location(ast.Expr(symbol), 
item.node))
-        self.annotation_aliases = {**aliases, **symbol_aliases}
-        constexpr_aliases = {}
-        ordered = sorted(
-            zip(parameters, annotations),
-            key=lambda pair: not (
-                isinstance(pair[1], ast.Attribute) and pair[1].attr == 
"constexpr"
+        return declaration, symbol_aliases
+
+    @staticmethod
+    def _specialized_value(
+        name: str, fallback: ast.expr, node: ast.AST, *, special: str, absent: 
str
+    ) -> ast.IfExp:
+        """Select a JIT value or explicit absence without evaluating the 
fallback."""
+        selected = ast.BoolOp(
+            ast.And(),
+            [
+                ast.Compare(ast.Name(special, ast.Load()), [ast.IsNot()], 
[ast.Constant(None)]),
+                ast.Compare(ast.Constant(name), [ast.In()], [ast.Name(special, 
ast.Load())]),
+            ],
+        )
+        absent_test = ast.Compare(ast.Constant(name), [ast.In()], 
[ast.Name(absent, ast.Load())])
+        return ast.copy_location(
+            ast.IfExp(
+                selected,
+                ast.Subscript(ast.Name(special, ast.Load()), 
ast.Constant(name), ast.Load()),
+                ast.IfExp(absent_test, ast.Constant(None), fallback),
             ),
+            node,
         )
-        for parameter, annotation in ordered:
+
+    def _parameter_declarations(
+        self,
+        parameters: list[ast.arg],
+        annotations: list[ast.expr | None],
+        *,
+        captures: str,
+        special: str,
+        absent: str,
+    ) -> tuple[list[ast.stmt], dict[str, str]]:
+        """Declare signature parameters, making constexpr values available 
first."""
+        from . import jit_support
+
+        declaration: list[ast.stmt] = []
+        constexpr_aliases: dict[str, str] = {}
+        constexpr_params: list[tuple[ast.arg, ast.expr | None, bool]] = []
+        other_params: list[tuple[ast.arg, ast.expr | None, bool]] = []
+        # Runtime annotations may depend on a later constexpr parameter. 
Preserve
+        # source order within each group and reuse this one syntactic 
classification.
+        for parameter, annotation in zip(parameters, annotations):
+            is_constexpr = isinstance(annotation, ast.Attribute) and 
annotation.attr == "constexpr"
+            group = constexpr_params if is_constexpr else other_params
+            group.append((parameter, annotation, is_constexpr))
+        for parameter, annotation, is_constexpr in [*constexpr_params, 
*other_params]:
             if annotation is None:
                 self._error(parameter, f"Parameter {parameter.arg!r} requires 
an annotation")
             name = parameter.arg
             alias = self.fresh("_parameter")
-            is_constexpr = isinstance(annotation, ast.Attribute) and 
annotation.attr == "constexpr"
             if is_constexpr:
                 constexpr_aliases[name] = alias
                 fallback = ast.Call(
@@ -1207,42 +1222,26 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                     [],
                 )
                 fallback = self._operation("arg", [ast.Constant(name), 
checked], parameter)
-            # Source parameters selected by JIT have no runtime ABI slot. Their
-            # annotation thunk is absent from the selected Python branch.
-            selected = ast.BoolOp(
-                ast.And(),
-                [
-                    ast.Compare(ast.Name(special, ast.Load()), [ast.IsNot()], 
[ast.Constant(None)]),
-                    ast.Compare(ast.Constant(name), [ast.In()], 
[ast.Name(special, ast.Load())]),
-                ],
-            )
-            absent_test = ast.Compare(
-                ast.Constant(name), [ast.In()], [ast.Name(absent, ast.Load())]
-            )
-            value = ast.IfExp(
-                selected,
-                ast.Subscript(ast.Name(special, ast.Load()), 
ast.Constant(name), ast.Load()),
-                ast.IfExp(absent_test, ast.Constant(None), fallback),
-            )
-            declaration.append(
-                ast.copy_location(ast.Assign([ast.Name(alias, ast.Store())], 
value), parameter)
+            # Specialized parameters have no runtime ABI slot. The annotation
+            # thunk is absent from the selected generated Python branch.
+            value = self._specialized_value(
+                name, fallback, parameter, special=special, absent=absent
             )
+            declaration.append(self._assign(alias, value, parameter))
             self.annotation_aliases[name] = alias
-        if returns is not None:
-            declaration.append(
-                ast.copy_location(
-                    ast.Expr(
-                        self._call(
-                            builder,
-                            "func_ret_type",
-                            [self._lambda([], self.visit(returns))],
-                            returns,
-                        )
-                    ),
-                    returns,
-                )
-            )
-        self.annotation_aliases = {}
+        return declaration, constexpr_aliases
+
+    def _function_frame(
+        self,
+        node: ast.FunctionDef,
+        options: ast.Dict,
+        declaration: list[ast.stmt],
+        *,
+        frame: str,
+        local: bool,
+        declare: bool,
+    ) -> list[ast.stmt]:
+        """Create the native frame and optionally execute its declaration 
pass."""
         # Source: @X.function def f(...): ...
         # Builder: with X.function(decl=True) as fn: signature
         with self._host():
@@ -1256,7 +1255,9 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             keywords.append(ast.keyword("span", self.span(node)))
         constructor = ast.copy_location(
             ast.Call(
-                ast.Attribute(ast.Name(builder, ast.Load()), "function", 
ast.Load()), [], keywords
+                ast.Attribute(ast.Name(self.dialect_prefix, ast.Load()), 
"function", ast.Load()),
+                [],
+                keywords,
             ),
             node,
         )
@@ -1264,109 +1265,168 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             declaration_scope = ast.With(
                 [ast.withitem(constructor, ast.Name(frame, ast.Store()))], 
declaration
             )
-            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(); with fn: define/call a helper for
-            # both signature and body, entering this ordinary frame just once.
-            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.
+            return [
+                ast.copy_location(declaration_scope, node),
+                self._assign(node.name, reference, node),
+            ]
+        # Ordinary standalone functions enter once for their signature and 
body.
+        return [self._assign(frame, constructor, node)]
+
+    def _body_parameters(
+        self,
+        node: ast.FunctionDef,
+        parameters: list[ast.arg],
+        constexpr_aliases: dict[str, str],
+        *,
+        frame: str,
+        special: str,
+        absent: str,
+    ) -> list[ast.stmt]:
+        """Read existing frame parameters into the body's Python lexical 
scope."""
         iterator = self.fresh("_arguments")
         body: list[ast.stmt] = [
-            ast.copy_location(
-                ast.Assign(
-                    [ast.Name(iterator, ast.Store())],
-                    ast.Call(
-                        self._inject(iter),
-                        [ast.Attribute(ast.Name(frame, ast.Load()), "params", 
ast.Load())],
-                        [],
-                    ),
+            self._assign(
+                iterator,
+                ast.Call(
+                    self._inject(iter),
+                    [ast.Attribute(ast.Name(frame, ast.Load()), "params", 
ast.Load())],
+                    [],
                 ),
                 node,
             )
         ]
         for parameter in parameters:
             name = parameter.arg
+            value: ast.expr
             if name in constexpr_aliases:
                 value = ast.Name(constexpr_aliases[name], ast.Load())
             else:
-                selected = ast.BoolOp(
-                    ast.And(),
-                    [
-                        ast.Compare(
-                            ast.Name(special, ast.Load()), [ast.IsNot()], 
[ast.Constant(None)]
-                        ),
-                        ast.Compare(
-                            ast.Constant(name), [ast.In()], [ast.Name(special, 
ast.Load())]
-                        ),
-                    ],
-                )
-                absent_test = ast.Compare(
-                    ast.Constant(name), [ast.In()], [ast.Name(absent, 
ast.Load())]
-                )
-                value = ast.IfExp(
-                    selected,
-                    ast.Subscript(ast.Name(special, ast.Load()), 
ast.Constant(name), ast.Load()),
-                    ast.IfExp(
-                        absent_test,
-                        ast.Constant(None),
-                        ast.Call(self._inject(next), [ast.Name(iterator, 
ast.Load())], []),
-                    ),
+                value = self._specialized_value(
+                    name,
+                    ast.Call(self._inject(next), [ast.Name(iterator, 
ast.Load())], []),
+                    parameter,
+                    special=special,
+                    absent=absent,
                 )
+            body.append(self._assign(name, value, parameter))
+        for parameter in getattr(node, "type_params", ()):
             body.append(
-                ast.copy_location(ast.Assign([ast.Name(name, ast.Store())], 
value), parameter)
-            )
-        for parameter in type_parameters:
-            body.append(
-                ast.copy_location(
-                    ast.Assign(
-                        [ast.Name(parameter.name, ast.Store())],
-                        self._operation(
-                            "resolve_type_var_", 
[ast.Constant(parameter.name)], parameter
-                        ),
-                    ),
+                self._assign(
+                    parameter.name,
+                    self._operation("resolve_type_var_", 
[ast.Constant(parameter.name)], parameter),
                     parameter,
                 )
             )
-        # Body annotations use definition aliases only for names without a real
-        # body binding. Ordinary body references are visited with no 
substitution.
-        body_annotation_aliases = {
-            name: alias for name, alias in aliases.items() if name not in 
local_names
-        }
-        old_body_annotations = self.body_annotation_aliases
-        self.body_annotation_aliases = body_annotation_aliases
-        body.extend(self.transform_statements(node.body))
-        self.body_annotation_aliases = old_body_annotations
-        definition = self._definition(body_name, body if declare else 
[*declaration, *body], node)
-        if not local:
-            definition._tvm_source_name = node.name
-            definition._tvm_signature_names = (
-                declared_names | {node.name} | set(self.module_functions)
-            )
-            if self.module_name:
-                definition._tvm_signature_names.add(self.module_name)
-        # Source: def f(...): body
-        # Builder:
-        #   with fn:
-        #       def build(): body
-        #       build()
-        # The helper owns only Python lexical scope; the enclosing with owns
-        # native frame entry/exit, including unwinding a failed body.
-        invocation = ast.copy_location(
-            ast.Expr(ast.Call(ast.Name(body_name, ast.Load()), [], [])), node
-        )
-        resumed = ast.copy_location(
-            ast.With([ast.withitem(ast.Name(frame, ast.Load()))], [definition, 
invocation]), node
-        )
-        self.dialect_prefix, self.current_scope, self.annotation_aliases = old
-        return statements, frame, resumed
+        return body
+
+    def function_program(
+        self, node: ast.FunctionDef, *, local: bool = False, declare: bool = 
True
+    ) -> tuple[list[ast.stmt], str, ast.With]:
+        """Declare a native frame and emit a lexical body helper inside its 
scope.
+
+        Definition aliases retain outer annotation values. Signature aliases 
add
+        declared symbols and each preceding parameter; constexpr aliases retain
+        compile-time values for the body. Body annotation aliases include only
+        definition names without a real body-local binding. None owns native 
IR.
+        """
+        kind, options = self.function_metadata(node)
+        builder = self.fresh("_X")
+        self.bindings[builder] = kind.builder
+        frame, body_name = self.fresh("_fn"), self.fresh("_build")
+        old = self.dialect_prefix, self.current_scope
+        self.dialect_prefix, self.current_scope = builder, node
+        try:
+            with self._aliases({}):
+                parameters = [*node.args.posonlyargs, *node.args.args, 
*node.args.kwonlyargs]
+                if node.args.vararg or node.args.kwarg:
+                    self._error(node, "IR signatures require ordinary named 
parameters")
+                facts = self.prescan.bindings.get(node, ()) if self.prescan 
else ()
+                annotations, returns, definition_aliases = 
self._function_annotations(
+                    node, parameters, facts
+                )
+                captures = self.fresh("_definition")
+                statements = self._definition_bindings(
+                    node, definition_aliases, captures=captures, local=local
+                )
+                special, absent = self.fresh("_specialization"), 
self.fresh("_absent")
+                statements.extend(
+                    self._specialization_bindings(node, special=special, 
absent=absent, local=local)
+                )
+                declaration: list[ast.stmt] = [
+                    ast.copy_location(
+                        ast.Expr(self._call(builder, "func_name", 
[ast.Constant(node.name)], node)),
+                        node,
+                    )
+                ]
+                symbols, symbol_aliases = self._symbol_declarations(node, 
facts)
+                declaration.extend(symbols)
+                with self._aliases({**definition_aliases, **symbol_aliases}):
+                    arguments, constexpr_aliases = 
self._parameter_declarations(
+                        parameters, annotations, captures=captures, 
special=special, absent=absent
+                    )
+                    declaration.extend(arguments)
+                    if returns is not None:
+                        declaration.append(
+                            ast.copy_location(
+                                ast.Expr(
+                                    self._call(
+                                        builder,
+                                        "func_ret_type",
+                                        [self._lambda([], 
self.visit(returns))],
+                                        returns,
+                                    )
+                                ),
+                                returns,
+                            )
+                        )
+                statements.extend(
+                    self._function_frame(
+                        node, options, declaration, frame=frame, local=local, 
declare=declare
+                    )
+                )
+                body = self._body_parameters(
+                    node, parameters, constexpr_aliases, frame=frame, 
special=special, absent=absent
+                )
+                local_names = {item.name for item in facts}
+                with self._aliases(
+                    {},
+                    body={
+                        name: alias
+                        for name, alias in definition_aliases.items()
+                        if name not in local_names
+                    },
+                ):
+                    body.extend(self.transform_statements(node.body))
+                definition = self._definition(
+                    body_name, body if declare else [*declaration, *body], node
+                )
+                if not local:
+                    definition._tvm_source_name = node.name
+                    definition._tvm_signature_names = (
+                        {item.name for item in getattr(node, "type_params", 
())}
+                        | {node.name}
+                        | set(self.module_functions)
+                    )
+                    if self.module_name:
+                        definition._tvm_signature_names.add(self.module_name)
+                # Source: def f(...): body
+                # Builder:
+                #   with fn:
+                #       def build(): body
+                #       build()
+                # The helper owns Python lexical scope; with owns native frame
+                # entry/exit, including unwinding a failed body.
+                invocation = ast.copy_location(
+                    ast.Expr(ast.Call(ast.Name(body_name, ast.Load()), [], 
[])), node
+                )
+                resumed = ast.copy_location(
+                    ast.With([ast.withitem(ast.Name(frame, ast.Load()))], 
[definition, invocation]),
+                    node,
+                )
+                return statements, frame, resumed
+        finally:
+            self.dialect_prefix, self.current_scope = old
 
     def program(self, tree: ast.Module) -> tuple[ast.Module, str]:
         """Emit direct native module construction, declarations, then 
bodies."""
@@ -1387,15 +1447,13 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         body: list[ast.stmt] = []
         for function in functions if declare else ():
             body.append(
-                ast.copy_location(
-                    ast.Assign(
-                        [ast.Name(function.name, ast.Store())],
-                        self._call(
-                            self.infrastructure_name,
-                            "reserve_function",
-                            [ast.Constant(function.name)],
-                            function,
-                        ),
+                self._assign(
+                    function.name,
+                    self._call(
+                        self.infrastructure_name,
+                        "reserve_function",
+                        [ast.Constant(function.name)],
+                        function,
                     ),
                     function,
                 )
@@ -1422,11 +1480,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                         [ast.Constant(target.id), ast.Name(target.id, 
ast.Load())],
                         target,
                     )
-                    body.append(
-                        ast.copy_location(
-                            ast.Assign([ast.Name(target.id, ast.Store())], 
value), target
-                        )
-                    )
+                    body.append(self._assign(target.id, value, target))
         definitions, frames, python_functions = [], [], []
         for function in functions:
             kind, _ = self.function_metadata(function)
@@ -1437,14 +1491,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                 function.decorator_list = []
                 body.append(function)
                 host = self.fresh("_python")
-                body.append(
-                    ast.copy_location(
-                        ast.Assign(
-                            [ast.Name(host, ast.Store())], 
ast.Name(function.name, ast.Load())
-                        ),
-                        function,
-                    )
-                )
+                body.append(self._assign(host, ast.Name(function.name, 
ast.Load()), function))
                 opaque = self.fresh("_opaque")
                 value = ast.Call(
                     self._inject(make_opaque_function),
@@ -1456,19 +1503,15 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                     ],
                     [],
                 )
+                body.append(self._assign(opaque, value, function))
                 body.append(
-                    ast.copy_location(ast.Assign([ast.Name(opaque, 
ast.Store())], value), function)
-                )
-                body.append(
-                    ast.copy_location(
-                        ast.Assign(
-                            [ast.Name(function.name, ast.Store())],
-                            self._call(
-                                self.infrastructure_name,
-                                "decl_function",
-                                [ast.Constant(function.name), ast.Name(opaque, 
ast.Load())],
-                                function,
-                            ),
+                    self._assign(
+                        function.name,
+                        self._call(
+                            self.infrastructure_name,
+                            "decl_function",
+                            [ast.Constant(function.name), ast.Name(opaque, 
ast.Load())],
+                            function,
                         ),
                         function,
                     )
@@ -1538,7 +1581,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         translated.extend(
             [
                 construction,
-                ast.copy_location(ast.Assign([ast.Name(result, ast.Store())], 
output), root),
+                self._assign(result, output, root),
             ]
         )
         translated.append(
diff --git a/tests/python/script/test_parser_alias_restoration.py 
b/tests/python/script/test_parser_alias_restoration.py
new file mode 100644
index 0000000000..4880f367a9
--- /dev/null
+++ b/tests/python/script/test_parser_alias_restoration.py
@@ -0,0 +1,85 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""A failed production visitor must restore its caller's annotation scope."""
+
+import ast
+
+import pytest
+
+from tvm.script.parser import entry
+
+
[email protected](
+    "source, expression",
+    [
+        pytest.param("lambda outer: (broken := 1)", True, id="lambda"),
+        pytest.param("[(broken := 1) for outer in values]", True, 
id="comprehension"),
+        pytest.param("value: (broken := 1) = 0", False, id="body-annotation"),
+        pytest.param(
+            "@X.script\ndef nested(x: (broken := 1)):\n    pass\n",
+            False,
+            id="function-signature",
+        ),
+        pytest.param(
+            "@X.script\ndef nested(x: X.tensor((4,))):\n    value: outer = 0\n 
   (broken := 1)\n",
+            False,
+            id="function-body",
+        ),
+    ],
+)
+def test_failed_visitor_restores_later_name_and_body_annotation(language, 
source, expression):
+    setup = "@X.script\ndef main():\n    pass\n"
+    tree, filename, _ = entry.acquire_source(setup, 
filename="alias_restoration.py")
+    transformer, namespace = entry._prepare_transpiler(
+        tree, setup, {"X": language.X}, {}, filename, track_span=False
+    )
+    outer, body = object(), object()
+    namespace.update(_outer_value=outer, _body_value=body)
+    aliases = {"outer": "_outer_value"}
+    body_aliases = {"outer": "_body_value"}
+    transformer.annotation_aliases = aliases
+    transformer.body_annotation_aliases = body_aliases
+    previous_scope = transformer.current_scope
+    previous_dialect = transformer.dialect_prefix
+    failing = ast.parse(source, mode="eval").body if expression else 
ast.parse(source).body[0]
+
+    with pytest.raises(SyntaxError, match="Unsupported expression: NamedExpr"):
+        transformer.visit(failing)
+
+    # Check the observable rewrite after failure, rather than the manager 
alone:
+    # ordinary reads use the caller's alias, body annotations use its other 
map.
+    reference = transformer.visit(ast.parse("outer", mode="eval").body)
+    assert (
+        eval(
+            compile(ast.fix_missing_locations(ast.Expression(reference)), 
filename, "eval"),
+            namespace,
+        )
+        is outer
+    )
+    assignment = transformer.visit(ast.parse("result: outer = 0").body[0])[0]
+    annotation = next(keyword.value for keyword in assignment.value.keywords 
if keyword.arg == "ty")
+    assert (
+        eval(
+            compile(ast.fix_missing_locations(ast.Expression(annotation)), 
filename, "eval"),
+            namespace,
+        )
+        is body
+    )
+    assert transformer.annotation_aliases is aliases
+    assert transformer.body_annotation_aliases is body_aliases
+    assert transformer.current_scope is previous_scope
+    assert transformer.dialect_prefix == previous_dialect

Reply via email to