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 9ac001e89d9b05dcece2f0541bd839cb96789f10
Author: Tianqi Chen <[email protected]>
AuthorDate: Wed Sep 23 13:18:11 2026 +0000

    [REFACTOR][Script] Define lexical helpers inside frame contexts
    
    Enter native frames before defining and invoking their Python scope 
helpers. Preserve declaration ordering, lexical captures and source metadata 
while removing synthetic frame arguments.
---
 python/tvm/script/parser/transpile.py            |  81 +++++-----
 tests/python/script/test_parser_scope_helpers.py | 198 +++++++++++++++++++++++
 2 files changed, 240 insertions(+), 39 deletions(-)

diff --git a/python/tvm/script/parser/transpile.py 
b/python/tvm/script/parser/transpile.py
index 9ae715526b..30a708ba43 100644
--- a/python/tvm/script/parser/transpile.py
+++ b/python/tvm/script/parser/transpile.py
@@ -133,15 +133,13 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         )
 
     @staticmethod
-    def _definition(
-        name: str, parameters: list[str], body: list[ast.stmt], node: ast.AST
-    ) -> ast.FunctionDef:
+    def _definition(name: str, body: list[ast.stmt], node: ast.AST) -> 
ast.FunctionDef:
         definition = ast.copy_location(
             ast.FunctionDef(
                 name,
                 ast.arguments(
                     posonlyargs=[],
-                    args=[ast.arg(name) for name in parameters],
+                    args=[],
                     kwonlyargs=[],
                     kw_defaults=[],
                     defaults=[],
@@ -775,7 +773,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         # Branch helper scoping keeps source names; mutable stores bind no 
locals.
         name = self.fresh(prefix)
         return [
-            self._definition(name, [], self.transform_statements(body), node),
+            self._definition(name, self.transform_statements(body), node),
             ast.copy_location(ast.Expr(ast.Call(ast.Name(name, ast.Load()), 
[], [])), node),
         ]
 
@@ -957,20 +955,18 @@ class IRBuilderTranspiler(ast.NodeTransformer):
 
     def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.FunctionDef | 
list[ast.stmt]:
         # Source: nested @X.function def f(...): body
-        # Builder: declare fn; def build(fn): with fn: body; build(fn).
+        # Builder:
+        #   declare fn
+        #   with fn:
+        #       def build(): body
+        #       build()
         if self.host_expression:
             return node
         kind, _ = self.function_metadata(node, allow_python=True)
         if kind.python:
             return node
-        declaration, frame, body = self.function_program(node, local=True)
-        return [
-            *declaration,
-            ast.copy_location(
-                ast.Expr(ast.Call(ast.Name(body, ast.Load()), [ast.Name(frame, 
ast.Load())], [])),
-                node,
-            ),
-        ]
+        declaration, _, body = self.function_program(node, local=True)
+        return [*declaration, body]
 
     def visit_Nonlocal(self, node: ast.Nonlocal) -> ast.Pass:
         # Source closure declarations do not mutate the host closure during 
build.
@@ -1021,12 +1017,13 @@ class IRBuilderTranspiler(ast.NodeTransformer):
 
     def function_program(
         self, node: ast.FunctionDef, *, local: bool = False, declare: bool = 
True
-    ) -> tuple[list[ast.stmt], str, str]:
-        """Declare one native frame and emit one body function taking that 
frame.
+    ) -> 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 frame when its body resumes. No factory,
-        callback record, copied parameter map or symbol owner is generated.
+        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
 
@@ -1274,7 +1271,8 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             )
         else:
             # Source: a standalone nonrecursive function.
-            # Builder: fn = X.function(); build(fn) enters once for 
signature+body.
+            # 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)
             )
@@ -1346,14 +1344,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         self.body_annotation_aliases = body_annotation_aliases
         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 if declare else [*declaration, *body],
-            ),
-            node,
-        )
-        definition = self._definition(body_name, [frame], [resumed], node)
+        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 = (
@@ -1361,9 +1352,21 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             )
             if self.module_name:
                 definition._tvm_signature_names.add(self.module_name)
-        statements.append(definition)
+        # 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, body_name
+        return statements, frame, resumed
 
     def program(self, tree: ast.Module) -> tuple[ast.Module, str]:
         """Emit direct native module construction, declarations, then 
bodies."""
@@ -1485,22 +1488,22 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                 )
                 python_functions.append((function.name, host))
                 continue
-            declaration, frame, callback = self.function_program(function, 
declare=declare)
+            declaration, frame, definition = self.function_program(function, 
declare=declare)
             body.extend(declaration)
-            definitions.append(
-                ast.copy_location(
-                    ast.Expr(
-                        ast.Call(ast.Name(callback, ast.Load()), 
[ast.Name(frame, ast.Load())], [])
-                    ),
-                    function,
-                )
-            )
+            definitions.append(definition)
             frames.append(frame)
         body.extend(definitions)
         # Source: class Module: functions...
         # Builder:
         #   with IRBuilder() as builder:
-        #       with I.ir_module(): declarations; build_f(fn); build_g(gn)
+        #       with I.ir_module():
+        #           declarations
+        #           with fn:
+        #               def build_f(): body_f
+        #               build_f()
+        #           with gn:
+        #               def build_g(): body_g
+        #               build_g()
         #   result = builder.get()
         module = ast.copy_location(
             ast.With(
diff --git a/tests/python/script/test_parser_scope_helpers.py 
b/tests/python/script/test_parser_scope_helpers.py
new file mode 100644
index 0000000000..bb1ea9fb59
--- /dev/null
+++ b/tests/python/script/test_parser_scope_helpers.py
@@ -0,0 +1,198 @@
+# 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.
+"""Lexical-only body helpers execute inside their already-entered frames."""
+
+import ast
+import copy
+
+import pytest
+
+from tvm.error import DiagnosticError
+from tvm.script.ir_builder import IRBuilder
+from tvm.script.parser import entry
+
+
[email protected](
+    "source, helper_count",
+    [
+        pytest.param(
+            "@X.script\ndef main(x: X.tensor((4,))):\n    X.record(x)\n",
+            1,
+            id="ordinary",
+        ),
+        pytest.param(
+            "@X.script\ndef main(x: X.tensor((4,))):\n    main(x)\n",
+            1,
+            id="recursive",
+        ),
+        pytest.param(
+            """
[email protected]_module
+class Module:
+    @X.script
+    def first(x: X.tensor((4,))):
+        second(x)
+    @X.script
+    def second(y: X.tensor((4,))):
+        X.record(y)
+""",
+            2,
+            id="module",
+        ),
+        pytest.param(
+            """
[email protected]
+def main(x: X.tensor((4,))):
+    @X.script
+    def inner(y: X.tensor((4,))):
+        X.record(x)
+        X.record(y)
+    X.record(x)
+""",
+            2,
+            id="nested",
+        ),
+        pytest.param(
+            """
[email protected]
+def main(condition: X.tensor(())):
+    if condition:
+        X.record(1)
+    else:
+        X.record(2)
+""",
+            3,
+            id="branches",
+        ),
+    ],
+)
+def test_lexical_helpers_are_defined_and_called_inside_frames(
+    language, monkeypatch, source, helper_count
+):
+    # The requested structure is with frame: def body(): ...; body().
+    # Observe the actual transpiled program and still execute its normal path.
+    programs = []
+    original = entry.recompose_builder
+
+    def capture(translated, **kwargs):
+        programs.append(copy.deepcopy(translated))
+        return original(translated, **kwargs)
+
+    monkeypatch.setattr(entry, "recompose_builder", capture)
+    language.parse(source)
+    assert len(programs) == 1
+    tree = programs[0]
+    parents = {child: parent for parent in ast.walk(tree) for child in 
ast.iter_child_nodes(parent)}
+    helpers = [node for node in ast.walk(tree) if isinstance(node, 
ast.FunctionDef)]
+    assert len(helpers) == helper_count
+    for helper in helpers:
+        scope = parents[helper]
+        assert isinstance(scope, ast.With)
+        assert len(scope.body) == 2 and scope.body[0] is helper
+        invocation = scope.body[1]
+        assert isinstance(invocation, ast.Expr) and 
isinstance(invocation.value, ast.Call)
+        assert isinstance(invocation.value.func, ast.Name)
+        assert invocation.value.func.id == helper.name
+        assert invocation.value.args == [] and invocation.value.keywords == []
+        assert helper.args.posonlyargs == [] and helper.args.args == []
+        assert helper.args.kwonlyargs == []
+        assert helper.args.vararg is None and helper.args.kwarg is None
+        assert (scope.lineno, invocation.lineno) == (helper.lineno, 
helper.lineno)
+        assert (scope.end_lineno, invocation.end_lineno) == 
(helper.end_lineno, helper.end_lineno)
+    assert language.stack == [] and language.source_stack == []
+    assert not IRBuilder.is_in_scope()
+
+
[email protected]("fail", [False, True])
+def test_nested_helper_frames_capture_parameters_and_unwind(language, fail):
+    seen = []
+    failure = ValueError("inner body failure")
+
+    def observe(label, outer, current, scope):
+        frame = language.frame()
+        assert language.stack[-1] is frame
+        seen.append((label, frame, outer, current, scope))
+        if fail and label == "inner":
+            raise failure
+
+    source = """
[email protected]
+def main(x: X.tensor((4,))):
+    scope = 1
+    observe("outer_before", x, x, scope)
+    @X.script
+    def inner(y: X.tensor((4,))):
+        scope = 2
+        observe("inner", x, y, scope)
+    observe("outer_after", x, x, scope)
+"""
+    if fail:
+        with pytest.raises(DiagnosticError) as caught:
+            language.parse(source, observe=observe)
+        assert caught.value.__cause__ is failure
+        source_line = next(
+            index for index, line in enumerate(source.splitlines(), 1) if 
'observe("inner"' in line
+        )
+        assert f"dummy.py:{source_line}:" in str(caught.value)
+    else:
+        result = language.parse(source, observe=observe)
+        assert result is language.functions["main"]
+
+    entries = [event for event in language.events if event[:2] == ("enter", 
"function")]
+    exits = [event for event in language.events if event[:2] == ("exit", 
"function")]
+    assert [event[2] for event in entries] == [False, True, False]
+    outer, inner = entries[0][3], entries[1][3]
+    assert entries[2][3] is inner
+    assert [event[3] for event in exits] == [inner, inner, outer]
+    assert len(outer.params) == len(inner.params) == 1
+    assert seen[0] == ("outer_before", outer, outer.params[0], 
outer.params[0], 1)
+    assert seen[1] == ("inner", inner, outer.params[0], inner.params[0], 2)
+    if fail:
+        assert len(seen) == 2
+    else:
+        assert seen[2] == ("outer_after", outer, outer.params[0], 
outer.params[0], 1)
+    assert language.stack == [] and language.source_stack == []
+    assert not IRBuilder.is_in_scope()
+
+
+def test_zero_argument_helper_retains_real_closure_defaults(language, 
monkeypatch):
+    token = object()
+    X = language.X
+    helpers = []
+    original = entry.recompose_builder
+
+    def capture(translated, **kwargs):
+        result = original(translated, **kwargs)
+        helpers.extend(node for node in ast.walk(translated) if hasattr(node, 
"_tvm_source_name"))
+        return result
+
+    monkeypatch.setattr(entry, "recompose_builder", capture)
+
+    @X.script
+    def main(x: X.tensor((4,))):
+        X.record(token)
+        X.record(x)
+
+    assert main.body[0][1] is token
+    assert main.body[1][1] is main.params[0]
+    assert len(helpers) == 1
+    helper = helpers[0]
+    assert helper.args.args == [] and helper.args.posonlyargs == []
+    captures = [argument.arg for argument in helper.args.kwonlyargs]
+    assert "token" in captures
+    assert len(helper.args.kw_defaults) == len(captures)
+    assert all(default is not None for default in helper.args.kw_defaults)

Reply via email to