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)

Reply via email to