This is an automated email from the ASF dual-hosted git repository.

tqchen pushed a commit to branch tvmscript-ast-only-transpiler
in repository https://gitbox.apache.org/repos/asf/tvm.git

commit eb4313fc5783fb9ee725f2059dab899f798b8088
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 22:02:13 2026 +0000

    [TVMScript] Resolve captured type variables before annotation arithmetic
---
 python/tvm/script/ir_builder/construction.py | 11 +++++--
 python/tvm/script/parser_v2/transpile.py     | 47 +++++++++++++++++++---------
 2 files changed, 41 insertions(+), 17 deletions(-)

diff --git a/python/tvm/script/ir_builder/construction.py 
b/python/tvm/script/ir_builder/construction.py
index a830f815bc..f0879a7c7b 100644
--- a/python/tvm/script/ir_builder/construction.py
+++ b/python/tvm/script/ir_builder/construction.py
@@ -23,6 +23,7 @@ transpiler state. Generated Python calls these operations in 
source order.
 
 from contextlib import contextmanager
 from types import SimpleNamespace
+from typing import TypeVar
 
 from tvm import ir
 
@@ -322,16 +323,22 @@ class FunctionRecord:
         Returns
         -------
         object
-            Existing canonical symbol, or fallback unchanged.
+            Existing canonical symbol, a symbol resolved from a host TypeVar,
+            or any other fallback unchanged.
 
         Notes
         -----
-        No symbols or frames are introduced and no fallback is cached.
+        A host TypeVar resolves in this record's retained symbol frame before
+        annotation arithmetic executes. Explicit predeclarations have already
+        run, so they retain their requested dtype. No frame is entered and
+        ordinary host values are neither converted nor cached.
 
         Examples
         --------
         >>> value = record.capture("n", outer_n)
         """
+        if isinstance(fallback, TypeVar):
+            return self.symbols.resolve(name)
         return self.symbols.symbols.get(name, fallback)
 
     def returns(self, annotation):
diff --git a/python/tvm/script/parser_v2/transpile.py 
b/python/tvm/script/parser_v2/transpile.py
index 320195a4ad..90137aafe3 100644
--- a/python/tvm/script/parser_v2/transpile.py
+++ b/python/tvm/script/parser_v2/transpile.py
@@ -1097,6 +1097,19 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         factory_body = [self._assign(record, record_expr, node)]
         if node.args.posonlyargs or node.args.kwonlyargs or node.args.vararg 
or node.args.kwarg:
             self._error(node, "IR signatures require ordinary named 
parameters")
+        annotation_names = {
+            item.id
+            for argument in node.args.args
+            if argument.annotation is not None
+            for item in ast.walk(argument.annotation)
+            if isinstance(item, ast.Name) and isinstance(item.ctx, ast.Load)
+        }
+        if node.returns is not None:
+            annotation_names.update(
+                item.id
+                for item in ast.walk(node.returns)
+                if isinstance(item, ast.Name) and isinstance(item.ctx, 
ast.Load)
+            )
         declaration = []
         introduced_names = set()
         for parameter in getattr(node, "type_params", []):
@@ -1156,6 +1169,7 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                     location,
                 )
             )
+        capture_position = len(declaration)
         for parameter in node.args.args:
             if parameter.annotation is None:
                 self._error(parameter, f"Parameter {parameter.arg!r} requires 
an annotation")
@@ -1210,10 +1224,23 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                     node.returns,
                 )
             )
-        for name in sorted(introduced_names - {argument.arg for argument in 
node.args.args}):
+        capture_statements = []
+        for name in sorted(
+            (annotation_names | introduced_names)
+            - {argument.arg for argument in node.args.args}
+            - {parameter.name for parameter in getattr(node, "type_params", 
[])}
+        ):
             # A real enclosing binding remains visible before a signature 
string
             # assigns the same spelling; absent names still raise before first 
use.
             condition = ast.Compare(ast.Constant(name), [ast.In()], 
[self._name(captures, node)])
+            # Establish Python locals before registered declaration expressions
+            # read namespace aliases; symbol resolution follows predeclaration.
+            initial = self._assign(
+                name, self._call(captures, "get", [ast.Constant(name)], node), 
node
+            )
+            factory_body.append(
+                self._located(ast.If(copy.deepcopy(condition), [initial], []), 
node)
+            )
             assignment = self._assign(
                 name,
                 self._call(
@@ -1224,7 +1251,10 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                 ),
                 node,
             )
-            factory_body.append(self._located(ast.If(condition, [assignment], 
[]), node))
+            capture_statements.append(self._located(ast.If(condition, 
[assignment], []), node))
+        # Opaque captures are resolved after explicit declarations and before
+        # signature expressions execute; builders alone identify host TypeVars.
+        declaration[capture_position:capture_position] = capture_statements
         factory_body.append(
             self._with(self._call(record, "declaration", [], node), 
declaration, node)
         )
@@ -1277,19 +1307,6 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         # Bare host TypeVars in annotations are resolved by constructor 
builders.
         # Capture their resulting values for free body names without inspecting
         # them here. Source locals/parameters keep Python's original shadowing.
-        annotation_names = {
-            item.id
-            for argument in node.args.args
-            if argument.annotation is not None
-            for item in ast.walk(argument.annotation)
-            if isinstance(item, ast.Name) and isinstance(item.ctx, ast.Load)
-        }
-        if node.returns is not None:
-            annotation_names.update(
-                item.id
-                for item in ast.walk(node.returns)
-                if isinstance(item, ast.Name) and isinstance(item.ctx, 
ast.Load)
-            )
         # Only syntactically referenced signature names can capture enclosing
         # symbols. An unrelated outer n must not constrain a local n 
declaration.
         capture_names = (annotation_names | introduced_names) - {

Reply via email to