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 801caf74989c2402662ffeb8784c933cfa517bb5
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 17:35:12 2026 +0000

    Collect registered declarations in the shared source prescan
---
 python/tvm/script/ir_builder/construction.py       |  15 ++-
 python/tvm/script/parser/protocol.py               |   9 +-
 python/tvm/script/parser/transpile.py              | 113 ++++++++++++++++-----
 .../relax/test_backend_dispatch_sort_scan.py       |  10 --
 .../relax/test_backend_transform_shape_lower.py    |  25 ++---
 tests/python/relax/test_blockbuilder_core.py       |   2 +-
 tests/python/relax/test_transform_normalize.py     |   2 -
 tests/python/relax/test_tvmscript_parser.py        |  74 +++++++-------
 tests/python/relax/test_vm_build.py                |   5 -
 tests/python/tvmscript/test_parser.py              |  49 +++++++++
 .../python/tvmscript/test_tvmscript_parser_tir.py  |  22 +++-
 11 files changed, 214 insertions(+), 112 deletions(-)

diff --git a/python/tvm/script/ir_builder/construction.py 
b/python/tvm/script/ir_builder/construction.py
index 5fec914630..d200dc4575 100644
--- a/python/tvm/script/ir_builder/construction.py
+++ b/python/tvm/script/ir_builder/construction.py
@@ -174,10 +174,11 @@ class FunctionRecord:
         return value
 
     def predeclare(self, name, annotation, location=None):
-        """Reserve an explicitly typed scalar parameter before dependent 
shapes.
+        """Reserve an explicit symbol type before dependent signature shapes.
 
         name is its source name; annotation is a re-evaluated declaration 
callable,
-        TypeVarDecl or primitive type. location is optional source data. 
Returns
+        TypeVarDecl, primitive type, or a registered body declaration's dtype
+        spelling. location is optional source data. Returns
         its canonical primitive symbol; raises TypeError for other annotations 
or
         incompatible prior declarations. Runs inside declaration() only.
         """
@@ -186,8 +187,14 @@ class FunctionRecord:
         if isinstance(annotation, TypeVarDecl):
             annotation = annotation.ty
         if ir.is_prim_var(annotation):
-            return self.symbols.bind(name, annotation)
-        return self.symbols.resolve(name, annotation, 
span=protocol.source_span(location))
+            value = self.symbols.bind(name, annotation)
+        else:
+            value = self.symbols.resolve(name, annotation, 
span=protocol.source_span(location))
+        # A body declaration is already an explicit binding, even in a function
+        # with no parameters. Only symbols newly introduced by the return
+        # annotation itself are unbound return-only symbols.
+        self.signature_symbols = set(self.symbols.symbols)
+        return value
 
     def symbol(self, name):
         """Return a declared symbol by source name after signature 
construction.
diff --git a/python/tvm/script/parser/protocol.py 
b/python/tvm/script/parser/protocol.py
index e8082cc41d..6dda5614f4 100644
--- a/python/tvm/script/parser/protocol.py
+++ b/python/tvm/script/parser/protocol.py
@@ -158,15 +158,18 @@ class DeclarationArguments(NamedTuple):
 
 
 def register_type_var_decl(constructor, *, value_parameter="expr", dtype=None):
-    """Register a declaration-capable constructor for parameter annotations.
+    """Register a constructor that declares a variable usable as a type 
variable.
 
     ``constructor`` is a callable supporting attribute assignment;
     ``value_parameter`` defaults to "expr" and names its optional value 
argument.
     ``dtype`` is the explicit dtype spelling, default None. Returns the same
     callable after attaching immutable DeclarationArguments. The signature
     prepass recognizes this marker so an explicit later parameter type is 
