This is an automated email from the ASF dual-hosted git repository. tqchen pushed a commit to branch tvmscript-generic-parser-builder in repository https://gitbox.apache.org/repos/asf/tvm.git
commit 46d1ea91a23b7aed69e6e736a1defad15d86d486 Author: Tianqi Chen <[email protected]> AuthorDate: Mon Sep 21 02:45:16 2026 +0000 Retain function metadata and Python module behavior --- python/tvm/script/parser_v2/frontend.py | 55 +++++++++++++++++++++++++++------ 1 file changed, 46 insertions(+), 9 deletions(-) diff --git a/python/tvm/script/parser_v2/frontend.py b/python/tvm/script/parser_v2/frontend.py index e39bdc41a1..6305e05bed 100644 --- a/python/tvm/script/parser_v2/frontend.py +++ b/python/tvm/script/parser_v2/frontend.py @@ -28,6 +28,7 @@ from functools import wraps from types import SimpleNamespace from typing import TypeVar +import tvm from tvm import ir from tvm.error import DiagnosticError from tvm.script.ir_builder import IRBuilder, protocol @@ -35,7 +36,13 @@ from tvm.script.ir_builder import ir as I from .annotations import AnnotationScope from .diagnostics import diagnostic_error -from .functions import FunctionGroup, attach_python, declare_python, is_python_function +from .functions import ( + FunctionGroup, + adapt_python_module, + attach_python, + declare_python, + is_python_function, +) _NAMESPACES = {} @@ -94,6 +101,9 @@ def make_decorator(builder, *, option_map=None, defaults=None): mapping, default_options = dict(option_map or {}), dict(defaults or {}) def decorator(function=None, **options): + if function is not None and not inspect.isfunction(function): + raise ValueError("Construction decorators require a function or keyword options") + def apply(function): function.__tvm_function_kind__ = decorator.__tvm_function_kind__ function.__tvm_function_options__ = options @@ -112,6 +122,9 @@ def make_helper(builder, *, preserve_return=True): """Create an explicit construction helper using the active shared builder.""" def decorator(function=None, **options): + if function is not None and not inspect.isfunction(function): + raise ValueError("Construction decorators require a function or keyword options") + def apply(function): definition_env = _capture(function) @@ -164,7 +177,7 @@ class Compiler: """Build from a copied original AST, retaining file and range information.""" def __init__(self, source, env=None, filename=None): - self.env = {"TypeVar": TypeVar, **_NAMESPACES, **(env or {})} + self.env = {"TypeVar": TypeVar, "tvm": tvm, **_NAMESPACES, **(env or {})} self.original = source members = vars(source).values() if inspect.isclass(source) else (source,) self.compile_flags = 0 @@ -174,7 +187,7 @@ class Compiler: self.compile_flags |= code.co_flags & __future__.annotations.compiler_flag if isinstance(source, str): text = source - self.filename = filename or "<tvmscript>" + self.filename = filename or "<str>" start, indent = 1, 0 linecache.cache[self.filename] = ( len(text), @@ -235,9 +248,27 @@ class Compiler: node, ) - def function_kind(self, node, env): + def function_kind(self, node, env, *, allow_python=False): if id(node) in self.function_kinds: return self.function_kinds[id(node)] + if inspect.isfunction(self.original) and node is self.tree.body[0]: + kind = protocol.function_kind(self.original) + if kind is not None: + options = { + **kind.metadata.get("defaults", {}), + **getattr(self.original, "__tvm_function_options__", {}), + } + mapping = kind.metadata.get("option_map", {}) + result = ( + kind, + { + mapping.get(key, key): value + for key, value in options.items() + if key != "check_well_formed" + }, + ) + self.function_kinds[id(node)] = result + return result for decorator in node.decorator_list: target = decorator.func if isinstance(decorator, ast.Call) else decorator value = _resolve(target, env, self.filename) @@ -260,6 +291,8 @@ class Compiler: } self.function_kinds[id(node)] = (kind, options) return kind, options + if allow_python: + return protocol.FunctionKind(None, {"python": True}), {} raise SyntaxError(f"Function {node.name!r} has no registered construction kind") def declare(self, node, env, *, local=False): @@ -287,7 +320,7 @@ class Compiler: ) spec.params[argument.arg] = scope.env[argument.arg] = value if node.returns is not None: - spec.result_type = scope.evaluate(node.returns) + spec.result_type = scope.evaluate(node.returns, introduce=False) X.func_ret_type(spec.result_type) spec.reference = frame.reference scope.env[node.name] = spec.reference @@ -310,8 +343,9 @@ class Compiler: set(spec.params) | set(spec.scope.symbols), scope=spec.scope, ) - frame.function.__name__ = spec.node.name - return frame.function + result = frame.function + result.__name__ = spec.node.name + return result def run_statements(self, body, builder, env, bound_names, *, scope=None, preserve_return=False): from .transform import Transformer @@ -328,7 +362,8 @@ class Compiler: nested_name = self.fresh("nested") def nested_statement(node): - if is_python_function(self, node, namespace): + kind, _ = self.function_kind(node, namespace, allow_python=True) + if kind.metadata.get("python"): return [copy.deepcopy(node)] index = len(nested) nested[index] = node @@ -504,7 +539,9 @@ class Compiler: attach_python(module, python_functions) if root is not None: module.__name__ = root.name - return module + original = self.original if inspect.isclass(self.original) else None + bases = tuple(_resolve(base, env, self.filename) for base in root.bases) + return adapt_python_module(module, original=original, bases=bases) return results[functions[0].name]
