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 4eb9effd10458ec4277eca784d2603fc1d5d83ef Author: Tianqi Chen <[email protected]> AuthorDate: Mon Sep 21 03:31:00 2026 +0000 Preserve Python helper binding and dialect source namespaces --- python/tvm/script/parser_v2/frontend.py | 44 ++++++++++++++++++++++---------- python/tvm/script/parser_v2/transform.py | 4 +++ python/tvm/tirx/script/builder/ir.py | 2 ++ python/tvm/tirx/script/v2.py | 3 +++ 4 files changed, 39 insertions(+), 14 deletions(-) diff --git a/python/tvm/script/parser_v2/frontend.py b/python/tvm/script/parser_v2/frontend.py index 6305e05bed..9229213baa 100644 --- a/python/tvm/script/parser_v2/frontend.py +++ b/python/tvm/script/parser_v2/frontend.py @@ -57,15 +57,24 @@ def _resolve(node, env, filename): return eval(compile(ast.fix_missing_locations(expression), filename, "eval"), env) +def _closure_values(function): + values = {} + for name, cell in zip(function.__code__.co_freevars, function.__closure__ or ()): + try: + values[name] = cell.cell_contents + except ValueError: + # Recursive and later-bound locals are empty until the helper is used. + pass + return values + + def _capture(obj): target = obj if inspect.isfunction(obj) else None module = inspect.getmodule(obj) env = dict(vars(module)) if module is not None else {} env.update(getattr(target, "__globals__", {})) if target is not None: - closure = inspect.getclosurevars(target) - env.update(closure.globals) - env.update(closure.nonlocals) + env.update(_closure_values(target)) filename = inspect.getsourcefile(obj) # Deferred annotations may be the only use of an enclosing local, so Python # need not put that value in the function's closure cells. @@ -133,7 +142,9 @@ def make_helper(builder, *, preserve_return=True): bound = inspect.signature(function).bind(*args, **kwargs) bound.apply_defaults() environment = ( - definition_env if options.get("hygienic", True) else _capture(function) + {**definition_env, **_closure_values(function)} + if options.get("hygienic", True) + else _capture(function) ) compiler = Compiler(function, environment) node = compiler.tree.body[0] @@ -394,16 +405,6 @@ class Compiler: def rewrite_expression(node): node = scope.rewrite(node) if scope is not None else copy.deepcopy(node) - if isinstance(node, ast.Call): - target = scope._resolve(node.func) if scope is not None else None - try: - replacement = getattr(builder, "__tvm_call_overrides__", {}).get(target) - except TypeError: - replacement = None - if replacement is not None: - name = self.fresh("call") - injected[name] = replacement - node.func = ast.copy_location(ast.Name(name, ast.Load()), node.func) method = None values = [] if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.Not): @@ -419,6 +420,20 @@ class Compiler: ) return node + def rewrite_iterable(node): + node = copy.deepcopy(node) + if isinstance(node, ast.Call): + target = scope._resolve(node.func) if scope is not None else None + try: + replacement = getattr(builder, "__tvm_call_overrides__", {}).get(target) + except TypeError: + replacement = None + if replacement is not None: + name = self.fresh("call") + injected[name] = replacement + node.func = ast.copy_location(ast.Name(name, ast.Load()), node.func) + return node + signature_values = {} if scope is not None: for name, value in scope.symbols.items(): @@ -436,6 +451,7 @@ class Compiler: signature_names=set(bound_names), signature_values=signature_values, expression_rewriter=rewrite_expression, + iterable_rewriter=rewrite_iterable, nested_function=nested_statement, preserve_return=preserve_return, ) diff --git a/python/tvm/script/parser_v2/transform.py b/python/tvm/script/parser_v2/transform.py index 674afc823b..4746d8a788 100644 --- a/python/tvm/script/parser_v2/transform.py +++ b/python/tvm/script/parser_v2/transform.py @@ -42,6 +42,7 @@ class Transformer(ast.NodeTransformer): nested_function=None, preserve_return=False, signature_values=None, + iterable_rewriter=None, ): self.filename = filename self.environment = dict(environment) @@ -54,6 +55,7 @@ class Transformer(ast.NodeTransformer): self.bound = set(self.signature_names) self.optional = {} self.expression_rewriter = expression_rewriter + self.iterable_rewriter = iterable_rewriter self.nested_function = nested_function self.preserve_return = preserve_return @@ -680,6 +682,8 @@ class Transformer(ast.NodeTransformer): if node.orelse: self._error(node, "A construction loop does not support an else clause") frame, values = self.fresh("loop"), self.fresh("indices") + if self.iterable_rewriter is not None: + node.iter = self.iterable_rewriter(node.iter) context = self._assign( frame, self._operation("For", [self._expression(node.iter)], node), node ) diff --git a/python/tvm/tirx/script/builder/ir.py b/python/tvm/tirx/script/builder/ir.py index 4364d5102d..4cb399c04f 100644 --- a/python/tvm/tirx/script/builder/ir.py +++ b/python/tvm/tirx/script/builder/ir.py @@ -3098,6 +3098,8 @@ def register_script_namespace(name: str, namespace: object) -> object: for module_name in [ "tvm.tirx.script.builder", + "tvm.tirx.script.builder_v2", + "tvm.tirx.script.v2", "tvm.tirx.script.parser", "tvm.tirx.script", "tvm.script.tirx", diff --git a/python/tvm/tirx/script/v2.py b/python/tvm/tirx/script/v2.py index 2b1668f082..de66be6956 100644 --- a/python/tvm/tirx/script/v2.py +++ b/python/tvm/tirx/script/v2.py @@ -22,6 +22,7 @@ import sys as _sys from tvm.script.parser_v2.frontend import make_decorator as _make_decorator from tvm.script.parser_v2.frontend import make_helper as _make_helper from tvm.script.parser_v2.frontend import register_namespace as _register_namespace +from tvm.tirx.layout import Axis as _Axis from . import builder_v2 as _builder from . import tile as _tile @@ -42,3 +43,5 @@ macro = _make_helper(_builder, preserve_return=False) _register_namespace("T", _sys.modules[__name__]) _register_namespace("tirx", _sys.modules[__name__]) + +_register_namespace("Axis", _Axis)