reserved
-    before an earlier shape string refers to it. This is not assignment-pattern
-    registration and never hoists body declarations. Calls are never
+    before an earlier shape string refers to it. The shared name prescan also
+    recognizes direct unconditional zero-argument body calls when dtype is a
+    string: it emits that spelling as builder predeclaration data, never 
evaluates
+    the call early. Argument-bearing/effectful expressions and nested control
+    scopes are excluded. The original call/bind remains in body order. Calls 
are never
     wrapped/evaluated and no frame is entered. Unsupported attribute assignment
     raises the ordinary Python error. Registration persists with the callable.
     """
diff --git a/python/tvm/script/parser/transpile.py 
b/python/tvm/script/parser/transpile.py
index 27d1926e72..f96a02a396 100644
--- a/python/tvm/script/parser/transpile.py
+++ b/python/tvm/script/parser/transpile.py
@@ -36,6 +36,11 @@ class NameCollector(ast.NodeVisitor):
     names is owned by the enclosing transpilation unit and populated in place;
     function/class/import/parameter names and parseable expression strings are
     reserved before any generated binding is allocated. No source name changes.
+    Each FunctionDef also owns _tvm_type_var_declarations: ordered tuples of
+    (source name, candidate expression AST, location AST, body-assignment 
flag).
+    These are syntax candidates, filtered by registered constructor metadata 
when
+    emitted. Signature candidates precede direct unconditional body candidates;
+    nested functions receive independent lists. No expression executes here.
     """
 
     def __init__(self, names):
@@ -51,10 +56,44 @@ class NameCollector(ast.NodeVisitor):
     def visit_FunctionDef(self, node):
         # Reserve definition names even if no ast.Name refers to them.
         self.names.setdefault(node.name, 0)
+        node._tvm_type_var_declarations = [
+            (parameter.arg, parameter.annotation, parameter, False)
+            for parameter in [*node.args.posonlyargs, *node.args.args, 
*node.args.kwonlyargs]
+            if parameter.annotation is not None
+        ]
+        # Only direct statements belong to this function's declaration prescan.
+        # Do not descend through conditions, loops, with scopes or nested defs.
+        for statement in node.body:
+            if isinstance(statement, ast.Assign):
+                for target in statement.targets:
+                    node._tvm_type_var_declarations.extend(
+                        self._body_declarations(target, statement.value)
+                    )
+            elif isinstance(statement, ast.AnnAssign) and statement.value is 
not None:
+                node._tvm_type_var_declarations.extend(
+                    self._body_declarations(statement.target, statement.value)
+                )
+            elif isinstance(statement, ast.Return | ast.Raise):
+                break
         self.generic_visit(node)
 
+    @staticmethod
+    def _body_declarations(target, value):
+        # Pattern: n = X.int64() or m, n = X.int64(), X.int64(). No argument
+        # expressions, starred unpacking or ordinary value-producing calls 
qualify.
+        if isinstance(target, ast.Name) and isinstance(value, ast.Call):
+            if not value.args and not value.keywords:
+                yield target.id, value, target, True
+        elif isinstance(target, ast.Tuple | ast.List) and isinstance(value, 
ast.Tuple | ast.List):
+            if len(target.elts) == len(value.elts):
+                for lhs, rhs in zip(target.elts, value.elts):
+                    yield from NameCollector._body_declarations(lhs, rhs)
+
     visit_AsyncFunctionDef = visit_FunctionDef
-    visit_ClassDef = visit_FunctionDef
+
+    def visit_ClassDef(self, node):
+        self.names.setdefault(node.name, 0)
+        self.generic_visit(node)
 
     def visit_alias(self, node):
         self.names.setdefault(node.asname or node.name.split(".")[0], 0)
@@ -990,35 +1029,41 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                     parameter,
                 )
             )
