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)
