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" )