-        # Pattern: f(A: X.Buffer(("n",)), n: X.int64) -> reserve the explicit
-        # scalar annotation before resolving signature strings. This is solely 
a
-        # parameter-annotation prepass; body assignments stay in execution 
order.
-        for parameter in node.args.args:
-            annotation = parameter.annotation
-            if annotation is None:
-                continue
-            constructor = self._resolve(
-                annotation.func if isinstance(annotation, ast.Call) else 
annotation
-            )
-            if (
-                getattr(constructor, "__tvm_type_var_decl__", None) is not None
-                or getattr(constructor, "__tvm_parameter_dtype__", None) is 
not None
+        # Pattern: f(A: X.Buffer(("n",)), n: X.int64), or a direct body
+        # n = X.int64(): reserve the declared type before signature shapes.
+        # The shared name prescan collects candidates; registered metadata 
alone
+        # selects declarations here. Body constructors are NOT executed early:
+        # emit their registered dtype spelling and let TypeVarFrame construct 
it.
+        # The original body call and uniform bind_ remain at their source 
position.
+        local_names = self._assigned_names(node.body) | {arg.arg for arg in 
node.args.args}
+        for name, annotation, location, in_body in getattr(node, 
"_tvm_type_var_declarations", ()):
+            target = annotation.func if isinstance(annotation, ast.Call) else 
annotation
+            if in_body and any(
+                isinstance(item, ast.Name) and item.id in local_names for item 
in ast.walk(target)
             ):
-                declaration.append(
-                    self._statement(
-                        self._call(
-                            record,
-                            "predeclare",
-                            [
-                                ast.Constant(parameter.arg),
-                                copy.deepcopy(annotation),
-                                self.span(parameter),
-                            ],
-                            parameter,
-                        ),
-                        parameter,
-                    )
+                continue  # A local namespace/callable binding shadows outer 
metadata.
+            constructor = self._resolve(target)
+            metadata = getattr(constructor, "__tvm_type_var_decl__", None)
+            parameter_dtype = getattr(constructor, "__tvm_parameter_dtype__", 
None)
+            if in_body:
+                if metadata is None or not isinstance(metadata.dtype, str):
+                    continue
+                value = ast.copy_location(ast.Constant(metadata.dtype), 
annotation)
+            elif metadata is not None or parameter_dtype is not None:
+                value = copy.deepcopy(annotation)
+            else:
+                continue
+            declaration.append(
+                self._statement(
+                    self._call(
+                        record,
+                        "predeclare",
+                        [ast.Constant(name), value, self.span(location)],
+                        location,
+                    ),
+                    location,
                 )
+            )
         for parameter in node.args.args:
             if parameter.annotation is None:
                 self._error(parameter, f"Parameter {parameter.arg!r} requires 
an annotation")
@@ -1055,6 +1100,18 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             annotation = rewrite_expression(
                 node.returns, self._resolve, builder, self.filename, 
annotation=True
             )
+            # Return expression strings also assign lexical names. Seed any
+            # enclosing binding before earlier parameter expressions read it.
+            introduced_names.update(
+                call.args[0].value
+                for call in ast.walk(annotation)
+                if isinstance(call, ast.Call)
+                and isinstance(call.func, ast.Attribute)
+                and call.func.attr == "resolve_type_var"
+                and call.args
+                and isinstance(call.args[0], ast.Constant)
+                and isinstance(call.args[0].value, str)
+            )
             declaration.append(
                 self._statement(
                     self._call(record, "returns", 
[self._expression(annotation)], node.returns),
diff --git a/tests/python/relax/test_backend_dispatch_sort_scan.py 
b/tests/python/relax/test_backend_dispatch_sort_scan.py
index 1ef80f0f1e..4ae1269c47 100644
--- a/tests/python/relax/test_backend_dispatch_sort_scan.py
+++ b/tests/python/relax/test_backend_dispatch_sort_scan.py
@@ -73,7 +73,6 @@ def test_dispatch_scanop_cuda():
     lowered to the packed func `"gpu_2d_continuous_cumsum"`.
     """
 
-    # Preserve the explicitly i64 shapes used by the reference lowering.
     m = T.int64()
 
     @I.ir_module
@@ -122,9 +121,6 @@ def test_dispatch_scanop_cuda():
 
 
 def test_dispatch_sort():
-    # Preserve the explicitly i64 shapes used by the reference lowering.
-    m = T.int64()
-
     @I.ir_module
     class Before:
         I.module_global_infos({"vdevice": [I.vdevice("llvm", 0)]})
@@ -222,9 +218,6 @@ def test_dispatch_sort_cuda():
 
 
 def test_dispatch_argsort():
-    # Preserve the explicitly i64 shapes used by the reference lowering.
-    m = T.int64()
-
     @I.ir_module
     class Before:
         I.module_global_infos({"vdevice": [I.vdevice("llvm", 0)]})
@@ -318,9 +311,6 @@ def test_dispatch_argsort_cuda():
 
 
 def test_dispatch_topk():
-    # Preserve the explicitly i64 shapes used by the reference lowering.
-    m = T.int64()
-
     @I.ir_module
     class Before:
         I.module_global_infos({"vdevice": [I.vdevice("llvm", 0)]})
diff --git a/tests/python/relax/test_backend_transform_shape_lower.py 
b/tests/python/relax/test_backend_transform_shape_lower.py
index b6dc0bb1cc..a2455c846a 100644
--- a/tests/python/relax/test_backend_transform_shape_lower.py
+++ b/tests/python/relax/test_backend_transform_shape_lower.py
@@ -173,9 +173,6 @@ def test_symbolic_compute():
     MS = MatchShapeCode
     MK = MakeShapeCode
 
-    # These fixtures exercise explicit i64 arithmetic, independent of the 
script default.
-    n, m = T.int64(), T.int64()
-
     @tvm.script.ir_module
     class Before:
         @R.function
@@ -463,8 +460,6 @@ def test_return_match_check_with_new_expr():
     expression to be computed.
     """
     MS = MatchShapeCode
-
-    # These fixtures exercise explicit i64 arithmetic, independent of the 
script default.
     n = T.int64()
 
     @tvm.script.ir_module
@@ -546,9 +541,6 @@ def test_symbolic_shape_multiple_function():
     MS = MatchShapeCode
     MK = MakeShapeCode
 
-    # These fixtures exercise explicit i64 arithmetic, independent of the 
script default.
-    m, n, m2, n2 = (T.int64() for _ in range(4))
-
     @I.ir_module
     class Before:
         @R.function
@@ -559,10 +551,10 @@ def test_symbolic_shape_multiple_function():
             return A
 
         @R.function
-        def fn2(A: R.Tensor(("n2", "m2"), dtype="float32")):
+        def fn2(A: R.Tensor(("n", "m"), dtype="float32")):
             R.func_attr({"relax.force_pure": True})
-            n2 = T.int64()
-            m2 = T.int64()
+            n = T.int64()
+            m = T.int64()
             return A
 
     # slot assignment:
@@ -610,12 +602,10 @@ def test_symbolic_shape_multiple_function():
             return A
 
         @R.function
-        def fn2(A: R.Tensor(("n2", "m2"), dtype="float32")) -> R.Tensor(
-            ("n2", "m2"), dtype="float32"
-        ):
+        def fn2(A: R.Tensor(("n", "m"), dtype="float32")) -> R.Tensor(("n", 
"m"), dtype="float32"):
             R.func_attr({"relax.force_pure": True})
-            n2 = T.int64()
-            m2 = T.int64()
+            n = T.int64()
+            m = T.int64()
             shape_heap: R.Tensor(dtype="int64", ndim=1) = 
R.call_builtin_with_ctx(
                 "vm.builtin.alloc_shape_heap",
                 (R.prim_value(2),),
@@ -741,9 +731,6 @@ def test_check_lifted_weights():
 def test_check_weights_with_dynamic_shape():
     MS = MatchShapeCode
 
-    # These fixtures exercise explicit i64 arithmetic, independent of the 
script default.
-    n = T.int64()
-
     @I.ir_module
     class Before:
         @R.function
diff --git a/tests/python/relax/test_blockbuilder_core.py 
b/tests/python/relax/test_blockbuilder_core.py
index f163ed1e8c..335284f575 100644
--- a/tests/python/relax/test_blockbuilder_core.py
+++ b/tests/python/relax/test_blockbuilder_core.py
@@ -14,7 +14,7 @@
 # KIND, either express or implied.  See the License for the
 # specific language governing permissions and limitations
 # under the License.
-# ruff: noqa: F841
+# ruff: noqa: F401, F841
 """Block builder unit test"""
 
 # The test here do not depend on tvmscript to cover most basic features
diff --git a/tests/python/relax/test_transform_normalize.py 
b/tests/python/relax/test_transform_normalize.py
index 6100f1201d..d136c80df3 100644
--- a/tests/python/relax/test_transform_normalize.py
+++ b/tests/python/relax/test_transform_normalize.py
@@ -105,8 +105,6 @@ def test_normalize_if():
 
 def test_normalize_no_op():
     # the normalize pass should be no-op for IR in ANF
-    m, n = T.int64(), T.int64()
-
     @tvm.script.ir_module
     class ANFMod1:
         @R.function
diff --git a/tests/python/relax/test_tvmscript_parser.py 
b/tests/python/relax/test_tvmscript_parser.py
index d52ac94b90..cf52a916b1 100644
--- a/tests/python/relax/test_tvmscript_parser.py
+++ b/tests/python/relax/test_tvmscript_parser.py
@@ -414,29 +414,31 @@ def test_relax_shape_to_tensor():
 def test_symbolic_shape():
     @R.function
     def foo(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), 
"float32"):
-        m = T.int32()
-        n = T.int32()
+        m = T.int64()
+        n = T.int64()
         gv0 = R.call_dps_packed("extern_func", x, R.Tensor((m, n), 
dtype="float32"))
         return gv0
 
     @R.function
     def bar(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), 
"float32"):
-        m = T.int32()
-        n = T.int32()
+        m = T.int64()
+        n = T.int64()
         gv0 = R.call_dps_packed("extern_func", x, R.Tensor((m, n), 
dtype="float32"))
         return gv0
 
     with pytest.raises(tvm.error.DiagnosticError):
 
         @R.function
-        def mismatch_dtype(x: R.Tensor(("m", "n"), "float32")) -> 
R.Tensor(None, "float32", ndim=2):
-            m = T.int32()
-            n = T.int64()  # Conflicts with the newly introduced i32 shape 
symbol.
+        def mismatch_dtype(x: R.Tensor(("m", "n"), "float32"), n: T.int64) -> 
R.Tensor(
+            None, "float32", ndim=2
+        ):
+            m = T.int64()
+            n = T.int32()  # Conflicts with the explicit parameter declaration.
             gv0 = R.call_dps_packed("extern_func", x, R.Tensor((m, n), 
dtype="float32"))
             return gv0
 
     def _expected(name: str):
-        n, m = tirx.Var("n", "int32"), tirx.Var("m", "int32")
+        n, m = tirx.Var("n", "int64"), tirx.Var("m", "int64")
         x = relax.Var("x", R.Tensor([m, n], "float32"))
         bb = relax.BlockBuilder()
         with bb.function(name, (x,)):
@@ -888,7 +890,7 @@ def test_annotation():
         y: R.Tensor(("m",), "float32"),
         r: R.Tensor(dtype="int64"),
     ) -> R.Any:
-        m = T.int32()
+        m = T.int64()
         z: R.Tensor((32, m), "float32") = R.multiply(x, y)
         w: R.Tensor(ndim=2) = R.multiply(z, z)
         q: R.Tensor = R.add(w, w)
@@ -980,13 +982,13 @@ def test_call_tir_with_tir_var():
         def main(
             dumb_param: R.Tensor(("n",), "float32"), x: R.Tensor(("n * 2",), 
"float32")
         ) -> R.Tensor(("n * 2",), "float32"):
-            n = T.int32()
+            n = T.int64()
             cls = Module
             y = R.call_tir(cls.copy, (x, n), R.Tensor((n * 2,), 
dtype="float32"))
             return y
 
         @T.prim_func(s_tir=True)
-        def copy(var_x: T.handle, n: T.int32, var_y: T.handle):
+        def copy(var_x: T.handle, n: T.int64, var_y: T.handle):
             X = T.match_buffer(var_x, (n * 2,), dtype="float32")
             Y = T.match_buffer(var_y, (n * 2,), dtype="float32")
             for i in T.grid(n * 2):
@@ -1344,7 +1346,7 @@ def test_computed_prim_value_as_branch_condition():
 
     @R.function
     def func(x: R.Tensor(["N"], "float32")):
-        N = T.int32()
+        N = T.int64()
         if R.prim_value(N % 16 == 0):
             out = R.call_pure_packed("fast_vectorized_impl", x, ty_args=[x.ty])
         else:
@@ -1363,7 +1365,7 @@ def test_tir_expr_as_branch_condition():
 
     @R.function(private=True)
     def sugared(x: R.Tensor(["N"], "float32")):
-        N = T.int32()
+        N = T.int64()
         if N % 16 == 0:
             out = R.call_pure_packed("fast_vectorized_impl", x, ty_args=[x.ty])
         else:
@@ -1372,7 +1374,7 @@ def test_tir_expr_as_branch_condition():
 
     @R.function(private=True)
     def unsugared(x: R.Tensor(["N"], "float32")):
-        N = T.int32()
+        N = T.int64()
         if R.prim_value(N % 16 == 0):
             out = R.call_pure_packed("fast_vectorized_impl", x, ty_args=[x.ty])
         else:
@@ -1417,7 +1419,7 @@ def test_computed_prim_value_as_assert_condition():
 
     @R.function(pure=False)
     def func(x: R.Tensor(["N"], "float32")):
-        N = T.int32()
+        N = T.int64()
         _ = R.assert_op(R.prim_value(N % 16 == 0))
         out = R.call_packed("fast_vectorized_impl", x, ty_args=[x.ty])
         return out
@@ -1435,14 +1437,14 @@ def test_tir_expr_as_assert_condition():
 
     @R.function(pure=False, private=True)
     def sugared(x: R.Tensor(["N"], "float32")):
-        N = T.int32()
+        N = T.int64()
         _ = R.assert_op(N % 16 == 0)
         out = R.call_packed("fast_vectorized_impl", x, ty_args=[x.ty])
         return out
 
     @R.function(pure=False, private=True)
     def unsugared(x: R.Tensor(["N"], "float32")):
-        N = T.int32()
+        N = T.int64()
         _ = R.assert_op(R.prim_value(N % 16 == 0))
         out = R.call_packed("fast_vectorized_impl", x, ty_args=[x.ty])
         return out
@@ -1468,7 +1470,7 @@ def 
test_erase_to_well_defined_keeps_variables_exposed_by_tensor_shape():
     @R.function
     def foo(x: R.Tensor(["m", "n"])):
         q = x
-        m, n = T.int32(), T.int32()
+        m, n = T.int64(), T.int64()
         z = R.match_cast(q, R.Tensor((m, n)))
         w = z
         return w
@@ -1481,7 +1483,7 @@ def 
test_erase_to_well_defined_keeps_variants_exposed_by_shape_expr():
     @R.function
     def foo(x: R.Tensor, _: R.Shape(["m", "n"])):
         q = x
-        m, n = T.int32(), T.int32()
+        m, n = T.int64(), T.int64()
         z = R.match_cast(q, R.Tensor((m, n)))
         w = z
         return w
@@ -1497,7 +1499,7 @@ def test_erase_to_well_defined_infers_from_shape_expr():
         @R.function
         def subroutine(x: R.Tensor, _: R.Shape(["m", "n"])) -> R.Tensor(["m", 
"n"]):
             q = x
-            m, n = T.int32(), T.int32()
+            m, n = T.int64(), T.int64()
             z = R.match_cast(q, R.Tensor((m, n)))
             w = z
             return w
@@ -1539,8 +1541,8 @@ def test_symbolic_vars_in_tensor_shape_with_usage_first():
         return z
 
     m = tirx.Var("m", "int32")
-    x = relax.Var("x", R.Tensor([m + 1], "float32"))
-    y = relax.Var("y", R.Tensor([m, 1], "float32"))
+    x = relax.Var("x", relax.TensorType([m + 1], "float32"))
+    y = relax.Var("y", relax.TensorType([m, 1], "float32"))
     bb = relax.BlockBuilder()
     with bb.function("foo", (x, y)):
         z = bb.emit(relax.op.add(x, y))
@@ -1556,13 +1558,13 @@ def 
test_symbolic_vars_in_tensor_shape_with_definition_first():
     def bar(x: R.Tensor(("m",), "float32"), y: R.Tensor(("T.max(m, 20)",), 
"float32")) -> R.Tensor(
         ("T.max(m, 20) + 1",), "float32"
     ):
-        m = T.int32()
+        m = T.int64()
         z = R.call_dps_packed("test_intrin", (x, y), R.Tensor((T.max(m, 20) + 
1,), dtype="float32"))
         return z
 
-    m = tirx.Var("m", "int32")
-    x = relax.Var("x", R.Tensor([m], "float32"))
-    y = relax.Var("y", R.Tensor([tirx.max(m, 20)], "float32"))
+    m = tirx.Var("m", "int64")
+    x = relax.Var("x", relax.TensorType([m], "float32"))
+    y = relax.Var("y", relax.TensorType([tirx.max(m, 20)], "float32"))
     bb = relax.BlockBuilder()
     with bb.function("bar", (x, y)):
         z = bb.emit(
@@ -1677,11 +1679,11 @@ def test_symbolic_vars_in_shape():
 
     @R.function
     def baz(x: R.Shape(("m",)), y: R.Tensor(("m * 2",), "float32")):
-        m = T.int32()
+        m = T.int64()
         z = R.call_dps_packed("test_intrin", y, R.Tensor((m * 2,), 
dtype="float32"))
         return z
 
-    m = tirx.Var("m", "int32")
+    m = tirx.Var("m", "int64")
     x = relax.Var("x", relax.ShapeType([m]))
     y = relax.Var("y", relax.TensorType([m * 2], "float32"))
     bb = relax.BlockBuilder()
@@ -1760,8 +1762,8 @@ def test_arith_operators():
 def test_memory_ops():
     @R.function
     def foo(x: R.Tensor(("m", "n"), dtype="float32")):
-        m = T.int32()
-        n = T.int32()
+        m = T.int64()
+        n = T.int64()
         storage = R.memory.alloc_storage(
             R.shape([4 * m * n]), virtual_device_index=0, 
storage_scope="global", dtype="float32"
         )
@@ -1776,8 +1778,8 @@ def test_memory_ops():
 def test_vm_ops():
     @R.function(pure=False)
     def foo(x: R.Tensor(("m", "n"), dtype="float32")):
-        m = T.int32()
-        n = T.int32()
+        m = T.int64()
+        n = T.int64()
         storage = R.vm.alloc_storage(R.shape([4 * m * n]), 
runtime_device_index=0, dtype="uint8")
         alloc = R.vm.alloc_tensor(storage, offset=0, shape=R.shape([m, n]), 
dtype="float32")
         tensor = R.builtin.alloc_tensor(R.shape([m, n]), dtype="float32", 
runtime_device_index=0)
@@ -2367,7 +2369,7 @@ def test_function_attributes_are_defined():
         @R.function
         def subroutine(x: R.Tensor, _: R.Shape(["m", "n"])) -> R.Tensor(["m", 
"n"]):
             q = x
-            m, n = T.int32(), T.int32()
+            m, n = T.int64(), T.int64()
             z = R.match_cast(q, R.Tensor((m, n)))
             w = z
             return w
@@ -2472,7 +2474,7 @@ def test_shared_meta_var_skips_relax_bindings():
 
     @R.function(private=True)
     def func(A: R.Tensor(["N"], "float32")):
-        N: R.Prim("int32") = T.int32()
+        N: R.Prim("int64") = T.int64()
         via_i = I.meta_var(N)
         via_t = T.meta_var(via_i)
         output = R.reshape(A, R.shape([via_t]))
@@ -2540,7 +2542,7 @@ def 
test_conditional_may_use_symbolic_variables_from_function_scope():
         B: R.Tensor(["N"], "float32"),
         cond: R.Prim("bool"),
     ) -> R.Tensor(["N"], "float32"):
-        N = T.int32()
+        N = T.int64()
 
         if cond:
             out: R.Tensor([N], "float32") = A + B
@@ -2555,7 +2557,7 @@ def 
test_conditional_may_use_symbolic_variables_from_function_scope():
         B: R.Tensor(["N"], "float32"),
         cond: R.Prim("bool"),
     ):
-        N = T.int32()
+        N = T.int64()
         if cond:
             out = A + B
         else:
diff --git a/tests/python/relax/test_vm_build.py 
b/tests/python/relax/test_vm_build.py
index 4fbd7221ad..7a10f81538 100644
--- a/tests/python/relax/test_vm_build.py
+++ b/tests/python/relax/test_vm_build.py
@@ -212,8 +212,6 @@ def test_vm_compile_e2e(exec_mode):
 
 
 def test_vm_compile_e2e_func_param_with_shape(exec_mode):
-    m, n, k = T.int64(), T.int64(), T.int64()
-
     @tvm.script.ir_module
     class TestVMCompileE2E2:
         @T.prim_func(s_tir=True)
@@ -562,8 +560,6 @@ def test_vm_relax_symbolic_shape(exec_mode):
 
 
 def test_vm_relax_symbolic_shape_tuple(exec_mode):
-    m, n = T.int64(), T.int64()
-
     @I.ir_module(s_tir=True)
     class mod:
         @R.function
@@ -931,7 +927,6 @@ class TestVMSetInput:
 
 
 def test_multi_systemlib(exec_mode):
-    m, n = T.int64(), T.int64()
     pytest.importorskip("cloudpickle")  # needed by popen_pool.PopenWorker
 
     @tvm.script.ir_module
diff --git a/tests/python/tvmscript/test_parser.py 
b/tests/python/tvmscript/test_parser.py
index 1e814a7c3e..06f81e5d72 100644
--- a/tests/python/tvmscript/test_parser.py
+++ b/tests/python/tvmscript/test_parser.py
@@ -30,6 +30,7 @@ from types import ModuleType, SimpleNamespace
 import pytest
 
 from tvm import ir
+from tvm.error import DiagnosticError
 from tvm.ir.prim import Cast
 from tvm.script import parser
 from tvm.script.ir_builder import IRBuilder, protocol
@@ -495,3 +496,51 @@ def test_parser_metadata_registration_boundary():
     assert protocol.expr_str_policy(constructor) is policy
     with IRBuilder(), pytest.raises(TypeError, match="concrete symbols"):
         constructor(("n",))
+
+
+def test_shared_prescan_preserves_body_effects_and_scope():
+    from tvm.script import tirx as T
+    from tvm.script.ir_builder import TypeVarDecl
+
+    calls = []
+
+    def scalar(expr=None):
+        calls.append("declaration")
+        return TypeVarDecl("int64")
+
+    def effect():
+        calls.append("effect")
+        return 0
+
+    syntax_protocol.register_type_var_decl(scalar, dtype="int64")
+    source = """
[email protected]_func
+def main(A: T.Buffer(("n",), "float32")):
+    T.evaluate(effect())
+    n = scalar()
+    T.evaluate(n)
+"""
+    function = parser.parse(source, {"T": T, "scalar": scalar, "effect": 
effect})
+    assert calls == ["effect", "declaration"]
+    assert str(function.params[0].ty.shape[0].ty.dtype) == "int64"
+    calls.clear()
+    with pytest.raises(DiagnosticError, match="incompatible declared type"):
+        parser.parse(
+            source.replace("n = scalar()", "n = scalar(effect())"),
+            {"T": T, "scalar": scalar, "effect": effect},
+        )
+    assert calls == ["effect", "effect", "declaration"]
+
+    compiler = Compiler("""
+def outer():
+    a = scalar()
+    if condition:
+        hidden = scalar()
+    def inner():
+        b = scalar()
+""")
+    outer = compiler.tree.body[0]
+    inner = outer.body[-1]
+    assert [item[0] for item in outer._tvm_type_var_declarations] == ["a"]
+    assert [item[0] for item in inner._tvm_type_var_declarations] == ["b"]
+    assert {"outer", "inner", "a", "b", "hidden"}.issubset(compiler.name_map)
diff --git a/tests/python/tvmscript/test_tvmscript_parser_tir.py 
b/tests/python/tvmscript/test_tvmscript_parser_tir.py
index a4baa9b456..e6cabc423b 100644
--- a/tests/python/tvmscript/test_tvmscript_parser_tir.py
+++ b/tests/python/tvmscript/test_tvmscript_parser_tir.py
@@ -118,15 +118,29 @@ def main(A: T.Buffer(("n",), "float32"), n: T.int64):
     assert str(n.ty.dtype) == "int64"
 
 
-def test_tir_string_defined_symbol_does_not_take_dtype_from_body():
-    with pytest.raises(tvm.error.DiagnosticError):
-        tvm.script.from_source(
-            """
+def test_tir_string_defined_symbol_preserves_direct_body_declaration():
+    func = tvm.script.from_source(
+        """
 @T.prim_func
 def main(A: T.Buffer(("n",), "float32")):
     n = T.int64()
     T.evaluate(n)
 """
+    )
+    n = func.params[0].ty.shape[0]
+    assert str(n.ty.dtype) == "int64"
+    assert func.body.value.same_as(n)
+
+
[email protected](
+    "statement", ["if True:\n        n = T.int64()", "for i in range(1):\n     
   n = T.int64()"]
+)
+def test_tir_control_flow_declarations_are_not_prescanned(statement):
+    with pytest.raises(tvm.error.DiagnosticError, match="incompatible declared 
type"):
+        tvm.script.from_source(
+            '@T.prim_func\ndef main(A: T.Buffer(("n",), "float32")):\n    '
+            + statement
+            + "\n    T.evaluate(0)\n"
         )
 
 

Reply via email to