From 66caf2fe4690f24f0317c3b87d4eeadb36868181 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 19:42:41 +0000 Subject: [PATCH 01/16] [Script] Introduce explicit dynamic symbolic variables Share scope-independent dynamic symbol construction between I and T, and replace implicit body declarations with explicit captures or typed function parameters. Preserve dtype information in scalar annotations and print symbols without reconstructing them from body assignments. --- .../relax/tutorials/relax_creation.py | 21 +- .../tensor_ir/tutorials/tir_creation.py | 11 +- .../mix_python_and_tvm_with_pymodule.py | 6 +- python/tvm/script/ir_builder/__init__.py | 2 + python/tvm/script/ir_builder/ir.py | 25 ++- .../tvm/script/ir_builder/parser_protocol.py | 7 +- python/tvm/script/parser/prescan.py | 17 +- python/tvm/script/parser/protocol_registry.py | 39 +--- python/tvm/script/parser/transpile.py | 50 +++-- python/tvm/tirx/script/ir_builder/__init__.py | 26 +-- python/tvm/tirx/script/ir_builder/ir.py | 30 +-- src/relax/script/printer/dependent_type.cc | 2 +- src/relax/script/printer/function.cc | 2 +- src/relax/script/printer/tir.cc | 5 +- src/script/printer/ir/ir.cc | 8 +- src/script/printer/utils.h | 37 ++-- src/tirx/script/printer/buffer.cc | 19 +- src/tirx/script/printer/expr.cc | 11 +- src/tirx/script/printer/function.cc | 13 +- src/tirx/script/printer/utils.h | 5 +- tests/python/relax/test_tvmscript_parser.py | 187 +++++++++++------- .../relax/test_tvmscript_printer_relax.py | 44 ++--- tests/python/relax/test_tvmscript_pyfunc.py | 3 +- .../python/relax/test_tvmscript_type_vars.py | 42 +++- tests/python/script/minilang.py | 11 +- tests/python/script/test_meta_programming.py | 2 - tests/python/script/test_symbolic_shape.py | 78 ++++++-- tests/python/tirx/test_tvmscript_type_vars.py | 60 +++++- .../tvmscript/test_tvmscript_parser_tir.py | 12 +- .../tvmscript/test_tvmscript_printer_tir.py | 62 +++--- .../tvmscript/test_tvmscript_roundtrip.py | 31 +-- .../tvmscript/test_tvmscript_syntax_sugar.py | 14 +- 32 files changed, 508 insertions(+), 374 deletions(-) diff --git a/docs/deep_dive/relax/tutorials/relax_creation.py b/docs/deep_dive/relax/tutorials/relax_creation.py index f1437b66e008..949313372de3 100644 --- a/docs/deep_dive/relax/tutorials/relax_creation.py +++ b/docs/deep_dive/relax/tutorials/relax_creation.py @@ -70,12 +70,14 @@ def forward( # TensorIR functions in Relax function. +n = T.dynamic("n", "int64") +m = T.dynamic("m", "int64") + + @I.ir_module class RelaxModuleWithTIR: @Ts.prim_func def relu(x: T.handle, y: T.handle): - n = T.int64() - m = T.int64() X = T.match_buffer(x, (n, m), "float32") Y = T.match_buffer(y, (n, m), "float32") for i, j in T.grid(n, m): @@ -85,13 +87,12 @@ def relu(x: T.handle, y: T.handle): @R.function def forward( - data: R.Tensor(("n", 784), dtype="float32"), + data: R.Tensor((n, 784), dtype="float32"), w0: R.Tensor((128, 784), dtype="float32"), b0: R.Tensor((128,), dtype="float32"), w1: R.Tensor((10, 128), dtype="float32"), b1: R.Tensor((10,), dtype="float32"), - ) -> R.Tensor(("n", 10), dtype="float32"): - n = T.int64() + ) -> R.Tensor((n, 10), dtype="float32"): cls = RelaxModuleWithTIR with R.dataflow(): lv0 = R.matmul(data, R.permute_dims(w0)) + b0 @@ -165,11 +166,13 @@ def forward(self, x): # Tensor Expression(TE), TensorIR functions or other TVM packed functions. +M = T.dynamic("M", "int64") +N = T.dynamic("N", "int64") +K = T.dynamic("K", "int64") + + @Ts.prim_func def tir_linear(x: T.handle, w: T.handle, b: T.handle, z: T.handle): - M = T.int64() - N = T.int64() - K = T.int64() X = T.match_buffer(x, (M, K), "float32") W = T.match_buffer(w, (N, K), "float32") B = T.match_buffer(b, (N,), "float32") @@ -227,7 +230,7 @@ def forward(self, x): # customized pass. bb = relax.BlockBuilder() -n = T.int64() +n = T.dynamic("n", "int64") x = relax.Var("x", R.Tensor((n, 784), "float32")) fc1_weight = relax.Var("fc1_weight", R.Tensor((128, 784), "float32")) fc1_bias = relax.Var("fc1_bias", R.Tensor((128,), "float32")) diff --git a/docs/deep_dive/tensor_ir/tutorials/tir_creation.py b/docs/deep_dive/tensor_ir/tutorials/tir_creation.py index 219ded0bc6c3..5aae86ec8e53 100644 --- a/docs/deep_dive/tensor_ir/tutorials/tir_creation.py +++ b/docs/deep_dive/tensor_ir/tutorials/tir_creation.py @@ -178,15 +178,16 @@ def mm_relu( # be used to ascertain the shape and data type of a TensorIR. +# Dynamic shape definition +M = T.dynamic("M", "int32") +N = T.dynamic("N", "int32") +K = T.dynamic("K", "int32") + + @I.ir_module class DynamicShapeModule: @Ts.prim_func def mm_relu(a: T.handle, b: T.handle, c: T.handle): - # Dynamic shape definition - M = T.int32() - N = T.int32() - K = T.int32() - # Bind the input buffers with the dynamic shapes A = T.match_buffer(a, [M, K], dtype) B = T.match_buffer(b, [K, N], dtype) diff --git a/docs/how_to/tutorials/mix_python_and_tvm_with_pymodule.py b/docs/how_to/tutorials/mix_python_and_tvm_with_pymodule.py index 12b93664eb3f..709f1ad07949 100644 --- a/docs/how_to/tutorials/mix_python_and_tvm_with_pymodule.py +++ b/docs/how_to/tutorials/mix_python_and_tvm_with_pymodule.py @@ -129,12 +129,12 @@ def forward(self, x, y): # immediately — no recompilation needed. if RUN_EXAMPLE: + n = T.dynamic("n", "int32") @R.py_module class DebugModule(BasePyModule): @Ts.prim_func def matmul_tir(var_A: T.handle, var_B: T.handle, var_C: T.handle): - n = T.int32() A = T.match_buffer(var_A, (n, 4), "float32") B = T.match_buffer(var_B, (4, 3), "float32") C = T.match_buffer(var_C, (n, 3), "float32") @@ -399,12 +399,12 @@ def main( # tensors at call time, so the same module handles different sizes without recompilation. if RUN_EXAMPLE: + n = T.dynamic("n", "int64") @R.py_module class DynamicModule(BasePyModule): @Ts.prim_func def scale_tir(var_x: T.handle, var_out: T.handle): - n = T.int64() x = T.match_buffer(var_x, (n,), "float32") out = T.match_buffer(var_out, (n,), "float32") for i in T.serial(n): @@ -435,7 +435,7 @@ def add_relax( print("add_relax(len=10):", out10) # Python → TIR with symbolic output shape - n = T.int64() + n = T.dynamic("n", "int64") x7 = torch.randn(7) scaled = mod.call_tir("scale_tir", [x7], relax.TensorType((n,), "float32")) print("scale_tir(len=7):", scaled) diff --git a/python/tvm/script/ir_builder/__init__.py b/python/tvm/script/ir_builder/__init__.py index 93bcec4b8d64..96ab99308f26 100644 --- a/python/tvm/script/ir_builder/__init__.py +++ b/python/tvm/script/ir_builder/__init__.py @@ -36,6 +36,7 @@ from .ir import ( decl_function, def_function, + dynamic, ir_module, lookup_name, meta_var, @@ -62,6 +63,7 @@ "constexpr", "decl_function", "def_function", + "dynamic", "ir_module", "lookup_name", "meta_var", diff --git a/python/tvm/script/ir_builder/ir.py b/python/tvm/script/ir_builder/ir.py index 5b6e0f3b5985..7230bdba71eb 100644 --- a/python/tvm/script/ir_builder/ir.py +++ b/python/tvm/script/ir_builder/ir.py @@ -19,16 +19,37 @@ import inspect from typing import TypeVar -from tvm.ir import BaseFunc, GlobalInfo, GlobalVar +from tvm.ir import BaseFunc, GlobalInfo, GlobalVar, Var from tvm.runtime import Object as tvm_Object from . import _ffi_api -from .base import IRBuilder +from .base import IRBuilder, source_span from .frame import IRModuleFrame T = TypeVar("T") +def dynamic(name: str, dtype: str = "int64", *, span=None) -> Var: + """Create a fresh primitive symbolic variable, independently of builder scope. + + Parameters + ---------- + name : str + The symbol's display name. Repeated names do not share identity. + dtype : str + Primitive dtype, defaulting to int64. + span : Optional[Span] + Source location of the symbol. + + Returns + ------- + Var + A fresh symbol. Reuse this object to share dimensions across annotations + and function bodies, including outside an ``I.ir_module`` definition. + """ + return Var(name, dtype, source_span(span)) + + def meta_var(value: T) -> T: """Return a Python metadata value without binding, naming or relocating it. diff --git a/python/tvm/script/ir_builder/parser_protocol.py b/python/tvm/script/ir_builder/parser_protocol.py index 91f3106ae687..ae5c115938dd 100644 --- a/python/tvm/script/ir_builder/parser_protocol.py +++ b/python/tvm/script/ir_builder/parser_protocol.py @@ -782,10 +782,9 @@ def resolve_type_var_( .. code:: python - # Source - n = T.int64() + # Source: def f[n: T.int32](...) # Generated builder - n = X.resolve_type_var_("n", dtype="int64") + n = X.resolve_type_var_("n", dtype="int32") """ raise NotImplementedError @@ -873,7 +872,7 @@ class Module: # Syntax markers live in tvm.script.parser.protocol_registry. # ``constexpr(value)`` selects host evaluation in marked control flow. # ``args_policy(path, fields)`` marks expression-string arguments. -# ``register_type_var_decl(path, constructor, dtype=...)`` marks symbolic declarations. +# ``register_scalar_annotation(path, constructor, dtype=...)`` describes scalar annotations. # ``mutable_cell_decl(path)`` marks mutable storage declarations. # ``result_span(path)`` permits attaching a call's result span without a call context. # ``module_decorator(path)`` marks module declaration decorators. diff --git a/python/tvm/script/parser/prescan.py b/python/tvm/script/parser/prescan.py index 3de503356391..36177e92c97b 100644 --- a/python/tvm/script/parser/prescan.py +++ b/python/tvm/script/parser/prescan.py @@ -393,7 +393,15 @@ def visit_FunctionDef(self, node: ast.FunctionDef | ast.AsyncFunctionDef) -> Non # Builder: # n = X.resolve_type_var_("n") # ------------------------------------------------- - self._record_binding(parameter.name, parameter, "symbol", dtype="int64") + bound = getattr(parameter, "bound", None) + dtype = ( + "int64" + if bound is None or (isinstance(bound, ast.Name) and bound.id == "int") + else protocol.SCALAR_ANNOTATION_DTYPE.get( + resolve_namespace_key(bound, self.environment) + ) + ) + self._record_binding(parameter.name, parameter, "symbol", dtype=dtype) for arg in [*node.args.posonlyargs, *node.args.args, *node.args.kwonlyargs]: # -------------------- Pattern -------------------- # Python source: @@ -412,7 +420,7 @@ def visit_FunctionDef(self, node: ast.FunctionDef | ast.AsyncFunctionDef) -> Non annotation.func if isinstance(annotation, ast.Call) else annotation, self.environment, ) - dtype = protocol.TYPE_VAR_DECL.get(constructor) + dtype = protocol.SCALAR_ANNOTATION_DTYPE.get(constructor) self._record_binding( arg.arg, arg, @@ -501,10 +509,7 @@ def _collect_target( ) -> None: constructor = resolve_constructor(value, self.environment, self.bindings[self.scope]) if isinstance(target, ast.Name): - dtype = protocol.TYPE_VAR_DECL.get(constructor) - if constructor in protocol.TYPE_VAR_DECL and not value.args and not value.keywords: - self._record_binding(target.id, target, "symbol", value, dtype) - elif getattr(self.builder, "supports_mutable_declarations", True) and ( + if getattr(self.builder, "supports_mutable_declarations", True) and ( "call" in protocol.MUTABLE_CELL_DECL.get(constructor, ()) or ( annotation is not None diff --git a/python/tvm/script/parser/protocol_registry.py b/python/tvm/script/parser/protocol_registry.py index c9c059510287..396a709bb462 100644 --- a/python/tvm/script/parser/protocol_registry.py +++ b/python/tvm/script/parser/protocol_registry.py @@ -71,7 +71,7 @@ class ArgsPolicy(NamedTuple): ARGS_POLICIES: dict[str, ArgsPolicy] = {} -TYPE_VAR_DECL: dict[str, object] = {} +SCALAR_ANNOTATION_DTYPE: dict[str, object] = {} MUTABLE_CELL_DECL: dict[str, frozenset[str]] = {} RESULT_SPAN: dict[str, bool] = {} MODULE_DECORATOR: dict[str, bool] = {} @@ -208,45 +208,18 @@ def handle_call_args_policy( return (policy, policy.positional_parameters) if policy is not None else None -def register_type_var_decl( +def register_scalar_annotation( namespace_path: str, constructor: _Callable, *, dtype: object = None, ) -> _Callable: - """Register symbolic declaration syntax and return the unchanged constructor. + """Register the dtype of a fixed scalar annotation without evaluating it. - Parameters - ---------- - namespace_path : str - Canonical registered namespace alias and exported callable path, such - as ``"T.int32"``. A later registration at this path replaces its dtype. - constructor : Callable - Callable providing the eager construction operation. Registration - records its syntax path without invoking or wrapping this callable. - dtype : object, optional - Static scalar dtype forwarded to the language variant's symbol resolver. - None (the default) leaves the dtype unspecified for that resolver. - - Returns - ------- - Callable - The exact ``constructor`` object. - - Notes - ----- - Zero-argument calls denote declarations. Dictionary membership distinguishes - an unspecified dtype from an unregistered - constructor. The parser predeclares symbols needed by signatures while the - native function owns their identity. No constructor executes at registration. - - .. code:: python - - register_type_var_decl("T.int32", T.int32, dtype="int32") - # Source: n = T.int32() - # Builder: n = X.resolve_type_var_("n", dtype="int32") + Used by scalar function parameters and explicit PEP 695 symbol bounds. + This metadata does not give constructor calls special assignment semantics. """ - TYPE_VAR_DECL[namespace_path] = dtype + SCALAR_ANNOTATION_DTYPE[namespace_path] = dtype return constructor diff --git a/python/tvm/script/parser/transpile.py b/python/tvm/script/parser/transpile.py index a09b21ca19be..b0d51957b2f5 100644 --- a/python/tvm/script/parser/transpile.py +++ b/python/tvm/script/parser/transpile.py @@ -927,7 +927,7 @@ def _binding_kind(self, target: ast.Name, *, frame_value: bool = False) -> str: return "ordinary" site = self.module.prescan.sites.get(target) kind = site.kind if site is not None else "ordinary" - if kind in ("symbol", "mutable", "module_alias"): + if kind in ("mutable", "module_alias"): return kind mutable = self.module.prescan.mutable_names.get(self.function.current_scope, ()) return "mutable_update" if target.id in mutable else "ordinary" @@ -962,28 +962,13 @@ def _bind( ) -> list[ast.stmt]: # Declaration syntax has precedence; no previous/existence tracking. if isinstance(target, ast.Name): - site = self.module.prescan.sites.get(target) kind = self._binding_kind(target, frame_value=frame_value) keywords = {"name": ast.Constant(target.id)} if ty is not None: keywords["ty"] = ty if frame_value: keywords["frame_value"] = ast.Constant(True) - if kind == "symbol" and not frame_value: - # -------------------- Pattern -------------------- - # Python source: - # n = X.int64() - # - # Builder: - # n = X.resolve_type_var_("n", dtype="int64") - # ------------------------------------------------- - value = self._call_dialect( - "resolve_type_var_", - [ast.Constant(target.id)], - target, - **({"dtype": ast.Constant(site.dtype)} if site.dtype else {}), - ) - elif kind == "mutable" and not frame_value: + if kind == "mutable" and not frame_value: # -------------------- Pattern -------------------- # Python source: # x = X.local_scalar(initial) @@ -1134,12 +1119,7 @@ def visit_Assign(self, node: ast.Assign) -> ast.stmt | list[ast.stmt]: value_span = self.module.span(node.value) if ordinary else None if len(node.targets) == 1 and isinstance(node.targets[0], ast.Name): target = node.targets[0] - site = self.module.prescan.sites.get(target) - value = ( - ast.Constant(None) - if site and site.kind == "symbol" - else self._rewrite_assignment_value(node.value, ordinary=ordinary) - ) + value = self._rewrite_assignment_value(node.value, ordinary=ordinary) return self._bind(target, value, node, value_span=value_span) if len(node.targets) == 1 and isinstance(node.targets[0], ast.Subscript): # -------------------- Pattern -------------------- @@ -1855,8 +1835,13 @@ def _create_symbol_declarations( if not isinstance(parameter, getattr(ast, "TypeVar", ())): self._raise_error(parameter, "Only scalar type parameters are supported") bound = getattr(parameter, "bound", None) + dtype = self.module.prescan.sites[parameter].dtype if bound is not None and not self._is_builtin(bound, int): - self._raise_error(parameter, "A symbolic type parameter bound must be int") + if protocol.SCALAR_ANNOTATION_DTYPE.get(self._resolve(bound)) is None: + self._raise_error( + parameter, + "A symbolic type parameter bound must be int or a registered scalar dtype", + ) if getattr(parameter, "default_value", None) is not None: self._raise_error(parameter, "A symbolic type parameter cannot have a default") alias = self.module.fresh("_symbol") @@ -1865,13 +1850,16 @@ def _create_symbol_declarations( self._assign( alias, self._call_dialect( - "resolve_type_var_", [ast.Constant(parameter.name)], parameter + "resolve_type_var_", + [ast.Constant(parameter.name)], + parameter, + dtype=ast.Constant(dtype), ), parameter, ) ) for item in facts: - if item.direct and item.dtype is not None and item.kind in ("symbol", "parameter"): + if item.direct and item.dtype is not None and item.kind == "parameter": symbol = self._call_dialect( "resolve_type_var_", [ast.Constant(item.name)], @@ -2273,10 +2261,20 @@ def create_function_builder_fragments( for item in ast.walk(statement) if isinstance(item, ast.Name) } + # Prefix statements in source text execute in the outer + # builder callable. Keep their values as Python closures; + # only externally supplied bindings belong to its globals. + prefix_names = { + item.name + for scope, bindings in self.module.prescan.bindings.items() + if isinstance(scope, ast.Module) + for item in bindings + } global_names = sorted( name for name, alias in definition_aliases.items() if name == alias + and name not in prefix_names and name in referenced and name not in self.module.prescan.namespaces and name not in {parameter.arg for parameter in parameters} diff --git a/python/tvm/tirx/script/ir_builder/__init__.py b/python/tvm/tirx/script/ir_builder/__init__.py index babeaf151d32..1852ba3ef232 100644 --- a/python/tvm/tirx/script/ir_builder/__init__.py +++ b/python/tvm/tirx/script/ir_builder/__init__.py @@ -25,6 +25,7 @@ from tvm import tirx as _tir from tvm.script.ir_builder.base import at as _at from tvm.script.ir_builder.base import source_span as _source_span +from tvm.script.ir_builder import dynamic as dynamic from tvm.script.parser.protocol_registry import ARGS_POLICIES as _ARGS_POLICIES from tvm.script.parser.protocol_registry import args_policy as _args_policy from tvm.script.parser.protocol_registry import constexpr as constexpr @@ -96,31 +97,6 @@ is_type_var = _ir.is_prim_var -def type_var(name, *, dtype=None, span=None): - """Construct a fresh standalone primitive symbol. - - Parameters - ---------- - name : str - Name of the symbol. - dtype : str or PrimType, optional - Primitive type of the symbol; None selects "int64". - span : Span or source location, optional - Source location attached to the constructed IR; None leaves it unspecified. - - Returns - ------- - result : Var - The newly constructed primitive variable. - - Notes - ----- - This constructor creates a new symbol on each call. Use the language variant - resolver for symbols shared by name within a function signature. - """ - return _ir.Var(name, "int64" if dtype is None else dtype, _source_span(span)) - - @_result_span("T.Buffer") @_mutable_cell_decl("T.Buffer", syntax="parameter") @_args_policy( diff --git a/python/tvm/tirx/script/ir_builder/ir.py b/python/tvm/tirx/script/ir_builder/ir.py index 5544ca32147c..da39f8737a77 100644 --- a/python/tvm/tirx/script/ir_builder/ir.py +++ b/python/tvm/tirx/script/ir_builder/ir.py @@ -48,7 +48,9 @@ from tvm.script.parser.protocol_registry import ( mutable_cell_decl as _mutable_cell_decl, ) -from tvm.script.parser.protocol_registry import register_type_var_decl as _register_type_var_decl +from tvm.script.parser.protocol_registry import ( + register_scalar_annotation as _register_scalar_annotation, +) from tvm.target import Target # pylint: disable=unused-import @@ -2038,7 +2040,7 @@ def func_gen(name: str): """ dtype = _ffi_name_to_dtype(name) constructor = DtypeConstructor(name, dtype) - _register_type_var_decl(f"T.{dtype}", constructor, dtype=dtype) + _register_scalar_annotation(f"T.{dtype}", constructor, dtype=dtype) _mutable_cell_decl(f"T.{dtype}", syntax="annotation")(constructor) return constructor @@ -2227,29 +2229,29 @@ def add_to_parent(stmt: tir.Stmt) -> None: bfloat16 = func_gen("BFloat16") # Shorthand aliases -f16 = _register_type_var_decl("T.f16", float16, dtype="float16") +f16 = _register_scalar_annotation("T.f16", float16, dtype="float16") _mutable_cell_decl("T.f16", syntax="annotation")(f16) -f32 = _register_type_var_decl("T.f32", float32, dtype="float32") +f32 = _register_scalar_annotation("T.f32", float32, dtype="float32") _mutable_cell_decl("T.f32", syntax="annotation")(f32) -f64 = _register_type_var_decl("T.f64", float64, dtype="float64") +f64 = _register_scalar_annotation("T.f64", float64, dtype="float64") _mutable_cell_decl("T.f64", syntax="annotation")(f64) -bf16 = _register_type_var_decl("T.bf16", bfloat16, dtype="bfloat16") +bf16 = _register_scalar_annotation("T.bf16", bfloat16, dtype="bfloat16") _mutable_cell_decl("T.bf16", syntax="annotation")(bf16) -i8 = _register_type_var_decl("T.i8", int8, dtype="int8") +i8 = _register_scalar_annotation("T.i8", int8, dtype="int8") _mutable_cell_decl("T.i8", syntax="annotation")(i8) -i16 = _register_type_var_decl("T.i16", int16, dtype="int16") +i16 = _register_scalar_annotation("T.i16", int16, dtype="int16") _mutable_cell_decl("T.i16", syntax="annotation")(i16) -i32 = _register_type_var_decl("T.i32", int32, dtype="int32") +i32 = _register_scalar_annotation("T.i32", int32, dtype="int32") _mutable_cell_decl("T.i32", syntax="annotation")(i32) -i64 = _register_type_var_decl("T.i64", int64, dtype="int64") +i64 = _register_scalar_annotation("T.i64", int64, dtype="int64") _mutable_cell_decl("T.i64", syntax="annotation")(i64) -u8 = _register_type_var_decl("T.u8", uint8, dtype="uint8") +u8 = _register_scalar_annotation("T.u8", uint8, dtype="uint8") _mutable_cell_decl("T.u8", syntax="annotation")(u8) -u16 = _register_type_var_decl("T.u16", uint16, dtype="uint16") +u16 = _register_scalar_annotation("T.u16", uint16, dtype="uint16") _mutable_cell_decl("T.u16", syntax="annotation")(u16) -u32 = _register_type_var_decl("T.u32", uint32, dtype="uint32") +u32 = _register_scalar_annotation("T.u32", uint32, dtype="uint32") _mutable_cell_decl("T.u32", syntax="annotation")(u32) -u64 = _register_type_var_decl("T.u64", uint64, dtype="uint64") +u64 = _register_scalar_annotation("T.u64", uint64, dtype="uint64") _mutable_cell_decl("T.u64", syntax="annotation")(u64) # pylint: enable=invalid-name diff --git a/src/relax/script/printer/dependent_type.cc b/src/relax/script/printer/dependent_type.cc index cef8bee1f720..027f64c417dc 100644 --- a/src/relax/script/printer/dependent_type.cc +++ b/src/relax/script/printer/dependent_type.cc @@ -49,7 +49,7 @@ ExprDoc PrintShapeVar(const PrimExpr& e, const AccessPath& e_p, const IRDocsifie if (f != nullptr) { auto walk_fn = [f, &func_var_mode](const tirx::Var& var) -> ffi::Expected { if (auto prim_var = var.as()) { - if (f->func_vars->count(prim_var.value().get())) { + if (f->prim_params->count(prim_var.value().get())) { func_var_mode = true; } } diff --git a/src/relax/script/printer/function.cc b/src/relax/script/printer/function.cc index 3063a7a740e7..47e249ed56e1 100644 --- a/src/relax/script/printer/function.cc +++ b/src/relax/script/printer/function.cc @@ -130,7 +130,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { // Step 6. Print body ffi::Array body = PrintSeqExpr(n->body, n_p->Attr("body"), d, /*use_ret=*/true); (*f)->stmts.insert((*f)->stmts.end(), body.begin(), body.end()); - auto type_var_docs = DefineTypeVarDocs(type_vars, ffi::GetRef((*f).get()), d); + auto type_var_docs = DefineTypeVarDocs(type_vars, d); return WrapFunctionDocWithTypeVars( d, FunctionDoc(func_name, params, {decorator}, ret_type, (*f)->stmts), type_var_docs); }); diff --git a/src/relax/script/printer/tir.cc b/src/relax/script/printer/tir.cc index 8508aef4de5a..8ae3a49a6945 100644 --- a/src/relax/script/printer/tir.cc +++ b/src/relax/script/printer/tir.cc @@ -65,10 +65,11 @@ Doc PrintCanonicalVar(Var n, AccessPath n_p, IRDocsifier d) { f->type_vars->insert(n.get()); } } - IdDoc var = d->Define(n, ffi::GetRef(f), n->name.empty() ? "v" : n->name); + Frame frame = d->frames.front(); + IdDoc var = d->Define(n, frame, n->name.empty() ? "v" : n->name); var->source_paths.push_back(n_p); if (!f->func_vars || f->prim_params->count(n.get()) || !f->type_vars->count(n.get())) { - f->stmts.push_back(AssignDoc(var, PrintVarCreation(prim_var, n_p, d), std::nullopt)); + frame->stmts.push_back(AssignDoc(var, PrintVarCreation(prim_var, n_p, d), std::nullopt)); } } if (ffi::Optional doc = d->GetVarDoc(n)) { diff --git a/src/script/printer/ir/ir.cc b/src/script/printer/ir/ir.cc index dcf059ed8d99..b260d0f3afa9 100644 --- a/src/script/printer/ir/ir.cc +++ b/src/script/printer/ir/ir.cc @@ -64,7 +64,7 @@ struct SortableFunction { } }; -ffi::Optional GetTypeVarDeclarationName(const StmtDoc& stmt) { +ffi::Optional GetDynamicDeclarationName(const StmtDoc& stmt) { const auto* assign = stmt.as(); if (assign == nullptr || !assign->rhs.has_value()) { return std::nullopt; @@ -74,8 +74,8 @@ ffi::Optional GetTypeVarDeclarationName(const StmtDoc& stmt) { if (lhs == nullptr || call == nullptr) { return std::nullopt; } - const auto* callee = call->callee.as(); - if (callee == nullptr || callee->name != "TypeVar") { + const auto* callee = call->callee.as(); + if (callee == nullptr || callee->name != "dynamic") { return std::nullopt; } return lhs->name; @@ -123,7 +123,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { d->cfg->binding_names.pop_back(); if (const auto* stmt_block = doc.as()) { for (const StmtDoc& stmt : stmt_block->stmts) { - if (ffi::Optional name = GetTypeVarDeclarationName(stmt)) { + if (ffi::Optional name = GetDynamicDeclarationName(stmt)) { if (!declared_type_vars.count(name.value())) { declared_type_vars.insert(name.value()); type_var_decls.push_back(stmt); diff --git a/src/script/printer/utils.h b/src/script/printer/utils.h index 1b311b38a6e7..9e9c1bcab3ed 100644 --- a/src/script/printer/utils.h +++ b/src/script/printer/utils.h @@ -134,9 +134,6 @@ inline std::string DType2Str(DLDataType dtype) { inline Doc HeaderWrapper(const IRDocsifier& d, const Doc& doc) { if (d->ir_usage.size()) { ffi::Array stmts; - if (d->ir_usage.count("type_var")) { - stmts.push_back(CommentDoc("from typing import TypeVar")); - } if (d->ir_usage.count("ir")) { stmts.push_back(CommentDoc("from tvm.script import ir as " + d->cfg->ir_prefix)); } @@ -168,7 +165,7 @@ inline Doc HeaderWrapper(const IRDocsifier& d, const Doc& doc) { } inline std::vector> DefineTypeVarDocs( - const std::unordered_set& type_vars, const Frame& frame, const IRDocsifier& d) { + const std::unordered_set& type_vars, const IRDocsifier& d) { std::vector> type_var_docs; type_var_docs.reserve(type_vars.size()); for (const VarNode* var_node : type_vars) { @@ -176,7 +173,7 @@ inline std::vector> DefineTypeVarDocs( ffi::Optional existing_doc = d->GetVarDoc(var); ExprDoc var_doc = existing_doc.has_value() ? existing_doc.value() - : d->Define(var, frame, var->name.empty() ? "v" : var->name); + : d->Define(var, d->frames.front(), var->name.empty() ? "v" : var->name); const auto* id_doc = var_doc.as(); TVM_FFI_ICHECK(id_doc != nullptr); type_var_docs.emplace_back(var, var_doc); @@ -188,6 +185,13 @@ inline std::vector> DefineTypeVarDocs( } inline bool UsePEP695TypeVars(const IRDocsifier& d) { + // Module symbols may be shared by multiple functions. Function-local + // generic parameters would create fresh identities on every definition. + for (const Frame& frame : d->frames) { + if (std::string(frame->GetTypeKey()) == "script.printer.IRFrame") { + return false; + } + } return d->cfg->GetExtraConfig("script.use_pep695", d->cfg->GetExtraConfig("relax.use_pep695", false)); } @@ -200,26 +204,27 @@ inline Doc WrapFunctionDocWithTypeVars(const IRDocsifier& d, FunctionDoc functio if (UsePEP695TypeVars(d)) { d->ir_usage.insert("future_annotations"); for (const auto& type_var_doc : type_var_docs) { - function_doc->type_params.push_back(type_var_doc.second); + PrimType dtype = type_var_doc.first->ty.as_or_throw(); + if (DType2Str(dtype->dtype) == "int64") { + function_doc->type_params.push_back(type_var_doc.second); + } else { + function_doc->type_params.push_back( + AssignDoc(type_var_doc.second, std::nullopt, TIR(d, DType2Str(dtype->dtype)))); + } } return HeaderWrapper(d, function_doc); } - d->ir_usage.insert("type_var"); ffi::Array stmts; - ffi::Array body; for (const auto& [var, var_doc] : type_var_docs) { - const ffi::String& name = var_doc.as()->name; + PrimType dtype = var->ty.as_or_throw(); stmts.push_back(AssignDoc( - var_doc, IdDoc("TypeVar")->Call({LiteralDoc::Str(name, ffi::Optional())}), + var_doc, + IR(d, "dynamic") + ->Call({LiteralDoc::Str(var->name, ffi::Optional())}, {"dtype"}, + {LiteralDoc::Str(DType2Str(dtype->dtype), ffi::Optional())}), std::nullopt)); - // TypeVar supplies the eager annotation's real Python binding. The body - // explicitly declares its native symbol instead of capturing that metadata. - PrimType dtype = var->ty.as_or_throw(); - body.push_back(AssignDoc(var_doc, TIR(d, DType2Str(dtype->dtype))->Call({}), std::nullopt)); } - body.insert(body.end(), function_doc->body.begin(), function_doc->body.end()); - function_doc->body = std::move(body); stmts.push_back(function_doc); return HeaderWrapper(d, StmtBlockDoc(stmts)); } diff --git a/src/tirx/script/printer/buffer.cc b/src/tirx/script/printer/buffer.cc index b99f8db61439..b943631a9fe3 100644 --- a/src/tirx/script/printer/buffer.cc +++ b/src/tirx/script/printer/buffer.cc @@ -32,8 +32,8 @@ namespace printer { ffi::Map BufferAttrs( tirx::BufferVar buffer, const AccessPath& buffer_p, const Frame& frame, const IRDocsifier& d, BufferVarDefinition var_definitions, ffi::Optional data = std::nullopt, - bool stringify_undefined_shape = false, std::unordered_set stringify_shape_vars = {}, - std::unordered_set stringify_compound_shape_vars = {}) { + bool stringify_undefined_shape = false, + std::unordered_set stringify_shape_vars = {}) { using tvm::tirx::Var; using tvm::tirx::VarNode; ffi::Map kwargs; @@ -64,16 +64,14 @@ ffi::Map BufferAttrs( } auto is_new_var = [&](const Expr& e) { return e->IsInstance() && !d->IsVarDefined(e); }; // All expression-string annotation fields use the same Python-binding rule. - // Bare TypeVars are real annotation bindings; compound expressions involving - // them, or any expression referring to a later parameter, must be quoted. + // Dynamic symbols are real Python bindings; expressions referring to later + // scalar parameters must still be quoted. auto expression_doc = [&](const PrimExpr& e, const AccessPath& e_p, bool was_undefined = false) -> ExprDoc { bool needs_quote = stringify_undefined_shape && was_undefined; auto walk_fn = [&](const Var& var) -> ffi::Expected { - needs_quote = needs_quote || - (stringify_undefined_shape && - (!d->IsVarDefined(var) || stringify_shape_vars.count(var))) || - (stringify_compound_shape_vars.count(var) && !e.same_as(var)); + needs_quote = needs_quote || (stringify_undefined_shape && + (!d->IsVarDefined(var) || stringify_shape_vars.count(var))); return ffi::WalkResult::Advance(); }; ffi::StructuralWalk(e, walk_fn); @@ -324,11 +322,10 @@ ExprDoc BufferDecl(const tirx::BufferVar& buffer, const ffi::String& method, } ExprDoc BufferAttn(const tirx::BufferVar& buffer, const AccessPath& p, const Frame& frame, - const IRDocsifier& d, std::unordered_set stringify_shape_vars, - std::unordered_set stringify_compound_shape_vars) { + const IRDocsifier& d, std::unordered_set stringify_shape_vars) { ffi::Map attrs = BufferAttrs(buffer, p, frame, d, BufferVarDefinition::MatchBuffer, std::nullopt, true, - std::move(stringify_shape_vars), std::move(stringify_compound_shape_vars)); + std::move(stringify_shape_vars)); if (!attrs.count("dtype")) { attrs.Set("dtype", LiteralDoc::DataType(buffer->dtype->dtype, p->Attr("dtype"))); } diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc index f694b844edce..ffc6145dd9c0 100644 --- a/src/tirx/script/printer/expr.cc +++ b/src/tirx/script/printer/expr.cc @@ -78,9 +78,9 @@ ExprDoc PrintVarCreation(const tirx::Var& var, const AccessPath& var_p, const IR rhs = TIR(d, "TensorMap")->Call({}, {}, {}); } } else { - rhs = TIR(d, DType2Str(var->ty.as_or_throw()->dtype)); - rhs->source_paths.push_back(var_p->Attr("dtype")); - rhs = rhs->Call({}, kwargs_keys, kwargs_values); + rhs = IR(d, "dynamic") + ->Call({LiteralDoc::Str(var->name, var_p->Attr("name"))}, {"dtype"}, + {LiteralDoc::Str(DType2Str(var->ty.as_or_throw()->dtype), type_p)}); } rhs->source_paths.push_back(type_p); return rhs; @@ -89,9 +89,10 @@ ExprDoc PrintVarCreation(const tirx::Var& var, const AccessPath& var_p, const IR Doc PrintVar(const tirx::Var& var, const AccessPath& var_p, const IRDocsifier& d) { if (!d->IsVarDefined(var)) { if (ffi::Optional opt_f = FindLowestVarDef(var, d)) { - ExprDoc lhs = DefineVar(var, opt_f.value(), d); + Frame frame = var->ty.as() ? d->frames.front() : opt_f.value(); + ExprDoc lhs = DefineVar(var, frame, d); ExprDoc rhs = PrintVarCreation(var, var_p, d); - opt_f.value()->stmts.push_back(AssignDoc(lhs, rhs, std::nullopt)); + frame->stmts.push_back(AssignDoc(lhs, rhs, std::nullopt)); } else { LOG(WARNING) << "Didn't find variable definition for: " << var->name; } diff --git a/src/tirx/script/printer/function.cc b/src/tirx/script/printer/function.cc index 34f655d2860c..59ed0ce2625c 100644 --- a/src/tirx/script/printer/function.cc +++ b/src/tirx/script/printer/function.cc @@ -49,8 +49,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { continue; } PrimType var_ty(var_ty_node->dtype); - if (!runtime_params.count(var.get()) && var_ty.IsScalar() && - var_ty.MatchesElementType(DLDataTypeCode::kDLInt, 64)) { + if (!runtime_params.count(var.get()) && var_ty.IsScalar()) { type_vars.insert(var.get()); } } @@ -71,8 +70,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { collect_type_vars(address); } } - auto type_var_docs = DefineTypeVarDocs(type_vars, ffi::GetRef((*f).get()), d); - bool use_postponed_annotations = UsePEP695TypeVars(d) && !type_vars.empty(); + auto type_var_docs = DefineTypeVarDocs(type_vars, d); int n_args = func->params.size(); // Step 1. Handle `func->params` ffi::Array args; @@ -120,15 +118,11 @@ TVM_FFI_STATIC_INIT_BLOCK() { continue; } std::unordered_set stringify_shape_vars; - std::unordered_set stringify_compound_shape_vars; auto walk_fn = [&](const tirx::Var& shape_var) -> ffi::Expected { bool is_type_var = type_vars.count(shape_var.get()); if (!bound_signature_vars.count(shape_var) && !is_type_var) { stringify_shape_vars.insert(shape_var); } - if (!use_postponed_annotations && is_type_var) { - stringify_compound_shape_vars.insert(shape_var); - } return ffi::WalkResult::Advance(); }; for (const PrimExpr& shape : buffer->shape) { @@ -143,8 +137,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { } IdDoc lhs = DefineBuffer(buffer, *f, d); ExprDoc annotation = - BufferAttn(buffer, var_p->Attr("ty"), *f, d, std::move(stringify_shape_vars), - std::move(stringify_compound_shape_vars)); + BufferAttn(buffer, var_p->Attr("ty"), *f, d, std::move(stringify_shape_vars)); args.push_back(AssignDoc(lhs, std::nullopt, annotation)); continue; } diff --git a/src/tirx/script/printer/utils.h b/src/tirx/script/printer/utils.h index fca09c45b511..919ddc6ba4c7 100644 --- a/src/tirx/script/printer/utils.h +++ b/src/tirx/script/printer/utils.h @@ -324,13 +324,10 @@ ExprDoc BufferDecl(const tirx::BufferVar& buffer, const ffi::String& method, * \param d The IRDocsifier * \param stringify_shape_vars Variables without a Python binding at this annotation. Every * shape expression containing one of these variables must be stringified. - * \param stringify_compound_shape_vars Variables whose compound shape expressions must be - * stringified while their bare-name uses remain direct. * \return The ExprDoc corresponding to the buffer declaration */ ExprDoc BufferAttn(const tirx::BufferVar& buffer, const AccessPath& p, const Frame& frame, - const IRDocsifier& d, std::unordered_set stringify_shape_vars = {}, - std::unordered_set stringify_compound_shape_vars = {}); + const IRDocsifier& d, std::unordered_set stringify_shape_vars = {}); /*! * \brief Print the creation of a Var diff --git a/tests/python/relax/test_tvmscript_parser.py b/tests/python/relax/test_tvmscript_parser.py index 5a0d39cfd68b..6e72a575b2c8 100644 --- a/tests/python/relax/test_tvmscript_parser.py +++ b/tests/python/relax/test_tvmscript_parser.py @@ -113,16 +113,17 @@ def f(x: R.Tensor(dtype="float32", ndim="1")): # error: dim is expected to be i def test_unexpected_tir_cast_args(): with pytest.raises(TypeError): + m = T.dynamic("m", "int64") @R.function - def f(x: R.Tensor(("m",), "float32")): - m = T.int64() + def f(x: R.Tensor((m,), "float32")): # tirx.cast expects 2 arguments, but got 3 return R.call_tir("foo", (x,), R.Tensor((T.cast("int32", m, 1),), dtype="float32")) def test_unexpected_tir_args(): with pytest.raises(TypeError): + m = T.dynamic("m", "int64") @tvm.script.ir_module class TestWellCallTIR: @@ -135,17 +136,16 @@ def tir_addone(A: T.Buffer((16, 16), "int32"), B: T.Buffer((16, 16), "int32")) - B[vi, vj] = A[vi, vj] + T.int32(1) @R.function - def foo(x: R.Tensor(("m", "m"), "float32")): - m = T.int64() + def foo(x: R.Tensor((m, m), "float32")): # tirx.max expects 2 arguments, but got 1 gv = R.call_tir(tir_addone, (x,), R.Tensor((T.max(16),), dtype="float32")) return gv with pytest.raises(TypeError): + m = T.dynamic("m", "int64") @R.function - def f(x: R.Tensor(("m", "n"), "float32")): - m = T.int64() + def f(x: R.Tensor((m, "n"), "float32")): # call_tir expected a tirx prim_func return relax.call_tir("extern_func", (x,), R.Tensor((T.max(m),), dtype="float32")) @@ -430,26 +430,28 @@ def foo(x: R.Shape((4, 4))): def test_symbolic_shape(): + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") + @R.function - def foo(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() + def foo(x: R.Tensor((m, n), "float32")) -> R.Tensor((m, n), "float32"): gv0 = R.call_dps_packed("extern_func", x, R.Tensor((m, n), dtype="float32")) return gv0 + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") + @R.function - def bar(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() + def bar(x: R.Tensor((m, n), "float32")) -> R.Tensor((m, n), "float32"): gv0 = R.call_dps_packed("extern_func", x, R.Tensor((m, n), dtype="float32")) return gv0 with pytest.raises(tvm.error.InternalError): + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int32") @R.function - def mismatch_dtype(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(None, "float32", ndim=2): - m = T.int64() - n = T.int32() # The shape dtype should be int64 + def mismatch_dtype(x: R.Tensor((m, n), "float32")) -> R.Tensor(None, "float32", ndim=2): gv0 = R.call_dps_packed("extern_func", x, R.Tensor((m, n), dtype="float32")) return gv0 @@ -494,10 +496,11 @@ def foo(x: R.Tensor((4, 4), "float32")): def test_match_cast(): + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") + @R.function def foo(x: R.Tensor("float32"), y: R.Tensor("float32")): - m = T.int64() - n = T.int64() x0 = R.match_cast(x, R.Tensor([m], "float32")) with R.dataflow(): y0 = R.match_cast(y, R.Tensor([n], "float32")) @@ -539,9 +542,11 @@ def foo(x: R.Tensor((4, 4), "float32")): def test_tuple_return_2(): + n = T.dynamic("n", "int64") + m = T.dynamic("m", "int64") + @R.function def foo(x: R.Tensor("float32", ndim=2)): - n, m = T.int64(), T.int64() x0 = R.match_cast(x, R.Tensor((n, m), "float32")) return (x0, R.shape([n + 1, m, 1])) @@ -556,9 +561,11 @@ def foo(x: R.Tensor("float32", ndim=2)): def test_tuple_binding(): + n = T.dynamic("n", "int64") + m = T.dynamic("m", "int64") + @R.function def foo(x: R.Tensor("float32", ndim=2)): - n, m = T.int64(), T.int64() x0 = R.match_cast(x, R.Tensor((n, m), "float32")) t0 = (x, x0) t1 = (x, R.shape([n, m]), t0) @@ -627,13 +634,14 @@ def foo(x: R.Tensor((128, 128), "float32")) -> R.Tensor(None, "float32", ndim=2) def test_dataflow_block_advanced(): + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") + @R.function def foo(x: R.Tensor((128, 128), "float32")) -> R.Tensor(None, "float32", ndim=2): gv0 = R.call_dps_packed("extern_func", x, R.Tensor((128, 128), dtype="float32")) gv1 = R.call_dps_packed("extern_func", gv0, R.Tensor((128, 128), dtype="float32")) with R.dataflow(): - m = T.int64() - n = T.int64() lv0 = R.call_dps_packed("extern_func", gv1, R.Tensor((128, 128), dtype="float32")) lv1 = R.match_cast(lv0, R.Tensor((m, n), "float32")) gv2 = R.call_dps_packed("extern_func", lv0, R.Tensor((128, 128), dtype="float32")) @@ -904,13 +912,14 @@ def foo(x: R.Object) -> R.Object: def test_annotation(): + m = T.dynamic("m", "int64") + @R.function(pure=False) def foo( - x: R.Tensor((32, "m"), "float32"), - y: R.Tensor(("m",), "float32"), + x: R.Tensor((32, m), "float32"), + y: R.Tensor((m,), "float32"), r: R.Tensor(dtype="int64"), ) -> R.Any: - 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) @@ -996,13 +1005,14 @@ def test_call_tir_empty_tuple_arg(): def test_call_tir_with_tir_var(): + n = T.dynamic("n", "int64") + @I.ir_module class Module: @R.function def main( - dumb_param: R.Tensor(("n",), "float32"), x: R.Tensor(("n * 2",), "float32") - ) -> R.Tensor(("n * 2",), "float32"): - n = T.int64() + dumb_param: R.Tensor((n,), "float32"), x: R.Tensor((n * 2,), "float32") + ) -> R.Tensor((n * 2,), "float32"): cls = Module y = R.call_tir(cls.copy, (x, n), R.Tensor((n * 2,), dtype="float32")) return y @@ -1356,9 +1366,10 @@ def func(cond: T.bool, x: R.Tensor((1,), "float32")): def test_computed_prim_value_as_branch_condition(): """The primitive scalar condition may be computed within the function""" + N = T.dynamic("N", "int64") + @R.function - def func(x: R.Tensor(["N"], "float32")): - N = T.int64() + def func(x: R.Tensor([N], "float32")): if R.prim_value(N % 16 == 0): out = R.call_pure_packed("fast_vectorized_impl", x, ty_args=[x.ty]) else: @@ -1375,18 +1386,20 @@ def func(x: R.Tensor(["N"], "float32")): def test_tir_expr_as_branch_condition(): """Syntactic sugar, use Expr directly""" + N = T.dynamic("N", "int64") + @R.function(private=True) - def sugared(x: R.Tensor(["N"], "float32")): - N = T.int64() + def sugared(x: R.Tensor([N], "float32")): if N % 16 == 0: out = R.call_pure_packed("fast_vectorized_impl", x, ty_args=[x.ty]) else: out = R.call_pure_packed("slow_non_vectorized_impl", x, ty_args=[x.ty]) return out + N = T.dynamic("N", "int64") + @R.function(private=True) - def unsugared(x: R.Tensor(["N"], "float32")): - N = T.int64() + def unsugared(x: R.Tensor([N], "float32")): if R.prim_value(N % 16 == 0): out = R.call_pure_packed("fast_vectorized_impl", x, ty_args=[x.ty]) else: @@ -1429,9 +1442,10 @@ def func(cond: T.bool, x: R.Tensor((1,), "float32")): def test_computed_prim_value_as_assert_condition(): """The primitive scalar condition may be computed within the function""" + N = T.dynamic("N", "int64") + @R.function(pure=False) - def func(x: R.Tensor(["N"], "float32")): - N = T.int64() + def func(x: R.Tensor([N], "float32")): _ = R.assert_op(R.prim_value(N % 16 == 0)) out = R.call_packed("fast_vectorized_impl", x, ty_args=[x.ty]) return out @@ -1447,16 +1461,18 @@ def func(x: R.Tensor(["N"], "float32")): def test_tir_expr_as_assert_condition(): """Syntactic sugar, use Expr directly""" + N = T.dynamic("N", "int64") + @R.function(pure=False, private=True) - def sugared(x: R.Tensor(["N"], "float32")): - N = T.int64() + def sugared(x: R.Tensor([N], "float32")): _ = R.assert_op(N % 16 == 0) out = R.call_packed("fast_vectorized_impl", x, ty_args=[x.ty]) return out + N = T.dynamic("N", "int64") + @R.function(pure=False, private=True) - def unsugared(x: R.Tensor(["N"], "float32")): - N = T.int64() + def unsugared(x: R.Tensor([N], "float32")): _ = R.assert_op(R.prim_value(N % 16 == 0)) out = R.call_packed("fast_vectorized_impl", x, ty_args=[x.ty]) return out @@ -1465,10 +1481,12 @@ def unsugared(x: R.Tensor(["N"], "float32")): def test_erase_to_well_defined_removes_internal_vars(): + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") + @R.function def foo(x: R.Tensor): q = x - m, n = T.int64(), T.int64() z = R.match_cast(q, R.Tensor((m, n))) w = z return w @@ -1479,10 +1497,12 @@ def foo(x: R.Tensor): def test_erase_to_well_defined_keeps_variables_exposed_by_tensor_shape(): + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") + @R.function - def foo(x: R.Tensor(["m", "n"])): + def foo(x: R.Tensor([m, n])): q = x - m, n = T.int64(), T.int64() z = R.match_cast(q, R.Tensor((m, n))) w = z return w @@ -1492,10 +1512,12 @@ def foo(x: R.Tensor(["m", "n"])): def test_erase_to_well_defined_keeps_variants_exposed_by_shape_expr(): + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") + @R.function - def foo(x: R.Tensor, _: R.Shape(["m", "n"])): + def foo(x: R.Tensor, _: R.Shape([m, n])): q = x - m, n = T.int64(), T.int64() z = R.match_cast(q, R.Tensor((m, n))) w = z return w @@ -1505,13 +1527,15 @@ def foo(x: R.Tensor, _: R.Shape(["m", "n"])): def test_erase_to_well_defined_infers_from_shape_expr(): + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") + @I.ir_module class Module: # The subroutine's symbolic variables are only in-scope for the subroutine. @R.function - def subroutine(x: R.Tensor, _: R.Shape(["m", "n"])) -> R.Tensor(["m", "n"]): + def subroutine(x: R.Tensor, _: R.Shape([m, n])) -> R.Tensor([m, n]): q = x - m, n = T.int64(), T.int64() z = R.match_cast(q, R.Tensor((m, n))) w = z return w @@ -1566,11 +1590,12 @@ def foo(x: R.Tensor(("m + 1",), "float32"), y: R.Tensor(("m", 1), "float32")): def test_symbolic_vars_in_tensor_shape_with_definition_first(): """Second param may use symbolic variable defined in first param""" + m = T.dynamic("m", "int64") + @R.function - def bar(x: R.Tensor(("m",), "float32"), y: R.Tensor(("T.max(m, 20)",), "float32")) -> R.Tensor( + 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.int64() z = R.call_dps_packed("test_intrin", (x, y), R.Tensor((T.max(m, 20) + 1,), dtype="float32")) return z @@ -1686,9 +1711,10 @@ def recurse(current: T.int64, value: R.Tensor([current], "float32")) -> R.Tensor def test_symbolic_vars_in_shape(): """Symbolic variable may be defined in R.Shape""" + m = T.dynamic("m", "int64") + @R.function - def baz(x: R.Shape(("m",)), y: R.Tensor(("m * 2",), "float32")): - m = T.int64() + def baz(x: R.Shape((m,)), y: R.Tensor((m * 2,), "float32")): z = R.call_dps_packed("test_intrin", y, R.Tensor((m * 2,), dtype="float32")) return z @@ -1772,10 +1798,11 @@ def foo(x: R.Tensor(("m", "n"), "float32"), y: R.Tensor(("m", "n"), "float32")): def test_memory_ops(): + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") + @R.function - def foo(x: R.Tensor(("m", "n"), dtype="float32")): - m = T.int64() - n = T.int64() + def foo(x: R.Tensor((m, n), dtype="float32")): storage = R.memory.alloc_storage( R.shape([4 * m * n]), virtual_device_index=0, storage_scope="global", dtype="float32" ) @@ -1788,10 +1815,11 @@ def foo(x: R.Tensor(("m", "n"), dtype="float32")): def test_vm_ops(): + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") + @R.function(pure=False) - def foo(x: R.Tensor(("m", "n"), dtype="float32")): - m = T.int64() - n = T.int64() + def foo(x: R.Tensor((m, n), dtype="float32")): 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) @@ -2371,6 +2399,9 @@ def main(A: R.Tensor, B: R.Tensor): def test_function_attributes_are_defined(): """func.attrs defaults to an empty DictAttrs""" + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") + @I.ir_module class Module: @R.function @@ -2379,9 +2410,8 @@ def main(x: R.Tensor, shape: R.Shape(["m", "n"])): return output @R.function - def subroutine(x: R.Tensor, _: R.Shape(["m", "n"])) -> R.Tensor(["m", "n"]): + def subroutine(x: R.Tensor, _: R.Shape([m, n])) -> R.Tensor([m, n]): q = x - m, n = T.int64(), T.int64() z = R.match_cast(q, R.Tensor((m, n))) w = z return w @@ -2399,15 +2429,17 @@ def test_function_symbolic_variables_are_annotated(): for simplifications must be provided to the analyzer. """ + extent = T.dynamic("extent", "int64") + @R.function(private=True) - def inferred_ty(A: R.Tensor(["extent"])): - extent = T.int64() + def inferred_ty(A: R.Tensor([extent])): output = R.strided_slice(A, [0], [0], [extent - 1]) return output + extent = T.dynamic("extent", "int64") + @R.function(private=True) - def expected(A: R.Tensor(["extent"])) -> R.Tensor(["extent-1"]): - extent = T.int64() + def expected(A: R.Tensor([extent])) -> R.Tensor([extent - 1]): output: R.Tensor([extent - 1]) = R.strided_slice(A, [0], [0], [extent - 1]) return output @@ -2415,7 +2447,7 @@ def expected(A: R.Tensor(["extent"])) -> R.Tensor(["extent-1"]): def test_non_declaration_prim_expr_emits_binding(): - """Only zero-argument dtype calls declare symbolic variables.""" + """Dtype casts emit ordinary bindings without replacing shape symbols.""" @R.function(private=True) def func(A: R.Tensor(["extent"], "float32")): @@ -2484,9 +2516,10 @@ def test_shared_meta_var_uses_ordinary_relax_bindings(): assert I.meta_var is T.meta_var + N = T.dynamic("N", "int64") + @R.function(private=True) - def func(A: R.Tensor(["N"], "float32")): - N: T.int64 = T.int64() + def func(A: R.Tensor([N], "float32")): via_i = I.meta_var(N) via_t = T.meta_var(via_i) output = R.reshape(A, R.shape([via_t])) @@ -2504,11 +2537,12 @@ def func(A: R.Tensor(["N"], "float32")): assert "meta_var" not in source _check(func) - with pytest.raises(ValueError): + symbol = T.dynamic("symbol", "int64") + with pytest.raises(tvm.error.InternalError, match="Invalid annotation"): @R.function(private=True) - def mismatched_declaration(): - value: T.float32 = T.int64() + def mismatched_binding(): + value: T.float32 = symbol return value @@ -2555,14 +2589,14 @@ def test_conditional_may_use_symbolic_variables_from_function_scope(): """ + N = T.dynamic("N", "int64") + @R.function(private=True) def explicit_ty( - A: R.Tensor(["N"], "float32"), - B: R.Tensor(["N"], "float32"), + A: R.Tensor([N], "float32"), + B: R.Tensor([N], "float32"), cond: T.bool, - ) -> R.Tensor(["N"], "float32"): - N = T.int64() - + ) -> R.Tensor([N], "float32"): if cond: out: R.Tensor([N], "float32") = A + B else: @@ -2570,13 +2604,14 @@ def explicit_ty( return out + N = T.dynamic("N", "int64") + @R.function(private=True) def inferred_ty( - A: R.Tensor(["N"], "float32"), - B: R.Tensor(["N"], "float32"), + A: R.Tensor([N], "float32"), + B: R.Tensor([N], "float32"), cond: T.bool, ): - N = T.int64() if cond: out = A + B else: diff --git a/tests/python/relax/test_tvmscript_printer_relax.py b/tests/python/relax/test_tvmscript_printer_relax.py index 50005330068d..7866ddb6e8c2 100644 --- a/tests/python/relax/test_tvmscript_printer_relax.py +++ b/tests/python/relax/test_tvmscript_printer_relax.py @@ -61,7 +61,7 @@ def test_function_dependent_shape_escaped_source_spans(): cast = tirx.Cast("int64", n) x = relax.Var("x", relax.TensorType([cast], "float32")) ret_ty = relax.TensorType(dtype="float32", ndim=1) - func = relax.Function([x], x, ret_ty=ret_ty).with_attr("global_symbol", "main") + func = relax.Function([x, n], x, ret_ty=ret_ty).with_attr("global_symbol", "main") cast_path = ( AccessPath.root() .attr("params") @@ -247,7 +247,7 @@ def test_shape_ty_2(): _assert_print( obj, """ -a = T.int64() +a = I.dynamic("a", dtype="int64") R.Shape([1, a, 3])""", ) @@ -260,7 +260,7 @@ def test_tensor_ty(): _assert_print( obj, """ -a = T.int64() +a = I.dynamic("a", dtype="int64") R.Tensor((1, a, 3), dtype="float32") """, ) @@ -302,7 +302,7 @@ def test_func_ty(): ) _assert_print( obj, - "a = T.int64()\n" + 'a = I.dynamic("a", dtype="int64")\n' "R.Callable((T.float32, R.Any, R.Shape([1, a, 3]), T.int64), " 'R.Tensor((1, 2, 3), dtype="float32"), True)', ) @@ -413,7 +413,7 @@ def test_var(): _assert_print( obj, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") a: R.Tensor((1, x, 3), dtype="float32") a""", ) @@ -424,7 +424,7 @@ def test_dataflow_var(): _assert_print( obj, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") a: R.Tensor((1, x, 3), dtype="float32") a""", ) @@ -441,11 +441,11 @@ def test_tuple(): _assert_print( obj, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") a: R.Tensor((1, x, 3), dtype="float32") -y = T.int64() +y = I.dynamic("y", dtype="int64") b: R.Tensor((1, y, 3), dtype="float32") -z = T.int64() +z = I.dynamic("z", dtype="int64") c: R.Tensor((1, z, 3), dtype="float32") (a, b, c) """, @@ -466,11 +466,11 @@ def test_tuple_get_item(): _assert_print( obj, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") a: R.Tensor((1, x, 3), dtype="float32") -y = T.int64() +y = I.dynamic("y", dtype="int64") b: R.Tensor((1, y, 3), dtype="float32") -z = T.int64() +z = I.dynamic("z", dtype="int64") c: R.Tensor((1, z, 3), dtype="float32") (a, b, c)[0] """, @@ -490,7 +490,7 @@ def test_call(): _assert_print( o0, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") a: R.Tensor((1, x, 3), dtype="float32") R.call_tir(tir_func, (a, x), out_ty=R.Tensor((1, x, 3), dtype="float32")) """, @@ -498,7 +498,7 @@ def test_call(): _assert_print( o1, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") a: R.Tensor((1, x, 3), dtype="float32") R.call_dps_packed("my_dps_func", (a,), out_ty=R.Tensor((1, x, 3), dtype="float32")) """, @@ -520,7 +520,7 @@ def test_call_tir_with_grad(): v1, """ v0: R.Tensor((54, 96), dtype="float32") -x = T.int64() +x = I.dynamic("x", dtype="int64") R.call_tir_with_grad(tir_func, (v0,), out_ty=R.Tensor((54, 96), dtype="float32"), te_grad_name="grad_func", te_grad_kwargs={"k": 1.0, "x": x}) """, ) @@ -545,7 +545,7 @@ def test_call_tir_inplace(): """ x: R.Tensor((32, 32), dtype="int32") y: R.Tensor((32, 32), dtype="int32") -t = T.int64() +t = I.dynamic("t", dtype="int64") R.call_tir_inplace(tir_func, (x, y, t), out_ty=[R.Tensor((32, 32), dtype="int32"), R.Tensor((32, 32), dtype="int32")], inplace_indices=[-1, 0]) """, ) @@ -571,7 +571,7 @@ def test_seq_expr(): _assert_print( obj, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") a: R.Tensor((1, x, 3), dtype="float32") with R.dataflow(): b: R.Tensor((1, x, 3), dtype="float32") = R.sin(a) @@ -596,7 +596,7 @@ def test_binding_block(): _assert_print( obj, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") a: R.Tensor((1, x, 3), dtype="float32") b: R.Tensor((1, x, 3), dtype="float32") = R.sin(a) c: R.Tensor((1, x, 3), dtype="float32") = R.sin(b) @@ -618,7 +618,7 @@ def test_dataflow_block(): _assert_print( obj, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") a: R.Tensor((1, x, 3), dtype="float32") with R.dataflow(): b: R.Tensor((1, x, 3), dtype="float32") = R.sin(a) @@ -640,7 +640,7 @@ def test_match_cast(): _assert_print( obj, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") a: R.Tensor((1, x, 3), dtype="float32") b: R.Tensor((1, 5, 3), dtype="float32") = R.match_cast(a, R.Tensor((1, 5, 3), dtype="float32")) """, @@ -655,7 +655,7 @@ def test_var_binding(): _assert_print( obj, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") a: R.Tensor((1, x, 3), dtype="float32") b: R.Tensor((1, x, 3), dtype="float32") = R.sin(a) """, @@ -693,7 +693,7 @@ def test_builtin_keywords(): _assert_print( obj, """ -x = T.int64() +x = I.dynamic("x", dtype="int64") R_1: R.Tensor((1, x, 3), dtype="float32") T_1: R.Tensor((1, x, 3), dtype="float32") = R.sin(R_1) """, diff --git a/tests/python/relax/test_tvmscript_pyfunc.py b/tests/python/relax/test_tvmscript_pyfunc.py index a9ffd1db4fb9..0a90f3b0eb96 100644 --- a/tests/python/relax/test_tvmscript_pyfunc.py +++ b/tests/python/relax/test_tvmscript_pyfunc.py @@ -37,6 +37,8 @@ from tvm.script import s_tir as Ts from tvm.script import tirx as T +n = T.dynamic("n", "int32") + @R.py_module class TestPyFuncModule(BasePyModule): @@ -65,7 +67,6 @@ def simple_tir_func( var_B: T.handle, ): T.func_attr({"tirx.noalias": True}) - n = T.int32() A = T.match_buffer(var_A, (n,), "float32") B = T.match_buffer(var_B, (n,), "float32") diff --git a/tests/python/relax/test_tvmscript_type_vars.py b/tests/python/relax/test_tvmscript_type_vars.py index 9dedab1bc023..e6dbafa3c0a5 100644 --- a/tests/python/relax/test_tvmscript_type_vars.py +++ b/tests/python/relax/test_tvmscript_type_vars.py @@ -16,14 +16,14 @@ # under the License. import sys -from typing import TypeVar import tvm import tvm.testing +from tvm.script import ir as I from tvm.script import relax as R -M = TypeVar("M") -UNUSED_GENERIC = TypeVar("UNUSED_GENERIC", bound=int) +M = I.dynamic("M") +UNUSED_GENERIC = I.dynamic("UNUSED_GENERIC") def test_type_vars_roundtrip(): @@ -49,15 +49,15 @@ def func[M: int](x: R.Tensor((M, M * 2), "float32")): tvm.ir.assert_structural_equal(func, typed) else: assert "from __future__ import annotations" not in script - assert 'M = TypeVar("M")' in script - assert "M = T.int64()" in script - assert 'R.Tensor((M, "M * 2"), dtype="float32")' in script + assert 'M = I.dynamic("M", dtype="int64")' in script + assert "M = T.int64()" not in script + assert 'R.Tensor((M, M * 2), dtype="float32")' in script portable = func.script(extra_config={"relax.use_pep695": False}) assert "from __future__ import annotations" not in portable - assert 'M = TypeVar("M")' in portable - assert 'R.Tensor((M, "M * 2"), dtype="float32")' in portable - assert "M = T.int64()" in portable + assert 'M = I.dynamic("M", dtype="int64")' in portable + assert 'R.Tensor((M, M * 2), dtype="float32")' in portable + assert "M = T.int64()" not in portable assert "UNUSED_GENERIC" not in script assert [param.name for param in func.params] == ["x"] assert not hasattr(func, "type_params") @@ -66,5 +66,29 @@ def func[M: int](x: R.Tensor((M, M * 2), "float32")): tvm.ir.assert_structural_equal(func, tvm.script.from_source(portable)) +def test_dynamic_module_symbol_identity(): + shared = I.dynamic("n") + independent = I.dynamic("n") + + @R.function(private=True) + def first(x: R.Tensor((shared,), "float32")): + return x + + @R.function(private=True) + def second(x: R.Tensor((shared, independent), "float32")): + return x + + mod = tvm.IRModule({"first": first, "second": second}) + source = mod.script() + assert source.count('I.dynamic("n", dtype="int64")') == 2 + restored = tvm.script.from_source(source, check_well_formed=False) + first_n = restored["first"].params[0].ty.shape.values[0] + second_shape = restored["second"].params[0].ty.shape.values + assert first_n.same_as(second_shape[0]) + assert not first_n.same_as(second_shape[1]) + assert str(second_shape[1].ty.dtype) == "int64" + tvm.ir.assert_structural_equal(mod, restored) + + if __name__ == "__main__": tvm.testing.main() diff --git a/tests/python/script/minilang.py b/tests/python/script/minilang.py index de9c99b30e77..2de387129b83 100644 --- a/tests/python/script/minilang.py +++ b/tests/python/script/minilang.py @@ -111,7 +111,9 @@ def __exit__(self, error_type, error, traceback): def resolve_type_var(self, name, dtype=None, *, value=None, **kwargs): if name not in self.type_var_map: self.type_var_map[name] = ( - value if value is not None else Value("symbol", (dtype,), name) + value + if value is not None + else Value("symbol", ("int64" if dtype is None else dtype,), name) ) return self.type_var_map[name] @@ -220,14 +222,15 @@ def operation(*args, name=name): def Tensor(shape=None, dtype="float32", device=None, placement="S[0]"): return Value("tensor", (shape, dtype, device, placement)) - def symbol(expr=None): - return Value("symbol", (expr,)) + def dynamic(name, dtype="int64"): + return Value("symbol", (dtype,), name) def cell(value=None): return Value("cell", (value,)) self.M.Tensor = Tensor - self.M.symbol = registry.register_type_var_decl("M.symbol", symbol, dtype="int64") + self.M.dynamic = dynamic + self.M.int32 = registry.register_scalar_annotation("M.int32", lambda: None, dtype="int32") self.M.cell = registry.mutable_cell_decl("M.cell")(cell) @contextmanager diff --git a/tests/python/script/test_meta_programming.py b/tests/python/script/test_meta_programming.py index 20d6e2e0eeb7..c5a30e282cbe 100644 --- a/tests/python/script/test_meta_programming.py +++ b/tests/python/script/test_meta_programming.py @@ -335,12 +335,10 @@ def _build_symbolic_functions(M): class Module: @M.function def first(x: M.Tensor(("n",), "float32")): - n = M.symbol() # noqa: F841 return x @M.function def second(x: M.Tensor(("n",), "float32")): - n = M.symbol() # noqa: F841 return x return Module diff --git a/tests/python/script/test_symbolic_shape.py b/tests/python/script/test_symbolic_shape.py index fe281434b569..e90b8598aa52 100644 --- a/tests/python/script/test_symbolic_shape.py +++ b/tests/python/script/test_symbolic_shape.py @@ -14,7 +14,7 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -"""Symbolic dimensions in function signatures share identity with body declarations. +"""Explicit symbolic dimensions retain identity in signatures and function bodies. Quoted dimensions must not create or replace ordinary Python bindings. """ @@ -23,25 +23,26 @@ # Invalid script examples deliberately contain unresolved or unused bindings. # ruff: noqa: F821, F841 import inspect +import sys import pytest -from tvm import ir +from tvm.script import ir as I from tvm.script import tirx as T +from tvm.script.parser import entry def test_signature_symbols_cross_nested_calls_parameters_return_and_body(language): - # Repeated quoted dimensions in nested annotations and body declarations must resolve to - # the same symbol. + # Nested annotations and body reads retain the externally constructed symbol. M = language.M M.tuple = lambda *fields: fields + n = M.dynamic("n") @M.function def main( - x: M.tuple(M.Tensor(("n",), "float32"), M.Tensor(("n",), "float32")), - y: M.Tensor(("n",), "float32"), - ) -> M.Tensor(("n",), "float32"): - n = M.symbol() + x: M.tuple(M.Tensor((n,), "float32"), M.Tensor((n,), "float32")), + y: M.Tensor((n,), "float32"), + ) -> M.Tensor((n,), "float32"): M.record(n) return y @@ -89,7 +90,7 @@ def main(x: T.Buffer(shape, "float32")): return main - n = ir.Var("n", "int64") + n = T.dynamic("n") function = build((n, 16)) assert function.params[0].ty.shape[0].same_as(n) with pytest.raises( @@ -137,16 +138,17 @@ def test_symbol_reassignment_reports_introduction_and_exact_write(language): # Ordinary writes to a symbolic dimension must point to its original introduction and # exact target. M = language.M + n = M.dynamic("n") with pytest.raises(SyntaxError) as caught: @M.function - def main(): - n = M.symbol() + def main(x: M.Tensor((n,))): n = 2 message = str(caught.value) introduction = _line_of( - test_symbol_reassignment_reports_introduction_and_exact_write, "n = M.symbol()" + test_symbol_reassignment_reports_introduction_and_exact_write, + "def main(x: M.Tensor((n,))):", ) offending = _line_of(test_symbol_reassignment_reports_introduction_and_exact_write, "n = 2") assert "Symbolic variable 'n' cannot be reassigned" in message @@ -156,15 +158,15 @@ def main(): assert (error.offset, error.end_offset) == (13, 14) -def test_repeated_symbol_declarations_reuse_identity(language): - # Repeated explicit declarations must reuse the annotation-owned symbol. +def test_external_dynamic_symbols_reuse_identity(language): + # Repeated captures share the exact externally constructed symbol. M = language.M + n = M.dynamic("n") + @M.function - def main(x: M.Tensor(("n",))): - n = M.symbol() + def main(x: M.Tensor((n,))): M.record(n) - n = M.symbol() M.record(n) symbol = main.params[0].args[0].args[0][0] @@ -176,10 +178,10 @@ def test_nested_scope_does_not_reassign_outer_symbol(language): # An ordinary nested local must not overwrite an enclosing symbolic dimension. M = language.M + n = M.dynamic("n") + @M.function def main(): - n = M.symbol() - @M.function def nested(): n = 2 @@ -189,3 +191,41 @@ def nested(): assert main.body[0][1].op == "symbol" assert language.functions["nested"].body == [("emit", 2)] + + +@pytest.mark.skipif(sys.version_info < (3, 12), reason="PEP 695 requires Python 3.12") +def test_mixed_generic_symbol_dtypes_and_identity(language): + source = """ +@M.function +def main[n, k: M.int32](x: M.Tensor((n, k))) -> M.Tensor((n, k)): + M.record(n) + M.record(k) + return x +""" + main = entry.parse(source, extra_vars={"M": language.M}, root_builder=language.M) + n, k = main.params[0].args[0].args[0] + assert n.args == ("int64",) + assert k.args == ("int32",) + assert main.ret_type.args[0][0] is n + assert main.ret_type.args[0][1] is k + assert main.body[0][1] is n + assert main.body[1][1] is k + + +def test_dynamic_symbols_are_fresh_and_scope_independent(): + assert T.dynamic is I.dynamic + n = T.dynamic("n") + same_name = I.dynamic("n") + k = I.dynamic("k", "int32") + assert n.ty.dtype == "int64" + assert k.ty.dtype == "int32" + assert not n.same_as(same_name) + + @I.ir_module + class Module: + @T.prim_func + def first(x: T.Buffer((n,), "float32")): + T.evaluate(n) + + assert Module["first"].params[0].ty.shape[0].same_as(n) + assert Module["first"].body.value.same_as(n) diff --git a/tests/python/tirx/test_tvmscript_type_vars.py b/tests/python/tirx/test_tvmscript_type_vars.py index eb35952d2b5c..76d5f8ff118d 100644 --- a/tests/python/tirx/test_tvmscript_type_vars.py +++ b/tests/python/tirx/test_tvmscript_type_vars.py @@ -24,8 +24,8 @@ def test_type_vars_roundtrip(): func = tvm.script.from_source( """ -M = TypeVar("M") -UNUSED = TypeVar("UNUSED") +M = I.dynamic("M") +UNUSED = I.dynamic("UNUSED") @T.prim_func(private=True) def func(A: T.Buffer((M, M * 2), "float32")): @@ -49,15 +49,15 @@ def func[M: int](A: T.Buffer((M, M * 2), "float32")): tvm.ir.assert_structural_equal(func, typed) else: assert "from __future__ import annotations" not in script - assert 'M = TypeVar("M")' in script - assert "M = T.int64()" in script + assert 'M = I.dynamic("M", dtype="int64")' in script + assert "M = T.int64()" not in script portable = func.script(extra_config={"script.use_pep695": False}) assert "from __future__ import annotations" not in portable - assert 'M = TypeVar("M")' in portable - assert 'T.Buffer((M, "M * T.int64(2)"), "float32")' in portable + assert 'M = I.dynamic("M", dtype="int64")' in portable + assert 'T.Buffer((M, M * T.int64(2)), "float32")' in portable assert "UNUSED" not in script - assert "M = T.int64()" in portable + assert "M = T.int64()" not in portable assert len(func.params) == 1 assert not hasattr(func, "type_params") assert func.attrs.get("tirx.type_vars") is None @@ -65,5 +65,51 @@ def func[M: int](A: T.Buffer((M, M * 2), "float32")): tvm.ir.assert_structural_equal(func, tvm.script.from_source(portable)) +def test_dynamic_int32_roundtrip(): + func = tvm.script.from_source( + """ +n = I.dynamic("n", "int32") +@T.prim_func(private=True) +def func(A: T.Buffer((n,), "float32")): + A[0] = T.float32(1) +""" + ) + source = func.script() + if sys.version_info >= (3, 12): + assert "def main[n: T.int32](" in source + else: + assert 'n = I.dynamic("n", dtype="int32")' in source + portable = func.script(extra_config={"script.use_pep695": False}) + assert 'n = I.dynamic("n", dtype="int32")' in portable + tvm.ir.assert_structural_equal(func, tvm.script.from_source(source)) + tvm.ir.assert_structural_equal(func, tvm.script.from_source(portable)) + + +def test_dynamic_module_body_identity(): + mod = tvm.script.from_source( + """ +n = I.dynamic("n", "int32") +m = I.dynamic("n", "int32") +@I.ir_module +class Module: + @T.prim_func(private=True) + def first(): + T.evaluate(n) + @T.prim_func(private=True) + def second(): + T.evaluate(n + m) +""", + check_well_formed=False, + ) + source = mod.script() + assert source.count('I.dynamic("n", dtype="int32")') == 2 + restored = tvm.script.from_source(source, check_well_formed=False) + shared = restored["first"].body.value + summed = restored["second"].body.value + assert shared.same_as(summed.a) + assert not shared.same_as(summed.b) + tvm.ir.assert_structural_equal(mod, restored, map_free_vars=True) + + if __name__ == "__main__": tvm.testing.main() diff --git a/tests/python/tvmscript/test_tvmscript_parser_tir.py b/tests/python/tvmscript/test_tvmscript_parser_tir.py index 372f25b105db..e659abbc264e 100644 --- a/tests/python/tvmscript/test_tvmscript_parser_tir.py +++ b/tests/python/tvmscript/test_tvmscript_parser_tir.py @@ -119,12 +119,12 @@ def main(A: T.Buffer(("n",), "float32"), n: T.int64): assert str(n.ty.dtype) == "int64" -def test_tir_string_defined_symbol_uses_prescanned_body_dtype(): +def test_tir_external_dynamic_symbol_preserves_dtype(): func = tvm.script.from_source( """ +n = T.dynamic("n", "int64") @T.prim_func -def main(A: T.Buffer(("n",), "float32")): - n = T.int64() +def main(A: T.Buffer((n,), "float32")): T.evaluate(n) """ ) @@ -173,7 +173,6 @@ def test_tir_return_annotation_does_not_define_symbolic_var(): """ @T.prim_func def main() -> T.Buffer(("n",), "float32"): - n = T.int32() A = T.alloc_buffer((n,), "float32") return A """ @@ -631,10 +630,11 @@ def func(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): def test_inferred_ty_with_dynamic_buffer(): """The inferred Type may contain dynamic shapes""" + M = T.dynamic("M", "int64") + N = T.dynamic("N", "int64") + @Ts.prim_func def func(a_handle: T.handle, b_handle: T.handle): - M = T.int64() - N = T.int64() A = T.match_buffer(a_handle, [M, N], "float32") B = T.match_buffer(b_handle, [M * N], "float32") for i, j in T.grid(M, N): diff --git a/tests/python/tvmscript/test_tvmscript_printer_tir.py b/tests/python/tvmscript/test_tvmscript_printer_tir.py index f4d0cd677ec2..aba2cee8fd13 100644 --- a/tests/python/tvmscript/test_tvmscript_printer_tir.py +++ b/tests/python/tvmscript/test_tvmscript_printer_tir.py @@ -26,6 +26,7 @@ from tvm.ir import Range from tvm.s_tir.script.ir_builder import prim_func as build_prim_func from tvm.script import s_tir as Ts +from tvm.script import ir as I from tvm.script.ir_builder import IRBuilder from tvm.tirx.script import ir_builder as T @@ -68,9 +69,9 @@ def test_prim_func_symbolic_buffer_param_roundtrip(): .with_attr("s_tir", True) ) - source = func.script() - assert 'T.Buffer(("n + 1", "n")' in source - assert source.index("n = T.int32()") < source.index("T.evaluate(n)") + source = func.script(extra_config={"script.use_pep695": False}) + assert "T.Buffer((n + 1, n)" in source + assert source.index('n = I.dynamic("n", dtype="int32")') < source.index("T.evaluate(n)") tvm.ir.assert_structural_equal(tvm.script.from_source(source), func) @@ -83,9 +84,9 @@ def test_prim_func_compound_buffer_shape_first_use_roundtrip(): .with_attr("s_tir", True) ) - source = func.script() - assert 'T.Buffer(("T.max(n, 1)",)' in source - assert source.index("n = T.int32()") < source.index("T.evaluate(n)") + source = func.script(extra_config={"script.use_pep695": False}) + assert "T.Buffer((T.max(n, 1),)" in source + assert source.index('n = I.dynamic("n", dtype="int32")') < source.index("T.evaluate(n)") tvm.ir.assert_structural_equal(tvm.script.from_source(source), func) @@ -171,9 +172,9 @@ def test_block_realize(): _assert_print( obj, """ -i = T.int32() -j = T.int32() -k = T.int32() +i = I.dynamic("i", dtype="int32") +j = I.dynamic("j", dtype="int32") +k = I.dynamic("k", dtype="int32") with Ts.sblock("block"): vi = Ts.axis.spatial(128, i) vj = Ts.axis.spatial(64, j) @@ -352,14 +353,14 @@ def test_assert_stmt(): def test_while(): with IRBuilder() as ib: - x = T.int32() + x = I.dynamic("v", "int32") with T.While(x < 10): T.evaluate(0) obj = ib.get() _assert_print( obj, """ -v = T.int32() +v = I.dynamic("v", dtype="int32") while v < 10: T.evaluate(0) """, @@ -474,7 +475,7 @@ def test_seq_stmt(): def test_if_then_else(): with IRBuilder() as ib: - with T.If(T.int32() == 1): + with T.If(I.dynamic("v", "int32") == 1): with T.Then(): T.evaluate(0) @@ -482,7 +483,7 @@ def test_if_then_else(): _assert_print( obj, """ -v = T.int32() +v = I.dynamic("v", dtype="int32") if v == 1: T.evaluate(0) """, @@ -506,7 +507,7 @@ def test_var(): _assert_print( a, """ -a = T.float32() +a = I.dynamic("a", dtype="float32") a""", ) @@ -532,7 +533,7 @@ def test_iter_var(): _assert_print( a, """ -a = T.int32() +a = I.dynamic("a", dtype="int32") T.iter_var(a, T.Range(0, 8), "DataPar", "") """, ) @@ -548,7 +549,7 @@ def test_cast(): _assert_print( obj, """ -a = T.float32() +a = I.dynamic("a", dtype="float32") T.Cast("float64", a) """, ) @@ -581,13 +582,13 @@ def test_binary_arith(): obj = op(a, b) if sign.isalpha(): expected = f""" -a = T.int32() -b = T.int32() +a = I.dynamic("a", dtype="int32") +b = I.dynamic("b", dtype="int32") T.{sign}(a, b)""" else: expected = f""" -a = T.int32() -b = T.int32() +a = I.dynamic("a", dtype="int32") +b = I.dynamic("b", dtype="int32") a {sign} b""" _assert_print(obj, expected) @@ -622,8 +623,8 @@ def test_int_div(): _assert_print( tirx.Div(a, b), """ -a = T.int32() -b = T.int32() +a = I.dynamic("a", dtype="int32") +b = I.dynamic("b", dtype="int32") T.Div(a, b) """, ) @@ -635,23 +636,23 @@ def test_logical(): _assert_print( tirx.And(a, b), """ -a = T.bool() -b = T.bool() +a = I.dynamic("a", dtype="bool") +b = I.dynamic("b", dtype="bool") a and b """, ) _assert_print( tirx.Or(a, b), """ -a = T.bool() -b = T.bool() +a = I.dynamic("a", dtype="bool") +b = I.dynamic("b", dtype="bool") a or b """, ) _assert_print( tirx.Not(a), """ -a = T.bool() +a = I.dynamic("a", dtype="bool") not a """, ) @@ -675,7 +676,7 @@ def test_ramp(lanes, scripted_lanes): _assert_print( obj, f""" -a = T.int32() +a = I.dynamic("a", dtype="int32") T.Ramp(a, 1, {scripted_lanes}) """, ) @@ -700,7 +701,7 @@ def test_let_expr(): _assert_print( obj, """ -x = T.int32() +x = I.dynamic("x", dtype="int32") T.Let(x + 1, where={x: 1}) """, ) @@ -917,9 +918,10 @@ def test_variable_with_cpp_address(): # The test function has all named objects suffixed with "_name", # to avoid spurious replacement when generating the expected # regex. + N_name = I.dynamic("N_name") + @Ts.prim_func def func(a_name: T.handle): - N_name = T.int64() A_name = T.match_buffer(a_name, N_name, "float32") for i_name in range(N_name): A_name[i_name] = A_name[i_name] + 1.0 diff --git a/tests/python/tvmscript/test_tvmscript_roundtrip.py b/tests/python/tvmscript/test_tvmscript_roundtrip.py index 557b5108d317..0c87b1f2c5fa 100644 --- a/tests/python/tvmscript/test_tvmscript_roundtrip.py +++ b/tests/python/tvmscript/test_tvmscript_roundtrip.py @@ -2033,12 +2033,13 @@ def constant_folding(a: T.handle) -> None: def simplify_bracket(): # uninitialized variables + a = T.dynamic("a", "int32") + b = T.dynamic("b", "int32") + c = T.dynamic("c", "int32") + d = T.dynamic("d", "int32") + @Ts.prim_func(check_well_formed=False) def simplify_bracket() -> None: - a = T.int32() - b = T.int32() - c = T.int32() - d = T.int32() T.evaluate(a + b * (c + d)) return simplify_bracket @@ -2165,10 +2166,11 @@ def multiple_commreducer() -> None: def func_div_mod(): # not well-formed: free variables + a = T.dynamic("a", "int32") + b = T.dynamic("b", "int32") + @Ts.prim_func(check_well_formed=False) def func_div_mod(): - a = T.int32() - b = T.int32() T.evaluate(a // b) T.evaluate(a % b) T.evaluate(T.truncmod(a, b)) @@ -2480,9 +2482,10 @@ def func(a: T.handle, b: T.handle): def let_expression(): + x = T.dynamic("x", "int32") + @Ts.prim_func def func(): - x = T.int32() T.evaluate(T.Let(x + 1, where={x: 1})) return func @@ -2640,9 +2643,10 @@ def func() -> None: def bool_cast(): # uninitialized var + a = T.dynamic("a", "bool") + @Ts.prim_func(check_well_formed=False) def func() -> None: - a = T.bool() T.evaluate(T.bool(T.int32(0))) T.evaluate(a == T.bool(False)) @@ -2966,9 +2970,10 @@ def func(): def undefined_shape_in_decl_buffer(): # uninitialized var + size = T.dynamic("size", "int32") + @Ts.prim_func(check_well_formed=False) def func(): - size = T.int32() buf = T.decl_buffer(shape=[size], dtype="float32") T.evaluate(buf[0]) @@ -2977,9 +2982,10 @@ def func(): def undefined_stride_in_decl_buffer(): # uninitialized var + stride = T.dynamic("stride", "int32") + @Ts.prim_func(check_well_formed=False) def func(): - stride = T.int32() data_ptr = T.handle("float32") buf = T.decl_buffer(shape=[1], dtype="float32", data=data_ptr, strides=[stride]) T.evaluate(buf[0]) @@ -2989,9 +2995,10 @@ def func(): def undefined_elem_offset_in_decl_buffer(): # uninitialized var + elem_offset = T.dynamic("elem_offset", "int32") + @Ts.prim_func(check_well_formed=False) def func(): - elem_offset = T.int32() data_ptr = T.handle("float32") buf = T.decl_buffer(shape=[1], dtype="float32", data=data_ptr, elem_offset=elem_offset) T.evaluate(buf[0]) @@ -3199,7 +3206,7 @@ def func(A: R.Any): def relax_symbolic_var(): """Relax tensors may use symbolic variables.""" - N = tvm.tirx.Var("N", "int64") + N = T.dynamic("N", "int64") @R.function def func(A: R.Tensor([N], "float16")): diff --git a/tests/python/tvmscript/test_tvmscript_syntax_sugar.py b/tests/python/tvmscript/test_tvmscript_syntax_sugar.py index d58447b57808..752a4b7b90f3 100644 --- a/tests/python/tvmscript/test_tvmscript_syntax_sugar.py +++ b/tests/python/tvmscript/test_tvmscript_syntax_sugar.py @@ -159,11 +159,13 @@ def func_with_sugar(A: T.Buffer(16, "float32")): # dynamic shape gemm +N = T.dynamic("N", "int32") +M = T.dynamic("M", "int32") +K = T.dynamic("K", "int32") + + @Ts.prim_func def gemm_dyn_shape(a: T.handle, b: T.handle, c: T.handle): - N = T.int32() - M = T.int32() - K = T.int32() A = T.match_buffer(a, (N, K), "float32") B = T.match_buffer(b, (K, M), "float32") C = T.match_buffer(c, (N, M), "float32") @@ -415,9 +417,10 @@ def test_preserve_trivial_let_binding(): builder API and the `j: T.let[T.dtype]` annotation produce the same LetStmt IR. """ + j = T.dynamic("j", "int32") + @Ts.prim_func def explicit(i: T.int32): - j = T.int32() T.bind(i, var=j) T.evaluate(j) @@ -432,9 +435,10 @@ def implicit(i: T.int32): def test_preserve_trivial_let_binding_of_value(): """Same as test_preserve_trivial_let_binding but with a constant RHS.""" + j = T.dynamic("j", "int32") + @Ts.prim_func def explicit(i: T.int32): - j = T.int32() T.bind(42, var=j) T.evaluate(j) From 2b1d91b291603b1071152dde254a210b83238e64 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 20:57:37 +0000 Subject: [PATCH 02/16] [Script] Use explicit expressions for symbolic dimensions Remove expression-string policies and implicit symbolic lookup from the parser and eager constructors. Keep whole annotation decoding and generic annotation adaptation, and print concrete expressions with explicit dynamic symbols. --- .../relax/tutorials/relax_creation.py | 6 +- .../mix_python_and_tvm_with_pymodule.py | 6 +- .../tvm/relax/script/ir_builder/__init__.py | 11 +- python/tvm/script/ir_builder/base.py | 92 +++++------ .../tvm/script/ir_builder/parser_protocol.py | 3 +- .../{expr_str_handling.py => annotation.py} | 11 +- python/tvm/script/parser/inspect_source.py | 2 +- python/tvm/script/parser/prescan.py | 4 +- python/tvm/script/parser/protocol_registry.py | 144 +++--------------- python/tvm/script/parser/transpile.py | 126 ++------------- python/tvm/tirx/script/ir_builder/__init__.py | 29 +--- python/tvm/tirx/script/ir_builder/ir.py | 12 +- src/relax/script/printer/dependent_type.cc | 46 +----- src/relax/script/printer/distributed.cc | 2 +- src/relax/script/printer/expr.cc | 2 +- src/relax/script/printer/function.cc | 3 + src/relax/script/printer/utils.h | 2 - src/script/printer/ir/ir.cc | 16 +- src/tirx/script/printer/buffer.cc | 60 +++----- src/tirx/script/printer/function.cc | 44 +++--- src/tirx/script/printer/utils.h | 4 +- tests/python/relax/test_tvmscript_parser.py | 78 +++++----- .../relax/test_tvmscript_printer_relax.py | 12 +- .../python/relax/test_tvmscript_type_vars.py | 4 +- tests/python/script/minilang.py | 1 - tests/python/script/test_basic_usage.py | 3 +- tests/python/script/test_meta_programming.py | 13 +- tests/python/script/test_parser_entry.py | 2 +- .../script/test_special_parser_protocol.py | 8 +- tests/python/script/test_symbolic_shape.py | 33 ++-- .../tvmscript/test_tvmscript_error_report.py | 16 +- .../tvmscript/test_tvmscript_parser_tir.py | 35 +++-- 32 files changed, 266 insertions(+), 564 deletions(-) rename python/tvm/script/parser/{expr_str_handling.py => annotation.py} (92%) diff --git a/docs/deep_dive/relax/tutorials/relax_creation.py b/docs/deep_dive/relax/tutorials/relax_creation.py index 949313372de3..4db2191fd3bf 100644 --- a/docs/deep_dive/relax/tutorials/relax_creation.py +++ b/docs/deep_dive/relax/tutorials/relax_creation.py @@ -43,17 +43,19 @@ from tvm.script import s_tir as Ts from tvm.script import tirx as T +n = T.dynamic("n") + @I.ir_module class RelaxModule: @R.function def forward( - data: R.Tensor(("n", 784), dtype="float32"), + data: R.Tensor((n, 784), dtype="float32"), w0: R.Tensor((128, 784), dtype="float32"), b0: R.Tensor((128,), dtype="float32"), w1: R.Tensor((10, 128), dtype="float32"), b1: R.Tensor((10,), dtype="float32"), - ) -> R.Tensor(("n", 10), dtype="float32"): + ) -> R.Tensor((n, 10), dtype="float32"): with R.dataflow(): lv0 = R.matmul(data, R.permute_dims(w0)) + b0 lv1 = R.nn.relu(lv0) diff --git a/docs/how_to/tutorials/mix_python_and_tvm_with_pymodule.py b/docs/how_to/tutorials/mix_python_and_tvm_with_pymodule.py index 709f1ad07949..a415d5fe43d8 100644 --- a/docs/how_to/tutorials/mix_python_and_tvm_with_pymodule.py +++ b/docs/how_to/tutorials/mix_python_and_tvm_with_pymodule.py @@ -412,9 +412,9 @@ def scale_tir(var_x: T.handle, var_out: T.handle): @R.function def add_relax( - x: R.Tensor(("n",), "float32"), - y: R.Tensor(("n",), "float32"), - ) -> R.Tensor(("n",), "float32"): + x: R.Tensor((n,), "float32"), + y: R.Tensor((n,), "float32"), + ) -> R.Tensor((n,), "float32"): return R.add(x, y) mod = DynamicModule(device=tvm.cpu(0), target="llvm") diff --git a/python/tvm/relax/script/ir_builder/__init__.py b/python/tvm/relax/script/ir_builder/__init__.py index 29987e713f99..f296028e5624 100644 --- a/python/tvm/relax/script/ir_builder/__init__.py +++ b/python/tvm/relax/script/ir_builder/__init__.py @@ -28,10 +28,10 @@ from tvm.relax.distributed import Placement as _Placement from tvm.relax.distributed import device_mesh as device_mesh from tvm.script.ir_builder import resolve_global_info_args as _resolve_global_info_args +from tvm.script.ir_builder import IRBuilder as _IRBuilder +from tvm.script.ir_builder.base import annotation_constructor as _annotation_constructor from tvm.script.ir_builder.base import at as _at from tvm.script.ir_builder.base import source_span as _source_span -from tvm.script.parser.protocol_registry import ARGS_POLICIES as _ARGS_POLICIES -from tvm.script.parser.protocol_registry import args_policy as _args_policy from tvm.script.parser.protocol_registry import constexpr as constexpr from . import distributed as dist @@ -84,7 +84,7 @@ @_resolve_global_info_args("vdevice", resolver=resolve_global_info_) -@_args_policy("R.Tensor", {"shape": "expr_str"}, scalar_strings=False) +@_annotation_constructor("shape") def Tensor(shape=None, dtype=None, vdevice=None, ndim=-1, *, span=None): """Construct a Relax tensor type. @@ -119,7 +119,7 @@ def Tensor(shape=None, dtype=None, vdevice=None, ndim=-1, *, span=None): @_resolve_global_info_args("device_mesh", resolver=resolve_global_info_) -@_args_policy("R.DTensor", {"shape": "expr_str"}, scalar_strings=False) +@_annotation_constructor("shape") def DTensor(shape=None, dtype=None, device_mesh=None, placement="", *, ndim=-1, span=None): """Construct a Relax distributed tensor type. @@ -157,7 +157,6 @@ def DTensor(shape=None, dtype=None, device_mesh=None, placement="", *, ndim=-1, # The distributed source spelling shares concrete constructors and argument policy. dist.DTensor = DTensor -_ARGS_POLICIES["R.dist.DTensor"] = _ARGS_POLICIES["R.DTensor"] dist.device_mesh = device_mesh Range = _ir.Range @@ -166,7 +165,7 @@ def DTensor(shape=None, dtype=None, device_mesh=None, placement="", *, ndim=-1, __tvm_value_if__ = True -@_args_policy("R.Shape", {"values": "expr_str"}, dtype="int64") +@_annotation_constructor("values") def Shape(values=None, ndim=-1, *, span=None): """Construct a Relax shape type. diff --git a/python/tvm/script/ir_builder/base.py b/python/tvm/script/ir_builder/base.py index 14a960e9ca4e..bcba3d2efacb 100644 --- a/python/tvm/script/ir_builder/base.py +++ b/python/tvm/script/ir_builder/base.py @@ -473,52 +473,53 @@ def _current_function_frame(): raise ValueError("Symbol resolution requires an active function frame") -def wrap_expression_constructor(constructor, call_signature, policy, *, as_type=False): - """Adapt an eager constructor using parser-owned expression-string metadata.""" - fields = policy.fields - - def unresolved(value, nested=False): - if isinstance(value, str): - return nested or policy.scalar_strings - if isinstance(value, TypeVar): - return True - if isinstance(value, tuple | list): - return any(unresolved(item, True) for item in value) - return False - - @wraps(constructor) - def invoke(*args, **kwargs): - bound = call_signature.bind(*args, **kwargs) - if IRBuilder.is_in_scope(): - # typing.TypeVar is ordinary eager Python metadata. Resolve it - # here, never in the syntax-only transpiler. - def resolve(value): - if isinstance(value, TypeVar): - if value.__bound__ is not None or value.__constraints__: - raise TypeError("A symbolic TypeVar cannot have constraints or a bound") - return _current_function_frame().resolve_type_var(value.__name__) - if isinstance(value, tuple): - return tuple(resolve(item) for item in value) - if isinstance(value, list): - return [resolve(item) for item in value] - return value +def annotation_constructor(*fields: str, as_type: bool = False): + """Adapt eager Python type parameters on concrete annotation constructors. + Unresolved ``typing.TypeVar`` values defer annotations outside a builder; + an active native function owns their resolution. Ordinary values, including + strings, are passed unchanged to the concrete API. + """ + + def decorate(constructor): + call_signature = signature(constructor) + + def unresolved(value): + if isinstance(value, TypeVar): + return True + if isinstance(value, tuple | list): + return any(unresolved(item) for item in value) + return False + + def resolve(value): + if isinstance(value, TypeVar): + if value.__bound__ is not None or value.__constraints__: + raise TypeError("A symbolic TypeVar cannot have constraints or a bound") + return _current_function_frame().resolve_type_var(value.__name__) + if isinstance(value, tuple): + return tuple(resolve(item) for item in value) + if isinstance(value, list): + return [resolve(item) for item in value] + return value + + @wraps(constructor) + def invoke(*args, **kwargs): + bound = call_signature.bind(*args, **kwargs) for field in fields: - if field in bound.arguments: - bound.arguments[field] = resolve(bound.arguments[field]) - if any(unresolved(bound.arguments[field]) for field in fields if field in bound.arguments): - if IRBuilder.is_in_scope(): - raise TypeError( - "Builder expression arguments require concrete symbols, not strings" - ) - return ir.Type.missing() - return constructor(*bound.args, **bound.kwargs) - - result = invoke - if as_type: - # The class is an annotation surface, not an IR or proxy type. - # __new__ returns the concrete construction result (or MissingType). - result = type( + if field not in bound.arguments: + continue + value = bound.arguments[field] + if IRBuilder.is_in_scope(): + bound.arguments[field] = resolve(value) + elif unresolved(value): + return ir.Type.missing() + return constructor(*bound.args, **bound.kwargs) + + if not as_type: + return invoke + # A real annotation class supports Python unions while constructing + # ordinary native types, with no proxy values or parser policy state. + return type( constructor.__name__, (), { @@ -528,7 +529,8 @@ def resolve(value): "__module__": constructor.__module__, }, ) - return result + + return decorate def _return_annotation(annotation): diff --git a/python/tvm/script/ir_builder/parser_protocol.py b/python/tvm/script/ir_builder/parser_protocol.py index ae5c115938dd..258d5b9b3b66 100644 --- a/python/tvm/script/ir_builder/parser_protocol.py +++ b/python/tvm/script/ir_builder/parser_protocol.py @@ -757,7 +757,7 @@ def resolve_type_var_( Parameters ---------- name : str - Function-local lookup key; quoted symbols do not create a Python binding. + Function-local lookup key for explicit header parameters and captured symbols. dtype : str, Type or Var, optional Explicit primitive type or supplied variable. None defaults new symbols to int64; existing symbols must agree with an explicit dtype. @@ -871,7 +871,6 @@ class Module: # # Syntax markers live in tvm.script.parser.protocol_registry. # ``constexpr(value)`` selects host evaluation in marked control flow. -# ``args_policy(path, fields)`` marks expression-string arguments. # ``register_scalar_annotation(path, constructor, dtype=...)`` describes scalar annotations. # ``mutable_cell_decl(path)`` marks mutable storage declarations. # ``result_span(path)`` permits attaching a call's result span without a call context. diff --git a/python/tvm/script/parser/expr_str_handling.py b/python/tvm/script/parser/annotation.py similarity index 92% rename from python/tvm/script/parser/expr_str_handling.py rename to python/tvm/script/parser/annotation.py index 25c37dc9c397..e95fb571e591 100644 --- a/python/tvm/script/parser/expr_str_handling.py +++ b/python/tvm/script/parser/annotation.py @@ -50,9 +50,9 @@ def _parse_string_expression(self, node: ast.Constant) -> ast.expr: node.end_col_offset + 1, ), ) from error - # Example: in R.Tensor(("n + 1",), "float32"), the generated resolve - # call for n retains the byte range of n inside the quoted literal, not - # the whole constructor. A triple-quoted physical newline advances lineno; + # In a quoted whole annotation such as "R.Tensor((n + 1,), 'float32')", + # n retains its byte range inside the literal, not the whole constructor. + # A triple-quoted physical newline advances lineno; # an escaped \n advances decoded input but maps back to the escape's real # source bytes. Generated call wrappers copy these four mapped fields. source = "".join(linecache.getlines(self.filename)) @@ -129,8 +129,3 @@ def parse_annotation(node: ast.expr, filename: str) -> ast.expr: if isinstance(node, ast.Constant) and isinstance(node.value, str): return _LiteralParser(filename)._parse_string_expression(node) return node - - -def parse_expression_string(node: ast.Constant, filename: str) -> ast.expr: - """Decode one policy-marked string, preserving its physical source ranges.""" - return _LiteralParser(filename)._parse_string_expression(node) diff --git a/python/tvm/script/parser/inspect_source.py b/python/tvm/script/parser/inspect_source.py index fb9da9177292..26b5603c9486 100644 --- a/python/tvm/script/parser/inspect_source.py +++ b/python/tvm/script/parser/inspect_source.py @@ -36,7 +36,7 @@ from tvm.ir import SourceName, Span -from .expr_str_handling import parse_annotation +from .annotation import parse_annotation from .prescan import collect_annotation_free_names diff --git a/python/tvm/script/parser/prescan.py b/python/tvm/script/parser/prescan.py index 36177e92c97b..59de07454999 100644 --- a/python/tvm/script/parser/prescan.py +++ b/python/tvm/script/parser/prescan.py @@ -25,7 +25,7 @@ from typing import NamedTuple, NoReturn from . import protocol_registry as protocol -from .expr_str_handling import parse_annotation +from .annotation import parse_annotation def collect_annotation_free_names( @@ -420,7 +420,6 @@ def visit_FunctionDef(self, node: ast.FunctionDef | ast.AsyncFunctionDef) -> Non annotation.func if isinstance(annotation, ast.Call) else annotation, self.environment, ) - dtype = protocol.SCALAR_ANNOTATION_DTYPE.get(constructor) self._record_binding( arg.arg, arg, @@ -430,7 +429,6 @@ def visit_FunctionDef(self, node: ast.FunctionDef | ast.AsyncFunctionDef) -> Non and "parameter" in protocol.MUTABLE_CELL_DECL.get(constructor, ()) else "parameter", arg.annotation, - dtype if not isinstance(annotation, ast.Call) else None, ) for statement in node.body: self.visit(statement) diff --git a/python/tvm/script/parser/protocol_registry.py b/python/tvm/script/parser/protocol_registry.py index 396a709bb462..f0119cf66289 100644 --- a/python/tvm/script/parser/protocol_registry.py +++ b/python/tvm/script/parser/protocol_registry.py @@ -25,8 +25,7 @@ or receiver inference. Tables retain static syntax facts only, never source functions, captures, frames -or constructed results. Registration decorators return the same callable, except -that ``args_policy`` preserves its existing builder-owned eager annotation adapter. +or constructed results. Registration decorators return the same callable. Shared builder operations are documented in ``tvm.script.ir_builder.parser_protocol``; concrete language variants own registration and namespace initialization. @@ -34,43 +33,12 @@ from __future__ import annotations -import ast -from collections.abc import Callable, Mapping -from inspect import signature -from types import MappingProxyType -from typing import Any, Literal, NamedTuple, NoReturn, TypeVar +from collections.abc import Callable +from typing import Any, Literal, NoReturn, TypeVar _Callable = TypeVar("_Callable", bound=Callable[..., Any]) -class ExprStrPolicy(NamedTuple): - """Expression-string fields, symbolic dtype and bare-string interpretation. - - Nested strings in marked fields always represent expressions. ``scalar_strings`` - controls bare strings. The dtype is forwarded to native symbol resolution; - this static record contains no resolved symbols or construction state. - """ - - fields: tuple[str, ...] - dtype: object = None - scalar_strings: bool = True - - -class ArgsPolicy(NamedTuple): - """Argument policies and positional names computed once at registration. - - ``fields`` maps parameter names to ``expr_str``. - ``expression`` describes the expression-string subset. ``positional_parameters`` - lists positional-only and positional-or-keyword names in signature order; - keyword-only parameters remain available through ``fields``. - """ - - fields: Mapping[str, str] - expression: ExprStrPolicy - positional_parameters: tuple[str, ...] - - -ARGS_POLICIES: dict[str, ArgsPolicy] = {} SCALAR_ANNOTATION_DTYPE: dict[str, object] = {} MUTABLE_CELL_DECL: dict[str, frozenset[str]] = {} RESULT_SPAN: dict[str, bool] = {} @@ -111,113 +79,41 @@ def constexpr(value: object) -> NoReturn: raise TypeError("constexpr is a parser syntax marker, not a runtime operation") -def args_policy( +def register_scalar_annotation( namespace_path: str, - fields: Mapping[str, str], + constructor: _Callable, *, - scalar_strings: bool = True, dtype: object = None, - as_type: bool = False, -) -> Callable[[Callable[..., Any]], Callable[..., Any]]: - """Register argument syntax at an explicit canonical namespace path. +) -> _Callable: + """Register the dtype of a fixed scalar annotation without evaluating it. Parameters ---------- namespace_path : str - Registered namespace alias followed by the exported member path. - fields : Mapping[str, str] - Parameter names mapped to ``expr_str``. - scalar_strings : bool, optional - Whether bare strings in expression fields denote expressions. Defaults to True. + Canonical registered namespace alias and exported callable path, such + as ``"T.int32"``. A later registration at this path replaces its dtype. + constructor : Callable + Callable providing the eager construction operation. Registration + records its syntax path without invoking or wrapping this callable. dtype : object, optional - Opaque dtype passed to the language variant's symbol resolver. None (the default) - leaves dtype selection to that resolver. - as_type : bool, optional - Preserve the eager annotation-class surface, including Python type unions. - Defaults to False. + Static scalar dtype for an explicit PEP 695 symbol bound. None (the + default) leaves the annotation without a supported scalar bound dtype. Returns ------- Callable - Decorator retaining the existing builder-owned annotation adapter where - needed. Concrete arguments still invoke the original constructor. + The exact ``constructor`` object. Notes ----- - Registration validates policy kinds and parameter names, and inspects the - signature once. Static policy and positional names are reused for every call. - Outside construction, unresolved expression annotations yield MissingType; - inside construction, unresolved strings raise TypeError and TypeVars resolve - through the active function frame. Those eager semantics belong to the builder. + This metadata does not give constructor calls special assignment semantics. + Scalar runtime parameters retain their ordinary annotation construction. .. code:: python - @args_policy("M.Tensor", {"shape": "expr_str"}) - def Tensor(shape): - ... - # Source: M.Tensor(("n",)) - # Builder: M.Tensor((M.resolve_type_var_("n"),)) - """ - fields = dict(fields) - unsupported = set(fields.values()).difference(("expr_str",)) - if unsupported: - raise ValueError(f"Unknown argument policies: {sorted(unsupported)}") - - def decorate(constructor: Callable[..., Any]) -> Callable[..., Any]: - call_signature = signature(constructor) - unknown = set(fields).difference(call_signature.parameters) - if unknown: - raise ValueError(f"Unknown argument policy fields: {sorted(unknown)}") - positional_parameters = tuple( - parameter.name - for parameter in call_signature.parameters.values() - if parameter.kind in (parameter.POSITIONAL_ONLY, parameter.POSITIONAL_OR_KEYWORD) - ) - expression_fields = tuple(name for name, kind in fields.items() if kind == "expr_str") - expression = ExprStrPolicy(expression_fields, dtype, bool(scalar_strings)) - if expression_fields or as_type: - from tvm.script.ir_builder.base import wrap_expression_constructor - - result = wrap_expression_constructor( - constructor, call_signature, expression, as_type=as_type - ) - else: - result = constructor - ARGS_POLICIES[namespace_path] = ArgsPolicy( - MappingProxyType(fields.copy()), expression, positional_parameters - ) - return result - - return decorate - - -def handle_call_args_policy( - node: ast.Call, resolve: Callable[[ast.expr], str | None] -) -> tuple[ArgsPolicy, tuple[str, ...]] | None: - """Select a fixed namespace call's policy before visiting its arguments. - - ``resolve`` returns only canonical namespace paths, without evaluating a - callee or receiver. Unregistered calls return None. Matching calls reuse the - positional names stored at registration; expression-string rewriting remains - in ``expr_str_handling`` and the normal call visitor traverses children once. - """ - namespace_path = resolve(node.func) - if namespace_path is None: - return None - policy = ARGS_POLICIES.get(namespace_path) - return (policy, policy.positional_parameters) if policy is not None else None - - -def register_scalar_annotation( - namespace_path: str, - constructor: _Callable, - *, - dtype: object = None, -) -> _Callable: - """Register the dtype of a fixed scalar annotation without evaluating it. - - Used by scalar function parameters and explicit PEP 695 symbol bounds. - This metadata does not give constructor calls special assignment semantics. + register_scalar_annotation("T.int32", T.int32, dtype="int32") + # Source: def f[n: T.int32](...): + # Builder: n = X.resolve_type_var_("n", dtype="int32") """ SCALAR_ANNOTATION_DTYPE[namespace_path] = dtype return constructor diff --git a/python/tvm/script/parser/transpile.py b/python/tvm/script/parser/transpile.py index b0d51957b2f5..2e4f15b33103 100644 --- a/python/tvm/script/parser/transpile.py +++ b/python/tvm/script/parser/transpile.py @@ -44,7 +44,7 @@ from typing import Any, NoReturn, TypeVar from . import protocol_registry as protocol -from .expr_str_handling import parse_annotation, parse_expression_string +from .annotation import parse_annotation from .prescan import Binding, PrescanContext, resolve_namespace_key, resolve_namespace_value _Node = TypeVar("_Node", bound=ast.AST) @@ -157,8 +157,6 @@ def __init__( # Bypass syntax-to-builder lowering. # Still traverse children and instrument source calls with spans. self.bypass_ast_rewrite = False - # Temporary policy only while this same visitor processes a decoded string. - self.expression_string: protocol.ExprStrPolicy | None = None # Only annotation syntax consults the function's one substitution map; # ordinary body globals keep Python lookup even when names coincide. self.annotation_expression = False @@ -361,28 +359,6 @@ def _read_constexpr_operand(self, node: ast.expr) -> ast.expr | None: def visit_Name(self, node: ast.Name) -> ast.expr: # -------------------- Pattern -------------------- # Python source: - # X.Tensor(("n",)) - # - # Builder: - # X.Tensor((X.resolve_type_var_("n"),)) - # ------------------------------------------------- - # Only the marked string is symbolic; no Python binding is introduced. - if self.expression_string is not None: - keywords = {} - if self.expression_string.dtype is not None: - keywords["dtype"] = ast.Constant(self.expression_string.dtype) - return self._attach_span( - self._call( - self.function.dialect_prefix, - "resolve_type_var_", - [ast.Constant(node.id)], - node, - **keywords, - ), - node, - ) - # -------------------- Pattern -------------------- - # Python source: # value: X.Tensor((n,)) # # Builder: @@ -428,11 +404,6 @@ def visit_Name(self, node: ast.Name) -> ast.expr: return node def visit_Attribute(self, node: ast.Attribute) -> ast.expr: - # Within a decoded string, known namespace attributes remain Python - # lookup; unknown roots still denote symbolic variables. - if self.expression_string is not None and self._resolve(node.value) is not None: - with self._use_string_policy(None): - return self.visit_Attribute(node) # -------------------- Pattern -------------------- # Python source: # Module.f @@ -576,50 +547,11 @@ def _is_module_owner(self, node: ast.expr) -> bool: ] return bool(records) and all(item.kind == "module_alias" for item in records) - @contextmanager - def _use_string_policy(self, policy: protocol.ExprStrPolicy | None) -> Iterator[None]: - """Restore decoded-string dtype context after normal or failed traversal.""" - previous, self.expression_string = self.expression_string, policy - try: - yield - finally: - self.expression_string = previous - - def _rewrite_policy_argument( - self, - node: ast.expr, - kind: str | None, - policy: protocol.ExprStrPolicy, - *, - nested: bool = False, - ) -> ast.expr: - """Visit an argument once, assembling policy lookups after source children.""" - if kind == "expr_str" and isinstance(node, ast.Tuple | ast.List): - # -------------------- Pattern -------------------- - # Python source: - # X.Tensor(("n", value)) - # - # Builder: - # X.Tensor((X.resolve_type_var_("n"), value)) - # ------------------------------------------------- - node.elts = [ - self._rewrite_policy_argument(item, kind, policy, nested=True) for item in node.elts - ] - return node if self.bypass_ast_rewrite else self._attach_span(node, node) - if isinstance(node, ast.Constant) and isinstance(node.value, str): - if kind == "expr_str" and (nested or policy.scalar_strings): - with self._use_string_policy(policy): - return self.visit(parse_expression_string(node, self.module.filename)) - return self.visit(node) def _visit_direct_operand(self, node: ast.expr) -> ast.expr: """Preserve an existing payload's span without bypassing child operations.""" if isinstance(node, ast.Name): - return ( - node - if not self.annotation_expression and self.expression_string is None - else self.visit(node) - ) + return self.visit(node) if self.annotation_expression else node if isinstance(node, ast.Attribute): node.value = self._visit_direct_operand(node.value) return node @@ -634,21 +566,16 @@ def _visit_direct_operand(self, node: ast.expr) -> ast.expr: def visit_Call(self, node: ast.Call, *, callee: ast.expr | None = None) -> ast.expr: # -------------------- Pattern -------------------- # Python source: - # X.Tensor(("n",), vdevice="cuda:0") + # X.Tensor((n,), vdevice="cuda:0") # # Builder: - # X.Tensor((X.resolve_type_var_("n"),), vdevice="cuda:0") + # X.Tensor((n,), vdevice="cuda:0") # ------------------------------------------------- binding_value = node is self.binding_expression marker = self._read_constexpr_operand(node) if marker is not None: with self._bypass_rewrite(): return self.visit(marker) - selected = ( - None - if self.bypass_ast_rewrite - else protocol.handle_call_args_policy(node, self._resolve) - ) constructor = self._resolve(node.func) global_call = ( isinstance(node.func, ast.Name) and node.func.id in self.module.module_functions @@ -660,28 +587,10 @@ def visit_Call(self, node: ast.Call, *, callee: ast.expr | None = None) -> ast.e # Visit source callee/arguments first. A normalized range callee is # assembled afterward, but its original arguments keep normal rewriting. if callee is None: - if self.expression_string is not None and constructor is not None: - with self._use_string_policy(None): - node.func = self._visit_direct_operand(node.func) - else: - node.func = self._visit_direct_operand(node.func) - if selected is None: - node.args = [self.visit(value) for value in node.args] - for keyword in node.keywords: - keyword.value = self.visit(keyword.value) - else: - policy, parameters = selected - known_position = True - for index, value in enumerate(node.args): - known_position = known_position and not isinstance(value, ast.Starred) - name = parameters[index] if known_position and index < len(parameters) else None - node.args[index] = self._rewrite_policy_argument( - value, policy.fields.get(name), policy.expression - ) - for keyword in node.keywords: - keyword.value = self._rewrite_policy_argument( - keyword.value, policy.fields.get(keyword.arg), policy.expression - ) + node.func = self._visit_direct_operand(node.func) + node.args = [self.visit(value) for value in node.args] + for keyword in node.keywords: + keyword.value = self.visit(keyword.value) if callee is not None: node.func = callee # -------------------- Pattern -------------------- @@ -1816,9 +1725,9 @@ def _create_specialization_bindings( return [self._assign(special, special_expr, node)] def _create_symbol_declarations( - self, node: ast.FunctionDef, facts: list[Binding] + self, node: ast.FunctionDef ) -> tuple[list[ast.stmt], dict[str, str]]: - """Predeclare symbol types and bind explicit signature type parameters.""" + """Bind explicit signature type parameters with their declared dtypes.""" declaration: list[ast.stmt] = [] symbol_aliases: dict[str, str] = {} # -------------------- Pattern -------------------- @@ -1829,8 +1738,6 @@ def _create_symbol_declarations( # Builder: # n = X.resolve_type_var_("n") # ------------------------------------------------- - # Explicit symbol dtypes precede quoted shapes; only explicit type - # parameters bind signature names. for parameter in getattr(node, "type_params", ()): if not isinstance(parameter, getattr(ast, "TypeVar", ())): self._raise_error(parameter, "Only scalar type parameters are supported") @@ -1858,17 +1765,6 @@ def _create_symbol_declarations( parameter, ) ) - for item in facts: - if item.direct and item.dtype is not None and item.kind == "parameter": - symbol = self._call_dialect( - "resolve_type_var_", - [ast.Constant(item.name)], - item.node, - dtype=ast.Constant(item.dtype), - ) - # A later Python parameter name does not enter annotation scope - # until its own arg, even though the native map knows its dtype. - declaration.append(ast.copy_location(ast.Expr(symbol), item.node)) return declaration, symbol_aliases @staticmethod @@ -2213,7 +2109,7 @@ def create_function_builder_fragments( node, ) ] - symbols, symbol_aliases = self._create_symbol_declarations(node, facts) + symbols, symbol_aliases = self._create_symbol_declarations(node) declaration.extend(symbols) with ( self._use_aliases({**definition_aliases, **symbol_aliases}), diff --git a/python/tvm/tirx/script/ir_builder/__init__.py b/python/tvm/tirx/script/ir_builder/__init__.py index 1852ba3ef232..cec51e5d2e33 100644 --- a/python/tvm/tirx/script/ir_builder/__init__.py +++ b/python/tvm/tirx/script/ir_builder/__init__.py @@ -23,11 +23,10 @@ from tvm import ir as _ir from tvm import tirx as _tir +from tvm.script.ir_builder.base import annotation_constructor as _annotation_constructor from tvm.script.ir_builder.base import at as _at from tvm.script.ir_builder.base import source_span as _source_span from tvm.script.ir_builder import dynamic as dynamic -from tvm.script.parser.protocol_registry import ARGS_POLICIES as _ARGS_POLICIES -from tvm.script.parser.protocol_registry import args_policy as _args_policy from tvm.script.parser.protocol_registry import constexpr as constexpr from tvm.script.parser.protocol_registry import ( mutable_cell_decl as _mutable_cell_decl, @@ -99,16 +98,8 @@ @_result_span("T.Buffer") @_mutable_cell_decl("T.Buffer", syntax="parameter") -@_args_policy( - "T.Buffer", - { - "shape": "expr_str", - "strides": "expr_str", - "elem_offset": "expr_str", - "byte_offset": "expr_str", - "allocated_addr": "expr_str", - }, - as_type=True, +@_annotation_constructor( + "shape", "strides", "elem_offset", "byte_offset", "allocated_addr", as_type=True ) def Buffer( shape, @@ -194,7 +185,6 @@ def Buffer( buffer = _mutable_cell_decl("T.buffer", syntax="parameter")(Buffer) -_ARGS_POLICIES["T.buffer"] = _ARGS_POLICIES["T.Buffer"] def Ptr(dtype, storage_scope="global", *, span=None): @@ -352,15 +342,6 @@ def shared_scalar(dtype="float32"): @_mutable_cell_decl("T.match_buffer") -@_args_policy( - "T.match_buffer", - { - "shape": "expr_str", - "strides": "expr_str", - "elem_offset": "expr_str", - "allocated_addr": "expr_str", - }, -) @_wraps(_native.match_buffer) def match_buffer(*args, **kwargs): """The buffer match function. @@ -427,8 +408,8 @@ def match_buffer(*args, **kwargs): Notes ----- - Shape, stride, element-offset and allocation-address expression strings are - resolved by the construction protocol before the native buffer match is created. + Shape, stride, element-offset and allocation-address expressions use + concrete primitive values, including externally constructed dynamic symbols. """ return _native.match_buffer(*args, **kwargs) diff --git a/python/tvm/tirx/script/ir_builder/ir.py b/python/tvm/tirx/script/ir_builder/ir.py index da39f8737a77..53efba58f30e 100644 --- a/python/tvm/tirx/script/ir_builder/ir.py +++ b/python/tvm/tirx/script/ir_builder/ir.py @@ -325,9 +325,7 @@ def buffer( """ shape = (shape,) if is_prim_expr(shape) or isinstance(shape, Integral) else shape shape = tuple(shape) - if strides is not None: - strides = [Var(s, "int64") if isinstance(s, str) else s for s in strides] - else: + if strides is None: strides = [] if allocated_addr is None: allocated_addr = [] @@ -537,9 +535,7 @@ def match_buffer( else: raise ValueError("Shape must be specified when binding input param") shape = (shape,) if is_prim_expr(shape) or isinstance(shape, Integral) else shape - if strides is not None: - strides = [Var(s, "int64") if isinstance(s, str) else s for s in strides] - else: + if strides is None: strides = [] if allocated_addr is None: allocated_addr = [] @@ -1527,9 +1523,7 @@ def decl_buffer( """ shape = (shape,) if is_prim_expr(shape) or isinstance(shape, Integral) else shape shape = tuple(shape) - if strides is not None: - strides = [Var(s, "int64") if isinstance(s, str) else s for s in strides] - else: + if strides is None: strides = [] dtype = _normalize_prim_type(dtype) decl_frame = _ffi_api.DeclBuffer( # type: ignore[attr-defined] # pylint: disable=no-member diff --git a/src/relax/script/printer/dependent_type.cc b/src/relax/script/printer/dependent_type.cc index 027f64c417dc..026819dcd9ab 100644 --- a/src/relax/script/printer/dependent_type.cc +++ b/src/relax/script/printer/dependent_type.cc @@ -17,8 +17,6 @@ * under the License. */ #include -#include -#include #include "../../../script/printer/ir/utils.h" #include "./utils.h" @@ -32,46 +30,6 @@ TVM_FFI_STATIC_INIT_BLOCK() { "", [](relax::AnyType n, AccessPath n_p, IRDocsifier d) -> Doc { return Relax(d, "Any"); }); } -ExprDoc PrintShapeVar(const PrimExpr& e, const AccessPath& e_p, const IRDocsifier& d) { - ExprDoc expr_doc = d->AsDoc(e, e_p); - // Step 1. Find if `func_vars` are being collected - const RelaxFrameNode* f = nullptr; - for (const Frame& frame : d->frames) { - if (const auto* relax_frame = frame.as()) { - if (relax_frame->func_vars) { - f = relax_frame; - break; - } - } - } - // Step 2. Figure out if the PrimExpr contains at least a func var - bool func_var_mode = false; - if (f != nullptr) { - auto walk_fn = [f, &func_var_mode](const tirx::Var& var) -> ffi::Expected { - if (auto prim_var = var.as()) { - if (f->prim_params->count(prim_var.value().get())) { - func_var_mode = true; - } - } - return ffi::WalkResult::Advance(); - }; - ffi::StructuralWalk(e, walk_fn); - } - // Step 3. Stringify the PrimExpr if func var exists - bool is_bare_type_var = false; - if (f != nullptr && f->type_vars != nullptr) { - if (auto var = e.as()) { - is_bare_type_var = f->type_vars->count(var.value().get()); - } - } - bool use_postponed_annotations = - UsePEP695TypeVars(d) && f != nullptr && f->type_vars != nullptr && !f->type_vars->empty(); - if (func_var_mode && !is_bare_type_var && !use_postponed_annotations) { - return ExprStringDoc(expr_doc, e_p); - } - return expr_doc; -} - TVM_FFI_STATIC_INIT_BLOCK() { IRDocsifier::vtable().set_dispatch( "", [](relax::ShapeType n, AccessPath n_p, IRDocsifier d) -> Doc { @@ -80,7 +38,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { AccessPath shape_p = n_p->Attr("values"); ffi::Array shape_docs; for (int i = 0, ndim = shape.size(); i < ndim; ++i) { - shape_docs.push_back(PrintShapeVar(shape[i], shape_p->ArrayItem(i), d)); + shape_docs.push_back(d->AsDoc(shape[i], shape_p->ArrayItem(i))); } return Relax(d, "Shape")->Call({ListDoc(shape_docs)}); } @@ -101,7 +59,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { AccessPath shape_p = n_p->Attr("shape")->Attr("values"); ffi::Array shape_docs; for (int i = 0, ndim = shape_expr->values.size(); i < ndim; ++i) { - shape_docs.push_back(PrintShapeVar(shape_expr->values[i], shape_p->ArrayItem(i), d)); + shape_docs.push_back(d->AsDoc(shape_expr->values[i], shape_p->ArrayItem(i))); } args.push_back(TupleDoc(shape_docs)); } else { diff --git a/src/relax/script/printer/distributed.cc b/src/relax/script/printer/distributed.cc index b7327bedfed3..cfc8dc59e833 100644 --- a/src/relax/script/printer/distributed.cc +++ b/src/relax/script/printer/distributed.cc @@ -49,7 +49,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { AccessPath shape_p = n_p->Attr("shape")->Attr("values"); ffi::Array shape_docs; for (int i = 0, ndim = shape_expr->values.size(); i < ndim; ++i) { - shape_docs.push_back(PrintShapeVar(shape_expr->values[i], shape_p->ArrayItem(i), d)); + shape_docs.push_back(d->AsDoc(shape_expr->values[i], shape_p->ArrayItem(i))); } args.push_back(TupleDoc(shape_docs)); } else { diff --git a/src/relax/script/printer/expr.cc b/src/relax/script/printer/expr.cc index a38290238b89..d30e93223480 100644 --- a/src/relax/script/printer/expr.cc +++ b/src/relax/script/printer/expr.cc @@ -66,7 +66,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { ffi::Array values_doc; AccessPath values_p = n_p->Attr("values"); for (int i = 0, l = n->values.size(); i < l; ++i) { - values_doc.push_back(PrintShapeVar(n->values[i], values_p->ArrayItem(i), d)); + values_doc.push_back(d->AsDoc(n->values[i], values_p->ArrayItem(i))); } return Relax(d, "shape")->Call({ListDoc(values_doc)}); }); diff --git a/src/relax/script/printer/function.cc b/src/relax/script/printer/function.cc index 47e249ed56e1..bd27d16ba550 100644 --- a/src/relax/script/printer/function.cc +++ b/src/relax/script/printer/function.cc @@ -76,6 +76,9 @@ TVM_FFI_STATIC_INIT_BLOCK() { prim_params.insert(param.get()); } } + if (!prim_params.empty()) { + d->ir_usage.insert("future_annotations"); + } // Step 1. Print params ffi::Array params; { diff --git a/src/relax/script/printer/utils.h b/src/relax/script/printer/utils.h index c83e15fe2cbf..a9db62fd7e0b 100644 --- a/src/relax/script/printer/utils.h +++ b/src/relax/script/printer/utils.h @@ -141,8 +141,6 @@ ffi::Array PrintSeqExpr(const relax::SeqExpr& n, const AccessPath& n_p, Doc PrintRelaxVar(tvm::Var n, AccessPath p, IRDocsifier d); -ExprDoc PrintShapeVar(const PrimExpr& e, const AccessPath& e_p, const IRDocsifier& d); - inline int FindVDeviceIndexByTargetKind(const relax::VDevice& vdevice, const IRDocsifier& d) { ffi::Array vdevices = d->global_infos["vdevice"]; int kind_index = 0; diff --git a/src/script/printer/ir/ir.cc b/src/script/printer/ir/ir.cc index b260d0f3afa9..f697bd9aaf0f 100644 --- a/src/script/printer/ir/ir.cc +++ b/src/script/printer/ir/ir.cc @@ -85,8 +85,8 @@ TVM_FFI_STATIC_INIT_BLOCK() { IRDocsifier::vtable().set_dispatch( "", [](IRModule mod, AccessPath p, IRDocsifier d) -> Doc { std::vector functions; - ffi::Array type_var_decls; - std::unordered_set declared_type_vars; + ffi::Array dynamic_decls; + std::unordered_set declared_dynamic_names; for (const auto& kv : mod->functions) { functions.push_back(SortableFunction(kv)); } @@ -124,9 +124,9 @@ TVM_FFI_STATIC_INIT_BLOCK() { if (const auto* stmt_block = doc.as()) { for (const StmtDoc& stmt : stmt_block->stmts) { if (ffi::Optional name = GetDynamicDeclarationName(stmt)) { - if (!declared_type_vars.count(name.value())) { - declared_type_vars.insert(name.value()); - type_var_decls.push_back(stmt); + if (!declared_dynamic_names.count(name.value())) { + declared_dynamic_names.insert(name.value()); + dynamic_decls.push_back(stmt); } } } @@ -148,11 +148,11 @@ TVM_FFI_STATIC_INIT_BLOCK() { } } ClassDoc class_doc(module_doc, {IR(d, "ir_module")}, (*f)->stmts); - if (type_var_decls.empty()) { + if (dynamic_decls.empty()) { return HeaderWrapper(d, class_doc); } - type_var_decls.push_back(class_doc); - return HeaderWrapper(d, StmtBlockDoc(type_var_decls)); + dynamic_decls.push_back(class_doc); + return HeaderWrapper(d, StmtBlockDoc(dynamic_decls)); }); } diff --git a/src/tirx/script/printer/buffer.cc b/src/tirx/script/printer/buffer.cc index b943631a9fe3..6f75ab082927 100644 --- a/src/tirx/script/printer/buffer.cc +++ b/src/tirx/script/printer/buffer.cc @@ -29,11 +29,10 @@ namespace script { namespace printer { -ffi::Map BufferAttrs( - tirx::BufferVar buffer, const AccessPath& buffer_p, const Frame& frame, const IRDocsifier& d, - BufferVarDefinition var_definitions, ffi::Optional data = std::nullopt, - bool stringify_undefined_shape = false, - std::unordered_set stringify_shape_vars = {}) { +ffi::Map BufferAttrs(tirx::BufferVar buffer, const AccessPath& buffer_p, + const Frame& frame, const IRDocsifier& d, + BufferVarDefinition var_definitions, + ffi::Optional data = std::nullopt) { using tvm::tirx::Var; using tvm::tirx::VarNode; ffi::Map kwargs; @@ -63,21 +62,6 @@ ffi::Map BufferAttrs( ffi::StructuralWalk(data.value(), count_data_var); } auto is_new_var = [&](const Expr& e) { return e->IsInstance() && !d->IsVarDefined(e); }; - // All expression-string annotation fields use the same Python-binding rule. - // Dynamic symbols are real Python bindings; expressions referring to later - // scalar parameters must still be quoted. - auto expression_doc = [&](const PrimExpr& e, const AccessPath& e_p, - bool was_undefined = false) -> ExprDoc { - bool needs_quote = stringify_undefined_shape && was_undefined; - auto walk_fn = [&](const Var& var) -> ffi::Expected { - needs_quote = needs_quote || (stringify_undefined_shape && - (!d->IsVarDefined(var) || stringify_shape_vars.count(var))); - return ffi::WalkResult::Advance(); - }; - ffi::StructuralWalk(e, walk_fn); - ExprDoc result = d->AsDoc(e, e_p); - return needs_quote ? ExprDoc(ExprStringDoc(result, e_p)) : result; - }; auto add_out_of_line_var_def = [&](const Var& var, const AccessPath& var_p) { TVM_FFI_ICHECK(!d->IsVarDefined(var)); ExprDoc lhs = DefineVar(var, frame, d); @@ -111,7 +95,7 @@ ffi::Map BufferAttrs( if (was_undefined) { add_out_of_line_var_def(e.as_or_throw(), e_p); } - results.push_back(expression_doc(e, e_p, was_undefined)); + results.push_back(d->AsDoc(e, e_p)); } kwargs.Set("shape", TupleDoc(results)); } @@ -152,19 +136,9 @@ ffi::Map BufferAttrs( PrimExpr e = strides[i]; AccessPath e_p = strides_p->ArrayItem(i); if (is_new_var(e)) { - // String stride declarations have int64 dtype. - PrimType stride_ty = e.ty(); - if (!stride_ty.IsScalar() || !stride_ty.MatchesElementType(DLDataTypeCode::kDLInt, 64)) { - add_out_of_line_var_def(e.as_or_throw(), e_p); - } else if (try_inline_def(e, e_p, [=]() { - return d->AsDoc(buffer, buffer_p) - ->Attr("strides")[{LiteralDoc::Int(i, std::nullopt)}]; - })) { - results.push_back(LiteralDoc::Str(e.as_or_throw()->name, e_p)); - continue; - } + add_out_of_line_var_def(e.as_or_throw(), e_p); } - results.push_back(expression_doc(e, e_p)); + results.push_back(d->AsDoc(e, e_p)); } kwargs.Set("strides", TupleDoc(results)); } @@ -173,14 +147,16 @@ ffi::Map BufferAttrs( if (const auto* int_imm = buffer->elem_offset.as()) { if (int_imm->value != 0 || int_imm->ty.as_or_throw()->dtype != buffer->DefaultIndexType()) { - kwargs.Set("elem_offset", expression_doc(buffer->elem_offset, buffer_p->Attr("elem_offset"))); + kwargs.Set("elem_offset", + d->AsDoc(buffer->elem_offset, buffer_p->Attr("elem_offset"))); } } else if (is_new_var(buffer->elem_offset)) { try_inline_def(buffer->elem_offset, buffer_p->Attr("elem_offset"), [=]() { return d->AsDoc(buffer, buffer_p)->Attr("elem_offset"); }); needs_print_factor = true; } else { - kwargs.Set("elem_offset", expression_doc(buffer->elem_offset, buffer_p->Attr("elem_offset"))); + kwargs.Set("elem_offset", + d->AsDoc(buffer->elem_offset, buffer_p->Attr("elem_offset"))); } // Step 6. Handle `buffer.scope` { @@ -230,13 +206,14 @@ ffi::Map BufferAttrs( // Unwrap single-element array: DeclBuffer expects Optional, not Array. // Use the normal expression printer so a bound scalar alias stays a scalar // load, while an ordinary buffer load retains its indices. - kwargs.Set("allocated_addr", expression_doc(buffer->allocated_addr[0], - buffer_p->Attr("allocated_addr")->ArrayItem(0))); + kwargs.Set("allocated_addr", + d->AsDoc(buffer->allocated_addr[0], + buffer_p->Attr("allocated_addr")->ArrayItem(0))); } else { ffi::Array addresses; for (size_t i = 0; i < buffer->allocated_addr.size(); ++i) { - addresses.push_back(expression_doc(buffer->allocated_addr[i], - buffer_p->Attr("allocated_addr")->ArrayItem(i))); + addresses.push_back(d->AsDoc(buffer->allocated_addr[i], + buffer_p->Attr("allocated_addr")->ArrayItem(i))); } kwargs.Set("allocated_addr", TupleDoc(addresses)); } @@ -322,10 +299,9 @@ ExprDoc BufferDecl(const tirx::BufferVar& buffer, const ffi::String& method, } ExprDoc BufferAttn(const tirx::BufferVar& buffer, const AccessPath& p, const Frame& frame, - const IRDocsifier& d, std::unordered_set stringify_shape_vars) { + const IRDocsifier& d) { ffi::Map attrs = - BufferAttrs(buffer, p, frame, d, BufferVarDefinition::MatchBuffer, std::nullopt, true, - std::move(stringify_shape_vars)); + BufferAttrs(buffer, p, frame, d, BufferVarDefinition::MatchBuffer); if (!attrs.count("dtype")) { attrs.Set("dtype", LiteralDoc::DataType(buffer->dtype->dtype, p->Attr("dtype"))); } diff --git a/src/tirx/script/printer/function.cc b/src/tirx/script/printer/function.cc index 59ed0ce2625c..12a64e5f2e99 100644 --- a/src/tirx/script/printer/function.cc +++ b/src/tirx/script/printer/function.cc @@ -78,9 +78,9 @@ TVM_FFI_STATIC_INIT_BLOCK() { std::unordered_map scalar_param_docs; // Define scalar docs up front so a preceding Buffer parameter can render // a reference to a later scalar parameter. `bound_signature_vars` - // separately tracks Python bindings in source order. Quoted shapes - // resolve native symbols without binding their names in Python. + // separately tracks Python bindings in source order. std::unordered_set bound_signature_vars; + bool has_dependent_annotations = false; for (const tirx::Var& param : func->params) { if (!param->ty.as()) { scalar_param_docs.emplace(param.get(), DefineVar(param, *f, d)); @@ -97,7 +97,10 @@ TVM_FFI_STATIC_INIT_BLOCK() { bool needs_body_declaration = false; auto check_annotation_var = [&](const tirx::Var& annotation_var) -> ffi::Expected { - if (!bound_signature_vars.count(annotation_var)) { + has_dependent_annotations = + has_dependent_annotations || runtime_params.count(annotation_var.get()); + if (!bound_signature_vars.count(annotation_var) && + !type_vars.count(annotation_var.get())) { needs_body_declaration = true; } return ffi::WalkResult::Advance(); @@ -107,6 +110,17 @@ TVM_FFI_STATIC_INIT_BLOCK() { tirx::TileLayoutNode::DefaultLayout(buffer->shape))) { ffi::StructuralWalk(buffer->layout, check_annotation_var); } + for (const PrimExpr& extent : buffer->shape) { + ffi::StructuralWalk(extent, check_annotation_var); + } + for (const PrimExpr& stride : buffer->strides) { + ffi::StructuralWalk(stride, check_annotation_var); + } + ffi::StructuralWalk(buffer->elem_offset, + check_annotation_var); + for (const PrimExpr& address : buffer->allocated_addr) { + ffi::StructuralWalk(address, check_annotation_var); + } if (needs_body_declaration) { tirx::Var handle(var->name + "_handle", PointerType::VoidPointerTy()); ExprDoc handle_doc = DefineVar(handle, *f, d); @@ -117,27 +131,8 @@ TVM_FFI_STATIC_INIT_BLOCK() { (*f)->stmts.push_back(AssignDoc(lhs, rhs, std::nullopt)); continue; } - std::unordered_set stringify_shape_vars; - auto walk_fn = [&](const tirx::Var& shape_var) -> ffi::Expected { - bool is_type_var = type_vars.count(shape_var.get()); - if (!bound_signature_vars.count(shape_var) && !is_type_var) { - stringify_shape_vars.insert(shape_var); - } - return ffi::WalkResult::Advance(); - }; - for (const PrimExpr& shape : buffer->shape) { - ffi::StructuralWalk(shape, walk_fn); - } - for (const PrimExpr& stride : buffer->strides) { - ffi::StructuralWalk(stride, walk_fn); - } - ffi::StructuralWalk(buffer->elem_offset, walk_fn); - for (const PrimExpr& address : buffer->allocated_addr) { - ffi::StructuralWalk(address, walk_fn); - } IdDoc lhs = DefineBuffer(buffer, *f, d); - ExprDoc annotation = - BufferAttn(buffer, var_p->Attr("ty"), *f, d, std::move(stringify_shape_vars)); + ExprDoc annotation = BufferAttn(buffer, var_p->Attr("ty"), *f, d); args.push_back(AssignDoc(lhs, std::nullopt, annotation)); continue; } @@ -145,6 +140,9 @@ TVM_FFI_STATIC_INIT_BLOCK() { args.push_back(AssignDoc(scalar_param_docs.at(var.get()), std::nullopt, a)); bound_signature_vars.insert(var); } + if (has_dependent_annotations) { + d->ir_usage.insert("future_annotations"); + } ffi::Optional ret_type = std::nullopt; if (!func->ret_type.IsMissing()) { const auto* as_tuple = func->ret_type.as(); diff --git a/src/tirx/script/printer/utils.h b/src/tirx/script/printer/utils.h index 919ddc6ba4c7..d6ae5a222ef9 100644 --- a/src/tirx/script/printer/utils.h +++ b/src/tirx/script/printer/utils.h @@ -322,12 +322,10 @@ ExprDoc BufferDecl(const tirx::BufferVar& buffer, const ffi::String& method, * \param p The object path * \param f The frame * \param d The IRDocsifier - * \param stringify_shape_vars Variables without a Python binding at this annotation. Every - * shape expression containing one of these variables must be stringified. * \return The ExprDoc corresponding to the buffer declaration */ ExprDoc BufferAttn(const tirx::BufferVar& buffer, const AccessPath& p, const Frame& frame, - const IRDocsifier& d, std::unordered_set stringify_shape_vars = {}); + const IRDocsifier& d); /*! * \brief Print the creation of a Var diff --git a/tests/python/relax/test_tvmscript_parser.py b/tests/python/relax/test_tvmscript_parser.py index 6e72a575b2c8..215608a2ce3f 100644 --- a/tests/python/relax/test_tvmscript_parser.py +++ b/tests/python/relax/test_tvmscript_parser.py @@ -145,7 +145,7 @@ def foo(x: R.Tensor((m, m), "float32")): m = T.dynamic("m", "int64") @R.function - def f(x: R.Tensor((m, "n"), "float32")): + def f(x: R.Tensor((m, m), "float32")): # call_tir expected a tirx prim_func return relax.call_tir("extern_func", (x,), R.Tensor((T.max(m),), dtype="float32")) @@ -1529,6 +1529,8 @@ def foo(x: R.Tensor, _: R.Shape([m, n])): def test_erase_to_well_defined_infers_from_shape_expr(): m = T.dynamic("m", "int64") n = T.dynamic("n", "int64") + main_m = T.dynamic("m", "int64") + main_n = T.dynamic("n", "int64") @I.ir_module class Module: @@ -1545,7 +1547,7 @@ def subroutine(x: R.Tensor, _: R.Shape([m, n])) -> R.Tensor([m, n]): # subroutine. Therefore, the shape of the tensor returned # from main can have a well-defined shape. @R.function - def main(x: R.Tensor, shape: R.Shape(["m", "n"])): + def main(x: R.Tensor, shape: R.Shape([main_m, main_n])): output = Module.subroutine(x, shape) return output @@ -1569,10 +1571,12 @@ def foo(x: R.Tuple()): def test_symbolic_vars_in_tensor_shape_with_usage_first(): - """First param may use symbolic variable defined in second param""" + """A captured symbol can first appear inside a compound dimension.""" + + m = T.dynamic("m") @R.function - def foo(x: R.Tensor(("m + 1",), "float32"), y: R.Tensor(("m", 1), "float32")): + def foo(x: R.Tensor((m + 1,), "float32"), y: R.Tensor((m, 1), "float32")): z = R.add(x, y) return z @@ -1588,13 +1592,13 @@ def foo(x: R.Tensor(("m + 1",), "float32"), y: R.Tensor(("m", 1), "float32")): def test_symbolic_vars_in_tensor_shape_with_definition_first(): - """Second param may use symbolic variable defined in first param""" + """A captured symbol is shared across direct and compound dimensions.""" m = T.dynamic("m", "int64") @R.function - def bar(x: R.Tensor((m,), "float32"), y: R.Tensor(("T.max(m, 20)",), "float32")) -> R.Tensor( - ("T.max(m, 20) + 1",), "float32" + def bar(x: R.Tensor((m,), "float32"), y: R.Tensor((T.max(m, 20),), "float32")) -> R.Tensor( + (T.max(m, 20) + 1,), "float32" ): z = R.call_dps_packed("test_intrin", (x, y), R.Tensor((T.max(m, 20) + 1,), dtype="float32")) return z @@ -1621,17 +1625,17 @@ def test_bound_prim_param_reused_in_dependent_annotations(): def main( n: T.int64, direct: R.Tensor([n], "float32"), - string_direct: R.Tensor(["n"], "float32"), - shape: R.Shape(["n"]), - compound: R.Tensor(["n + 1"], "float32"), -) -> R.Tensor(["n + 1"], "float32"): + repeated: R.Tensor([n], "float32"), + shape: R.Shape([n]), + compound: R.Tensor([n + 1], "float32"), +) -> R.Tensor([n + 1], "float32"): return compound """ ) - n, direct, string_direct, shape, compound = func.params + n, direct, repeated, shape, compound = func.params assert direct.ty.shape[0].same_as(n) - assert string_direct.ty.shape[0].same_as(n) + assert repeated.ty.shape[0].same_as(n) assert shape.ty.values[0].same_as(n) assert compound.ty.shape[0].a.same_as(n) assert func.ret_ty.shape[0].a.same_as(n) @@ -1644,8 +1648,8 @@ def test_bound_prim_param_reused_in_declared_function_signature(): @I.ir_module class Module: @R.function - def main(n: T.int64, x: R.Tensor(["n + 1"], "float32")) -> R.Tensor( - ["n + 1"], "float32" + def main(n: T.int64, x: R.Tensor([n + 1], "float32")) -> R.Tensor( + [n + 1], "float32" ): return x """ @@ -1658,12 +1662,12 @@ def main(n: T.int64, x: R.Tensor(["n + 1"], "float32")) -> R.Tensor( _check(mod) -def test_later_prim_param_not_adopted_by_usage_first_symbol(): - with pytest.raises(ValueError): +def test_later_prim_param_requires_external_shape_symbol(): + with pytest.raises(NameError): tvm.script.from_source( """ @R.function -def main(x: R.Tensor(["n"], "float32"), n: T.int64): +def main(x: R.Tensor([n], "float32"), n: T.int64): return x """ ) @@ -1674,7 +1678,7 @@ def test_non_int64_prim_param_rejected_in_shape_annotation(): tvm.script.from_source( """ @R.function -def main(n: T.int32, x: R.Tensor(["n"], "float32")): +def main(n: T.int32, x: R.Tensor([n], "float32")): return x """ ) @@ -1729,26 +1733,21 @@ def baz(x: R.Shape((m,)), y: R.Tensor((m * 2,), "float32")): _check(baz, bb.get()["baz"]) -def test_undefined_symbolic_var_raises_error(): - """An undefined symbolic variable in an error - - A symbolic variables is defined at the first site where it appears - as a shape parameter without any modification. TVMScript does not - support solving for a symbolic variable in terms of the argument - shape. That is, this test case raises an error, and will not - attempt to define `m` as either `x.shape[0]-1` or `x.shape[1]//2`. - """ - with pytest.raises(ValueError): +def test_string_shape_expression_is_not_resolved(): + """A quoted expression is ordinary data, not a symbol declaration.""" + with pytest.raises(TypeError, match="Array"): @R.function - def foo(x: R.Tensor(("m + 1", "m * 2"), "float32")): # name 'm' is not defined - z = R.add(x, x) - return z + def foo(x: R.Tensor(("m + 1", "m * 2"), "float32")): + return x def test_arith_operators(): + m = T.dynamic("m") + n = T.dynamic("n") + @R.function - def foo(x: R.Tensor(("m", "n"), "float32"), y: R.Tensor(("m", "n"), "float32")): + def foo(x: R.Tensor((m, n), "float32"), y: R.Tensor((m, n), "float32")): a0 = -x a1 = x + y a2 = x - y @@ -1830,8 +1829,11 @@ def foo(x: R.Tensor((m, n), dtype="float32")): def test_builtin_ops(): + m = T.dynamic("m") + n = T.dynamic("n") + @R.function - def foo(x: R.Tensor(("m", "n"), dtype="float32")): + def foo(x: R.Tensor((m, n), dtype="float32")): tensor = R.builtin.stop_lift_params(x) gv = tensor return gv @@ -2401,11 +2403,13 @@ def test_function_attributes_are_defined(): m = T.dynamic("m", "int64") n = T.dynamic("n", "int64") + main_m = T.dynamic("m", "int64") + main_n = T.dynamic("n", "int64") @I.ir_module class Module: @R.function - def main(x: R.Tensor, shape: R.Shape(["m", "n"])): + def main(x: R.Tensor, shape: R.Shape([main_m, main_n])): output = Module.subroutine(x, shape) return output @@ -2449,8 +2453,10 @@ def expected(A: R.Tensor([extent])) -> R.Tensor([extent - 1]): def test_non_declaration_prim_expr_emits_binding(): """Dtype casts emit ordinary bindings without replacing shape symbols.""" + symbol = T.dynamic("extent") + @R.function(private=True) - def func(A: R.Tensor(["extent"], "float32")): + def func(A: R.Tensor([symbol], "float32")): extent = T.int64(4) output = A return output diff --git a/tests/python/relax/test_tvmscript_printer_relax.py b/tests/python/relax/test_tvmscript_printer_relax.py index 7866ddb6e8c2..a86682250474 100644 --- a/tests/python/relax/test_tvmscript_printer_relax.py +++ b/tests/python/relax/test_tvmscript_printer_relax.py @@ -56,7 +56,7 @@ def func(a: R.Tensor((10, 10))) -> R.Tensor((10, 10)): ) -def test_function_dependent_shape_escaped_source_spans(): +def test_function_dependent_shape_source_spans(): n = tirx.Var("n", "int64") cast = tirx.Cast("int64", n) x = relax.Var("x", relax.TensorType([cast], "float32")) @@ -86,17 +86,17 @@ def render(path): assert "Access path:" not in lines[definition_index] return lines[definition_index], lines[definition_index + 1] - expression = r'"T.Cast(\"int64\", n)"' + expression = 'T.Cast("int64", n)' definition, underline = render(cast_path) expression_start = definition.index(expression) assert underline[expression_start : expression_start + len(expression)] == "^" * len(expression) assert underline.strip() == "^" * len(expression) - escaped_dtype = r"\"int64\"" + dtype_literal = '"int64"' definition, underline = render(cast_path.attr("dtype")) - dtype_start = definition.index(escaped_dtype) - assert underline[dtype_start : dtype_start + len(escaped_dtype)] == "^" * len(escaped_dtype) - assert underline.strip() == "^" * len(escaped_dtype) + dtype_start = definition.index(dtype_literal) + assert underline[dtype_start : dtype_start + len(dtype_literal)] == "^" * len(dtype_literal) + assert underline.strip() == "^" * len(dtype_literal) definition, underline = render(cast_path.attr("value")) variable_start = definition.index(expression) + expression.rindex("n") diff --git a/tests/python/relax/test_tvmscript_type_vars.py b/tests/python/relax/test_tvmscript_type_vars.py index e6dbafa3c0a5..00499acb192b 100644 --- a/tests/python/relax/test_tvmscript_type_vars.py +++ b/tests/python/relax/test_tvmscript_type_vars.py @@ -29,8 +29,8 @@ def test_type_vars_roundtrip(): @R.function(private=True) def func( - x: R.Tensor((M, "M * 2"), "float32"), - ) -> R.Tensor((M, "M * 2"), "float32"): + x: R.Tensor((M, M * 2), "float32"), + ) -> R.Tensor((M, M * 2), "float32"): return x script = func.script() diff --git a/tests/python/script/minilang.py b/tests/python/script/minilang.py index 2de387129b83..25905fdf6c8f 100644 --- a/tests/python/script/minilang.py +++ b/tests/python/script/minilang.py @@ -218,7 +218,6 @@ def operation(*args, name=name): ) @resolve_global_info_args("device", resolver=self.resolve_global_info) - @registry.args_policy("M.Tensor", {"shape": "expr_str"}, scalar_strings=False) def Tensor(shape=None, dtype="float32", device=None, placement="S[0]"): return Value("tensor", (shape, dtype, device, placement)) diff --git a/tests/python/script/test_basic_usage.py b/tests/python/script/test_basic_usage.py index 1b088e84af5a..6d05987a2cdb 100644 --- a/tests/python/script/test_basic_usage.py +++ b/tests/python/script/test_basic_usage.py @@ -179,6 +179,7 @@ def test_nested_policy_and_starred_calls_are_evaluated_once(language): mesh = object() language.global_infos["mesh[0]"] = mesh values = (1, 2) + n = M.dynamic("n") def outer(value): seen.append(value) @@ -189,7 +190,7 @@ def collect(*values, other): @M.function def main(): - outer(M.Tensor(("n",), device="mesh[0]")) + outer(M.Tensor((n,), device="mesh[0]")) collect(*values, other=3) assert len(seen) == 1 and seen[0].args[2] is mesh diff --git a/tests/python/script/test_meta_programming.py b/tests/python/script/test_meta_programming.py index c5a30e282cbe..17f82ea0454a 100644 --- a/tests/python/script/test_meta_programming.py +++ b/tests/python/script/test_meta_programming.py @@ -31,7 +31,6 @@ import pytest -from tvm import ir from tvm.script import ir as I from tvm.script.parser import entry, protocol_registry @@ -331,23 +330,25 @@ def main(x: M.Tensor((4,))): def _build_symbolic_functions(M): + first_n = M.dynamic("n") + second_n = M.dynamic("n") + @I.ir_module class Module: @M.function - def first(x: M.Tensor(("n",), "float32")): + def first(x: M.Tensor((first_n,), "float32")): return x @M.function - def second(x: M.Tensor(("n",), "float32")): + def second(x: M.Tensor((second_n,), "float32")): return x return Module def test_dynamic_caller_symbol_does_not_join_function_declarations(language): - # Before: two function signatures declare "n" while a caller has its own n. - # Expected builder: each function frame resolves its own n, independent of the caller. - n = ir.Var("n", "int64") + # Two independently constructed symbols remain distinct from an unrelated caller capture. + n = I.dynamic("n") module = _build_symbolic_functions(language.M) first = module["first"].params[0].args[0].args[0][0] second = module["second"].params[0].args[0].args[0][0] diff --git a/tests/python/script/test_parser_entry.py b/tests/python/script/test_parser_entry.py index 3d69be916767..47256a16546a 100644 --- a/tests/python/script/test_parser_entry.py +++ b/tests/python/script/test_parser_entry.py @@ -32,7 +32,7 @@ def test_parse_string_returns_fresh_symbols(language): # The direct string API must create independent symbol/parameter objects on repeated calls. # This is the suite's single dedicated parse(str) API case. - source = '@M.function\ndef main(x: M.Tensor(("n + 1",))):\n M.record(x)\n' + source = '@M.function\ndef main(x: M.Tensor((M.dynamic("n") + 1,))):\n M.record(x)\n' first = entry.parse(source, extra_vars={"M": language.M}, root_builder=language.M) second = entry.parse(source, extra_vars={"M": language.M}, root_builder=language.M) assert first.params[0] is not second.params[0] diff --git a/tests/python/script/test_special_parser_protocol.py b/tests/python/script/test_special_parser_protocol.py index 1c1da9d8ef89..d85d11df81ea 100644 --- a/tests/python/script/test_special_parser_protocol.py +++ b/tests/python/script/test_special_parser_protocol.py @@ -67,19 +67,19 @@ def main(): assert calls == [first, first, invalid, second, second] -def test_argument_policy_preserves_expression_dtype(language): - # Shape-expression parsing must retain the declared int32 symbol dtype. +def test_external_expression_preserves_symbol_dtype(language): + # Ordinary expressions retain the externally declared int32 symbol dtype. M = language.M - @registry.args_policy("M.shape", {"values": "expr_str"}, dtype="int32") def shape(values): return values M.shape = shape + n = M.dynamic("n", "int32") @M.function def main(): - M.shape(("n", "n + 1")) + M.shape((n, n + 1)) n, increment = main.body[0][1] assert n.args == ("int32",) diff --git a/tests/python/script/test_symbolic_shape.py b/tests/python/script/test_symbolic_shape.py index e90b8598aa52..4e6a5ae6b584 100644 --- a/tests/python/script/test_symbolic_shape.py +++ b/tests/python/script/test_symbolic_shape.py @@ -15,7 +15,7 @@ # specific language governing permissions and limitations # under the License. """Explicit symbolic dimensions retain identity in signatures and function bodies. -Quoted dimensions must not create or replace ordinary Python bindings. +Dimensions use ordinary Python expressions over explicit symbols. """ from __future__ import annotations @@ -54,22 +54,23 @@ def main( def test_signature_read_before_introduction_remains_unbound(language): - # A quoted dimension must not introduce an ordinary Python name before its declaration. + # An undeclared dimension must remain an ordinary unbound Python name. M = language.M with pytest.raises(NameError, match="n"): @M.function - def main(x: M.Tensor((n, "n"), "float32")): + def main(x: M.Tensor((n,), "float32")): return x -def test_signature_strings_do_not_replace_captured_python_names(language): - # A captured Python dimension and a same-spelling quoted symbol must remain distinct. +def test_same_named_symbol_does_not_replace_captured_python_value(language): + # A captured Python dimension and a same-spelling external symbol remain distinct. M = language.M n = 7 + symbol = M.dynamic("n") @M.function - def main(x: M.Tensor((n, "n"), "float32"), y: M.Tensor((n,), "float32")): + def main(x: M.Tensor((n, symbol), "float32"), y: M.Tensor((n,), "float32")): M.record(n) return y @@ -81,8 +82,7 @@ def main(x: M.Tensor((n, "n"), "float32"), y: M.Tensor((n,), "float32")): def test_captured_shape_requires_concrete_symbols(): - # Captured tuples bypass literal decoding; native shape construction must - # preserve concrete symbols and reject captured strings instead of inventing vars. + # Native shape construction preserves concrete symbols and rejects strings. def build(shape): @T.prim_func def main(x: T.Buffer(shape, "float32")): @@ -93,11 +93,8 @@ def main(x: T.Buffer(shape, "float32")): n = T.dynamic("n") function = build((n, 16)) assert function.params[0].ty.shape[0].same_as(n) - with pytest.raises( - TypeError, match="^Builder expression arguments require concrete symbols, not strings$" - ) as caught: + with pytest.raises(AssertionError, match="data must be int or Expr, but got n"): build(("n", 16)) - assert type(caught.value) is TypeError def _line_of(function, statement): @@ -106,13 +103,14 @@ def _line_of(function, statement): def test_argument_policies_reuse_symbols_and_resolve_only_marked_literals(language): - # Quoted shapes must reuse one symbol while only marked device strings resolve. + # Shape expressions reuse the external symbol while marked device strings resolve. M = language.M device = object() language.global_infos["cuda:1"] = device + n = M.dynamic("n") @M.function - def main(x: M.Tensor(shape=("n + 1", "n"), dtype="float32", device="cuda:1")): + def main(x: M.Tensor(shape=(n + 1, n), dtype="float32", device="cuda:1")): M.record(x) annotation = main.params[0].args[0] @@ -124,13 +122,14 @@ def main(x: M.Tensor(shape=("n + 1", "n"), dtype="float32", device="cuda:1")): assert main.body == [("emit", main.params[0])] -def test_quoted_symbols_do_not_introduce_python_bindings(language): - # Resolving a quoted dimension must not silently define its unquoted body name. +def test_external_symbol_does_not_introduce_same_named_python_binding(language): + # Capturing a symbol under another name must not introduce its IR name in Python. M = language.M + symbol = M.dynamic("n") with pytest.raises(NameError, match="n"): @M.function - def main(x: M.Tensor(("n", "n"))): + def main(x: M.Tensor((symbol, symbol))): M.record(n) diff --git a/tests/python/tvmscript/test_tvmscript_error_report.py b/tests/python/tvmscript/test_tvmscript_error_report.py index bce0cde2af18..85bfe9231788 100644 --- a/tests/python/tvmscript/test_tvmscript_error_report.py +++ b/tests/python/tvmscript/test_tvmscript_error_report.py @@ -64,7 +64,7 @@ def test_buffer_bind(): def buffer_bind_missing_args(a: T.handle) -> None: A = T.match_buffer((16, 16), "float32") # error - check_error(buffer_bind_missing_args, 2, ValueError) + check_error(buffer_bind_missing_args, 2, TypeError) def test_undefined_buffer(): @@ -574,26 +574,26 @@ def non_integer_typed_block_iter(): def test_illegal_buffer_slice(): def strided_buffer_region(A: T.handle): # do not allow stride in buffer region - A = T.match_buffer((128, 128), "int32") - with Ts.sblock(): + A = T.match_buffer(A, (128, 128), "int32") + with Ts.sblock("block"): Ts.reads([]) Ts.writes([A[0:128:2, 0:128:3]]) # error T.evaluate(T.call_extern("strided_compute", dtype="")) def access_reversed_slice(A: T.handle): # do not allow reversed slice step - A = T.match_buffer((128,), "int32") + A = T.match_buffer(A, (128,), "int32") A[0:128:-1] = T.broadcast(1, 128) # error def access_non_const_slice_length(A: T.handle): # do not allow non-constant slice length - A = T.match_buffer((128,), "int32") + A = T.match_buffer(A, (128,), "int32") for i in range(4): T.evaluate(A[0:i:1]) # error - check_error(strided_buffer_region, 3, ValueError) - check_error(access_reversed_slice, 3, ValueError) - check_error(access_non_const_slice_length, 3, ValueError) + check_error(strided_buffer_region, 6, ValueError) + check_error(access_reversed_slice, 4, tvm.error.InternalError) + check_error(access_non_const_slice_length, 5, TypeError) def test_syntax_sugar_fail(): diff --git a/tests/python/tvmscript/test_tvmscript_parser_tir.py b/tests/python/tvmscript/test_tvmscript_parser_tir.py index e659abbc264e..44c8c24c7f35 100644 --- a/tests/python/tvmscript/test_tvmscript_parser_tir.py +++ b/tests/python/tvmscript/test_tvmscript_parser_tir.py @@ -48,16 +48,16 @@ def test_tir_bound_prim_param_reused_in_dependent_annotations(): def main( n: T.int32, direct: T.Buffer((n,), "float32"), - string_direct: T.Buffer(("n",), "float32"), - compound: T.Buffer(("n + 1",), "float32"), -) -> T.Buffer(("n",), "float32"): - return string_direct + repeated: T.Buffer((n,), "float32"), + compound: T.Buffer((n + 1,), "float32"), +) -> T.Buffer((n,), "float32"): + return repeated """ ) - n, direct, string_direct, compound = func.params + n, direct, repeated, compound = func.params assert direct.ty.shape[0].same_as(n) - assert string_direct.ty.shape[0].same_as(n) + assert repeated.ty.shape[0].same_as(n) assert compound.ty.shape[0].a.same_as(n) assert func.ret_type.shape[0].same_as(n) @@ -68,7 +68,7 @@ def test_tir_bound_prim_param_reused_in_declared_function_signature(): @I.ir_module class Module: @T.prim_func - def main(n: T.int32, A: T.Buffer(("n + 1",), "float32")): + def main(n: T.int32, A: T.Buffer((n + 1,), "float32")): T.evaluate(n) """ ) @@ -77,11 +77,12 @@ def main(n: T.int32, A: T.Buffer(("n + 1",), "float32")): assert A.ty.shape[0].a.same_as(n) -def test_tir_string_defined_symbol_adopted_by_later_prim_param(): +def test_tir_external_symbol_adopted_by_later_prim_param(): func = tvm.script.from_source( """ +n = T.dynamic("n", "int32") @T.prim_func -def main(A: T.Buffer(("n",), "float32"), n: T.int32): +def main(A: T.Buffer((n,), "float32"), n: T.int32): T.evaluate(n) """ ) @@ -92,10 +93,11 @@ def main(A: T.Buffer(("n",), "float32"), n: T.int32): mod = tvm.script.from_source( """ +n = T.dynamic("n", "int32") @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer(("n",), "float32"), n: T.int32): + def main(A: T.Buffer((n,), "float32"), n: T.int32): T.evaluate(n) """ ) @@ -105,11 +107,12 @@ def main(A: T.Buffer(("n",), "float32"), n: T.int32): assert str(n.ty.dtype) == "int32" -def test_tir_string_defined_symbol_preserves_later_prim_param_dtype(): +def test_tir_external_symbol_preserves_later_prim_param_dtype(): func = tvm.script.from_source( """ +n = T.dynamic("n", "int64") @T.prim_func -def main(A: T.Buffer(("n",), "float32"), n: T.int64): +def main(A: T.Buffer((n,), "float32"), n: T.int64): T.evaluate(n) """ ) @@ -134,12 +137,12 @@ def main(A: T.Buffer((n,), "float32")): assert func.body.value.same_as(n) -def test_tir_direct_use_before_string_definition_is_undefined(): +def test_tir_undeclared_shape_symbol_is_undefined(): with pytest.raises(NameError): tvm.script.from_source( """ @T.prim_func -def main(A: T.Buffer((n, "n"), "float32")): +def main(A: T.Buffer((n, n), "float32")): T.evaluate(0) """ ) @@ -168,11 +171,11 @@ def test_tir_direct_later_prim_param_is_undefined(source): def test_tir_return_annotation_does_not_define_symbolic_var(): - with pytest.raises(ValueError): + with pytest.raises(NameError): tvm.script.from_source( """ @T.prim_func -def main() -> T.Buffer(("n",), "float32"): +def main() -> T.Buffer((n,), "float32"): A = T.alloc_buffer((n,), "float32") return A """ From 6cb9bf206bd51a943d98135af11b3a1711c0a70b Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 21:01:27 +0000 Subject: [PATCH 03/16] [Script] Migrate symbolic dimensions to dynamic captures Declare symbolic dimensions and strides explicitly in script callers, intrinsic factories, and examples. Keep independent function symbols distinct and preserve shared captures and deferred whole annotations. --- .../tvm/relax/backend/gpu_generic/cumsum.py | 9 +- .../tvm/relax/backend/gpu_generic/sampling.py | 11 +- python/tvm/relax/block_builder.py | 8 +- .../relax/frontend/nn/llm/_decode_kernels.py | 50 +- .../relax/frontend/nn/llm/_page_kernels.py | 73 +- .../relax/frontend/nn/llm/_prefill_kernels.py | 107 +- .../frontend/nn/llm/position_embedding.py | 25 +- python/tvm/relax/frontend/nn/llm/tree_attn.py | 92 +- python/tvm/relax/frontend/nn/op.py | 14 +- python/tvm/s_tir/tensor_intrin/cuda.py | 74 +- python/tvm/s_tir/tensor_intrin/metal.py | 24 +- python/tvm/s_tir/tensor_intrin/rocm.py | 12 +- .../codegen/test_codegen_error_handling.py | 12 +- .../codegen/test_target_codegen_aarch64.py | 45 +- .../python/codegen/test_target_codegen_arm.py | 6 +- .../codegen/test_target_codegen_cuda.py | 9 +- .../codegen/test_target_codegen_device.py | 3 +- .../codegen/test_target_codegen_llvm.py | 30 +- .../test_target_codegen_static_init.py | 3 +- .../codegen/test_target_codegen_vulkan.py | 24 +- .../contrib/test_tir_triton_integration.py | 26 +- .../python/relax/backend/adreno/mod_utils.py | 36 +- .../test_transform_annotate_custom_scope.py | 19 +- ...istributed_transform_propagate_sharding.py | 230 ++-- tests/python/relax/test_analysis.py | 38 +- ...est_analysis_computable_at_compile_time.py | 29 +- ...test_analysis_suggest_layout_transforms.py | 18 +- .../relax/test_analysis_type_analysis.py | 9 +- .../python/relax/test_analysis_well_formed.py | 48 +- tests/python/relax/test_ast_printer.py | 47 +- .../relax/test_backend_dispatch_sampling.py | 12 +- .../relax/test_backend_dispatch_sort_scan.py | 57 +- .../test_backend_transform_shape_lower.py | 104 +- .../test_base_py_module_symbolic_shape.py | 49 +- tests/python/relax/test_bind_symbolic_vars.py | 95 +- tests/python/relax/test_blockbuilder_core.py | 66 +- .../python/relax/test_blockbuilder_emit_te.py | 9 +- tests/python/relax/test_codegen_cutlass.py | 59 +- tests/python/relax/test_contrib_vllm.py | 50 +- tests/python/relax/test_dataflow_inplace.py | 71 +- tests/python/relax/test_dataflow_pattern.py | 29 +- tests/python/relax/test_dataflow_rewriter.py | 73 +- tests/python/relax/test_e2e_op_dynamic.py | 9 +- .../test_frontend_from_exported_program.py | 97 +- .../python/relax/test_frontend_nn_exporter.py | 63 +- .../relax/test_frontend_nn_extern_module.py | 11 +- .../python/relax/test_frontend_nn_modules.py | 19 +- tests/python/relax/test_frontend_nn_op.py | 96 +- .../relax/test_frontend_nn_subroutines.py | 16 +- tests/python/relax/test_frontend_onnx.py | 294 ++-- tests/python/relax/test_frontend_stablehlo.py | 15 +- tests/python/relax/test_frontend_tflite.py | 23 +- tests/python/relax/test_inline_functions.py | 23 +- tests/python/relax/test_op_image.py | 2 +- tests/python/relax/test_op_index.py | 28 +- tests/python/relax/test_op_size.py | 6 +- tests/python/relax/test_op_take.py | 5 +- tests/python/relax/test_op_view.py | 47 +- .../relax/test_optimize_layout_transform.py | 10 +- tests/python/relax/test_pipeline.py | 9 +- .../python/relax/test_pytorch_integration.py | 9 +- tests/python/relax/test_relax_operators.py | 29 +- .../relax/test_relax_to_pyfunc_converter.py | 25 +- .../relax/test_runtime_builtin_rnn_state.py | 6 +- tests/python/relax/test_testing_nn.py | 18 +- .../relax/test_tir_call_source_kernel.py | 26 +- tests/python/relax/test_transform.py | 30 +- .../test_transform_adjust_matmul_order.py | 146 +- .../relax/test_transform_alter_op_impl.py | 5 +- .../test_transform_annotate_tir_op_pattern.py | 26 +- .../test_transform_attach_global_symbol.py | 58 +- .../relax/test_transform_bind_params.py | 22 +- .../test_transform_bind_symbolic_vars.py | 125 +- .../test_transform_bundle_model_params.py | 6 +- .../test_transform_canonicalize_bindings.py | 120 +- .../relax/test_transform_codegen_pass.py | 43 +- .../test_transform_combine_parallel_matmul.py | 46 +- .../test_transform_compute_prim_value.py | 20 +- .../relax/test_transform_convert_layout.py | 114 +- tests/python/relax/test_transform_cse.py | 18 +- .../test_transform_dead_code_elimination.py | 37 +- .../relax/test_transform_decompose_ops.py | 7 +- .../relax/test_transform_fold_constant.py | 30 +- tests/python/relax/test_transform_fuse_ops.py | 139 +- .../test_transform_fuse_ops_by_pattern.py | 48 +- tests/python/relax/test_transform_fuse_tir.py | 167 ++- tests/python/relax/test_transform_gradient.py | 10 +- .../test_transform_gradient_te_register.py | 51 +- .../test_transform_ipc_allreduce_rewrite.py | 37 +- .../relax/test_transform_lambda_lift.py | 51 +- .../test_transform_lazy_transform_params.py | 71 +- .../test_transform_legalize_ops_binary.py | 600 +++++---- ..._transform_legalize_ops_create_datatype.py | 265 ++-- ...test_transform_legalize_ops_distributed.py | 3 +- .../relax/test_transform_legalize_ops_grad.py | 48 +- .../test_transform_legalize_ops_image.py | 54 +- ...sform_legalize_ops_index_linear_algebra.py | 308 +++-- .../test_transform_legalize_ops_manipulate.py | 576 ++++---- .../relax/test_transform_legalize_ops_nn.py | 1195 +++++++++-------- .../relax/test_transform_legalize_ops_qdq.py | 74 +- ...ansform_legalize_ops_search_statistical.py | 365 +++-- .../test_transform_legalize_ops_unary.py | 10 +- .../test_transform_lift_transform_params.py | 158 ++- ...t_transform_lower_gpu_ipc_alloc_storage.py | 28 +- .../python/relax/test_transform_normalize.py | 14 +- ...test_transform_remove_unused_parameters.py | 113 +- ...est_transform_reorder_take_after_matmul.py | 99 +- .../test_transform_rewrite_cuda_graph.py | 84 +- ...test_transform_rewrite_dataflow_reshape.py | 24 +- ...test_transform_static_plan_block_memory.py | 153 ++- tests/python/relax/test_utils.py | 4 +- tests/python/relax/test_vm_build.py | 86 +- tests/python/relax/test_vm_builtin_lower.py | 18 +- tests/python/relax/test_vm_codegen_only.py | 8 +- .../python/runtime/test_runtime_extension.py | 3 +- tests/python/s_tir/dlight/test_benchmark.py | 37 +- tests/python/s_tir/dlight/test_cpu_gemv.py | 12 +- .../python/s_tir/dlight/test_gpu_fallback.py | 31 +- tests/python/s_tir/dlight/test_gpu_gemv.py | 15 +- .../dlight/test_gpu_general_reduction.py | 33 +- .../s_tir/dlight/test_gpu_low_batch_gemv.py | 36 +- tests/python/s_tir/dlight/test_gpu_matmul.py | 36 +- .../s_tir/dlight/test_gpu_matmul_tensorize.py | 212 ++- .../python/s_tir/dlight/test_gpu_reduction.py | 18 +- tests/python/s_tir/dlight/test_gpu_rmsnorm.py | 12 +- ...tproc_rewrite_parallel_vectorize_unroll.py | 3 +- ..._meta_schedule_postproc_verify_gpu_code.py | 13 +- .../test_meta_schedule_trace_apply.py | 31 +- .../schedule/test_tir_schedule_compute_at.py | 6 +- .../test_tir_schedule_compute_inline.py | 14 +- .../schedule/test_tir_schedule_pad_einsum.py | 18 +- .../schedule/test_tir_schedule_rfactor.py | 4 +- .../schedule/test_tir_schedule_sampling.py | 3 +- .../schedule/test_tir_schedule_split_fuse.py | 9 +- .../schedule/test_tir_schedule_tensorize.py | 28 +- .../test_tir_schedule_transform_layout.py | 12 +- .../python/s_tir/test_arith_domain_touched.py | 3 +- tests/python/s_tir/test_s_tir_renew_defs.py | 11 +- ...t_s_tir_transform_compact_buffer_region.py | 3 +- ...st_s_tir_transform_default_gpu_schedule.py | 9 +- ...tir_transform_force_narrow_index_to_i32.py | 3 +- .../test_s_tir_transform_hoist_if.py | 4 +- ...t_s_tir_transform_inject_ptx_async_copy.py | 99 +- ..._tir_transform_inject_software_pipeline.py | 14 +- ...est_s_tir_transform_lift_thread_binding.py | 6 +- ..._transform_lower_cross_thread_reduction.py | 10 +- ...test_s_tir_transform_lower_match_buffer.py | 38 +- ...test_s_tir_transform_lower_opaque_block.py | 8 +- ...tir_transform_memhammer_lower_auto_copy.py | 39 +- tests/python/te/test_te_create_primfunc.py | 37 +- .../test_tir_analysis_verify_well_formed.py | 22 +- tests/python/tirx-base/test_tir_intrin.py | 12 +- tests/python/tirx-base/test_tir_specialize.py | 26 +- .../test_tir_transform_convert_ssa.py | 32 +- ...tir_transform_force_narrow_index_to_i32.py | 6 +- .../test_tir_transform_make_packed_api.py | 3 +- .../test_tir_transform_simplify.py | 12 +- .../test_tir_transform_split_host_device.py | 3 +- .../test_tir_transform_vectorize.py | 3 +- tests/python/tirx/test_parser_printer.py | 4 +- .../transform/test_transform_lower_tirx.py | 5 +- 161 files changed, 5464 insertions(+), 3939 deletions(-) diff --git a/python/tvm/relax/backend/gpu_generic/cumsum.py b/python/tvm/relax/backend/gpu_generic/cumsum.py index ded3fb239275..cada8521efdd 100644 --- a/python/tvm/relax/backend/gpu_generic/cumsum.py +++ b/python/tvm/relax/backend/gpu_generic/cumsum.py @@ -174,10 +174,12 @@ def update_cross_block( bx > 0, source[by, src_offset + bx - 1], 0 ) + m = T.dynamic("m") + n = T.dynamic("n") + @Ts.prim_func(private=True) def cumsum(var_a: T.handle, var_out: T.handle): T.func_attr({"tirx.is_scheduled": True}) # prevent further scheduling - m, n = T.int64(), T.int64() A = T.match_buffer(var_a, [m, n], dtype=in_dtype) Out = T.match_buffer(var_out, [m, n], dtype=out_dtype) Tmp = T.alloc_buffer([m, n], dtype=out_dtype) @@ -253,10 +255,13 @@ def gpu_3d_axis_1_cumsum( out_dtype = out_dtype or in_dtype TX = T.int64(tx_len) + outer = T.dynamic("outer") + scan = T.dynamic("scan") + inner = T.dynamic("inner") + @Ts.prim_func(private=True) def cumsum(var_a: T.handle, var_out: T.handle): T.func_attr({"tirx.is_scheduled": True}) - outer, scan, inner = T.int64(), T.int64(), T.int64() A = T.match_buffer(var_a, [outer, scan, inner], dtype=in_dtype) Out = T.match_buffer(var_out, [outer, scan, inner], dtype=out_dtype) diff --git a/python/tvm/relax/backend/gpu_generic/sampling.py b/python/tvm/relax/backend/gpu_generic/sampling.py index 84431f91b724..f481e221eb30 100644 --- a/python/tvm/relax/backend/gpu_generic/sampling.py +++ b/python/tvm/relax/backend/gpu_generic/sampling.py @@ -259,6 +259,10 @@ def single_batch_sampling( aggregate[()] += step_aggregate[()] + n = T.dynamic("n") + vocab_size = T.dynamic("vocab_size") + batch_size = T.dynamic("batch_size") + @Ts.prim_func def parallel_sampling_from_prob( var_prob: T.handle, @@ -267,7 +271,6 @@ def parallel_sampling_from_prob( var_sampled_token_ids: T.handle, ): T.func_attr({"tirx.is_scheduled": True}) - n, vocab_size, batch_size = T.int64(), T.int64(), T.int64() # match buffers prob = T.match_buffer(var_prob, (n, vocab_size), prob_dtype) uniform_samples = T.match_buffer(var_uniform_samples, (batch_size, 1), sample_dtype) @@ -318,11 +321,13 @@ def generic_get_sample_index( ): """Generate a generic get_sample_index kernel.""" + batch = T.dynamic("batch") + vocab_size = T.dynamic("vocab_size") + out_batch = T.dynamic("out_batch") + @Ts.prim_func(private=True) def _get_sample_index(A: T.handle, B: T.handle, C: T.handle, D: T.handle): - batch, vocab_size = T.int64(), T.int64() prob = T.match_buffer(A, (batch, vocab_size), prob_dtype) - out_batch = T.int64() usample = T.match_buffer(B, (out_batch, 1), sample_dtype) sample_indices = T.match_buffer(C, (out_batch, 1), sample_indices_dtype) output_index = T.match_buffer(D, (out_batch, 1), dtype) diff --git a/python/tvm/relax/block_builder.py b/python/tvm/relax/block_builder.py index 13a50bfa138c..07e055408348 100644 --- a/python/tvm/relax/block_builder.py +++ b/python/tvm/relax/block_builder.py @@ -470,6 +470,9 @@ def te_func(args, args_dict, msg): .. code-block:: python + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Module: @Ts.prim_func @@ -477,8 +480,6 @@ def te_func(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_compute: T.handle) -> None: # function attr dict T.func_attr({"tirx.noalias": True}) - m = T.int64() - n = T.int64() rxplaceholder = T.match_buffer(var_rxplaceholder, [n, m], dtype="float32") rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [n, m], dtype="float32") compute = T.match_buffer(var_compute, [128, 128], dtype="float32") @@ -519,6 +520,9 @@ def te_func(A): .. code-block:: python + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Module: @Ts.prim_func diff --git a/python/tvm/relax/frontend/nn/llm/_decode_kernels.py b/python/tvm/relax/frontend/nn/llm/_decode_kernels.py index adbd414639f7..56040388ba96 100644 --- a/python/tvm/relax/frontend/nn/llm/_decode_kernels.py +++ b/python/tvm/relax/frontend/nn/llm/_decode_kernels.py @@ -57,6 +57,15 @@ def _attention_decode_cpu(num_kv_heads, num_qo_heads, head_dim, qkv_dtype, slidi if sliding_window: global_symbol += "_sliding_window" + B = T.dynamic("B", "int32") + nnz_pages = T.dynamic("nnz_pages", "int32") + max_num_pages = T.dynamic("max_num_pages", "int32") + page_indptr_elem_offset = T.dynamic("page_indptr_elem_offset", "int32") + page_values_elem_offset = T.dynamic("page_values_elem_offset", "int32") + k_rope_pos_offset_elem_offset = T.dynamic("k_rope_pos_offset_elem_offset", "int32") + q_rope_position_elem_offset = T.dynamic("q_rope_position_elem_offset", "int32") + length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") + @Ts.prim_func def batch_decode_paged_kv( Q_handle: T.handle, @@ -74,14 +83,6 @@ def batch_decode_paged_kv( sm_scale: T.float32, ): T.func_attr({"tirx.is_scheduled": True, "global_symbol": global_symbol}) - B = T.int32() - nnz_pages = T.int32() - max_num_pages = T.int32() - page_indptr_elem_offset = T.int32() - page_values_elem_offset = T.int32() - k_rope_pos_offset_elem_offset = T.int32() - q_rope_position_elem_offset = T.int32() - length_info_elem_offset = T.int32() Q = T.match_buffer(Q_handle, (B, H_qo, D), qkv_dtype) pages = T.match_buffer(pages_handle, (max_num_pages, 2, H_kv, page_size, D), qkv_dtype) @@ -212,6 +213,16 @@ def _attention_decode(num_kv_heads, num_qo_heads, head_dim, qkv_dtype, sliding_w global_symbol += "_sliding_window" # pylint: disable=too-many-branches + B = T.dynamic("B", "int32") + nnz_pages = T.dynamic("nnz_pages", "int32") + max_num_pages = T.dynamic("max_num_pages", "int32") + pages_elem_offset = T.dynamic("pages_elem_offset") + page_indptr_elem_offset = T.dynamic("page_indptr_elem_offset", "int32") + page_values_elem_offset = T.dynamic("page_values_elem_offset", "int32") + k_rope_pos_offset_elem_offset = T.dynamic("k_rope_pos_offset_elem_offset", "int32") + q_rope_position_elem_offset = T.dynamic("q_rope_position_elem_offset", "int32") + length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") + @Ts.prim_func def batch_decode_paged_kv( Q_handle: T.handle, @@ -229,15 +240,6 @@ def batch_decode_paged_kv( sm_scale: T.float32, ): T.func_attr({"tirx.is_scheduled": True, "global_symbol": global_symbol}) - B = T.int32() - nnz_pages = T.int32() - max_num_pages = T.int32() - pages_elem_offset = T.int64() - page_indptr_elem_offset = T.int32() - page_values_elem_offset = T.int32() - k_rope_pos_offset_elem_offset = T.int32() - q_rope_position_elem_offset = T.int32() - length_info_elem_offset = T.int32() Q = T.match_buffer(Q_handle, (B, H_qo, D), qkv_dtype) pages = T.match_buffer(pages_handle, (max_num_pages, 2, H_kv, page_size, D), qkv_dtype, elem_offset=pages_elem_offset) @@ -413,6 +415,10 @@ def batch_decode_paged_kv( def _merge_state_inplace_cpu(v_dtype): + N = T.dynamic("N", "int32") + H = T.dynamic("H", "int32") + D = T.dynamic("D", "int32") + @Ts.prim_func def merge_state_inplace_cpu( v: T.handle, @@ -421,9 +427,6 @@ def merge_state_inplace_cpu( s_other: T.handle, ): T.func_attr({"tirx.is_scheduled": True}) - N = T.int32() - H = T.int32() - D = T.int32() V = T.match_buffer(v, (N, H, D), v_dtype) S = T.match_buffer(s, (N, H), "float32") @@ -464,6 +467,10 @@ def _merge_state_inplace(num_heads, head_dim, v_dtype, target: Target, global_sy gdy = num_heads // bdy check_thread_limits(target, bdx=bdx, bdy=bdy, bdz=1, gdz=1) + N = T.dynamic("N", "int32") + H = T.dynamic("H", "int32") + D = T.dynamic("D", "int32") + @Ts.prim_func def merge_state_inplace( v: T.handle, @@ -472,9 +479,6 @@ def merge_state_inplace( s_other: T.handle, ): T.func_attr({"tirx.is_scheduled": True}) - N = T.int32() - H = T.int32() - D = T.int32() V = T.match_buffer(v, (N, H, D), v_dtype) S = T.match_buffer(s, (N, H), "float32") diff --git a/python/tvm/relax/frontend/nn/llm/_page_kernels.py b/python/tvm/relax/frontend/nn/llm/_page_kernels.py index da43ddd79f53..d416f96a6ebc 100644 --- a/python/tvm/relax/frontend/nn/llm/_page_kernels.py +++ b/python/tvm/relax/frontend/nn/llm/_page_kernels.py @@ -41,6 +41,11 @@ def _kv_cache_transpose_append(num_key_value_heads, head_dim, dtype, page_size: int = 16): """Return the TIR function that appends new k/v data to PagedKVCache.""" + ntoken = T.dynamic("ntoken") + num_pages = T.dynamic("num_pages") + pages_elem_offset = T.dynamic("pages_elem_offset") + position_map_elem_offset = T.dynamic("position_map_elem_offset", "int32") + @Ts.prim_func def tir_kv_cache_transpose_append( var_pages: T.handle, @@ -49,10 +54,6 @@ def tir_kv_cache_transpose_append( var_position_map: T.handle, ): T.func_attr({"tirx.noalias": True}) - ntoken = T.int64() - num_pages = T.int64() - pages_elem_offset = T.int64() - position_map_elem_offset = T.int32() pages = T.match_buffer(var_pages, (num_pages, 2, num_key_value_heads, page_size, head_dim), dtype, elem_offset=pages_elem_offset) k_data = T.match_buffer(var_k_data, (ntoken, num_key_value_heads, head_dim), dtype) v_data = T.match_buffer(var_v_data, (ntoken, num_key_value_heads, head_dim), dtype) @@ -78,6 +79,11 @@ def tir_kv_cache_transpose_append( def _kv_cache_transpose_append_mla(d_qk: int, dtype, page_size: int = 16): """Return the TIR function that appends new compressed KV data to PagedKVCache for MLA.""" + ntoken = T.dynamic("ntoken") + num_pages = T.dynamic("num_pages") + pages_elem_offset = T.dynamic("pages_elem_offset") + position_map_elem_offset = T.dynamic("position_map_elem_offset", "int32") + @Ts.prim_func def tir_kv_cache_transpose_append_mla( var_pages: T.handle, @@ -85,10 +91,6 @@ def tir_kv_cache_transpose_append_mla( var_position_map: T.handle, ): T.func_attr({"tirx.noalias": True}) - ntoken = T.int64() - num_pages = T.int64() - pages_elem_offset = T.int64() - position_map_elem_offset = T.int32() pages = T.match_buffer(var_pages, (num_pages, page_size, d_qk), dtype, elem_offset=pages_elem_offset) kv_data = T.match_buffer(var_kv_data, (ntoken, d_qk), dtype) position_map = T.match_buffer(var_position_map, (ntoken,), "int32", elem_offset=position_map_elem_offset) @@ -107,6 +109,12 @@ def tir_kv_cache_transpose_append_mla( def _kv_cache_debug_get_kv(num_hidden_layers, num_key_value_heads, head_dim, dtype): """Return the TIR function that fetches the k/v data on given positions and layer.""" + seqlen = T.dynamic("seqlen") + page_size = T.dynamic("page_size") + num_pages = T.dynamic("num_pages") + pages_elem_offset = T.dynamic("pages_elem_offset") + position_map_elem_offset = T.dynamic("position_map_elem_offset") + @Ts.prim_func def tir_kv_cache_debug_get_kv( var_pages: T.handle, @@ -116,11 +124,6 @@ def tir_kv_cache_debug_get_kv( layer_id: T.int64, ): T.func_attr({"tirx.noalias": True}) - seqlen = T.int64() - page_size = T.int64() - num_pages = T.int64() - pages_elem_offset = T.int64() - position_map_elem_offset = T.int64() pages = T.match_buffer(var_pages, (num_pages, 2, num_key_value_heads, page_size, head_dim), dtype,elem_offset=pages_elem_offset) position_map = T.match_buffer(var_position_map, (seqlen,), "int32", elem_offset=position_map_elem_offset) k_data = T.match_buffer(var_k_data, (num_hidden_layers, seqlen, num_key_value_heads, head_dim), dtype) @@ -140,6 +143,12 @@ def tir_kv_cache_debug_get_kv( def _kv_cache_debug_get_kv_mla(num_hidden_layers, d_qk, dtype): """Return the TIR function that fetches the k/v data on given positions and layer.""" + seqlen = T.dynamic("seqlen") + page_size = T.dynamic("page_size") + num_pages = T.dynamic("num_pages") + pages_elem_offset = T.dynamic("pages_elem_offset") + position_map_elem_offset = T.dynamic("position_map_elem_offset") + @Ts.prim_func def tir_kv_cache_debug_get_kv_mla( var_pages: T.handle, @@ -148,11 +157,6 @@ def tir_kv_cache_debug_get_kv_mla( layer_id: T.int64, ): T.func_attr({"tirx.noalias": True}) - seqlen = T.int64() - page_size = T.int64() - num_pages = T.int64() - pages_elem_offset = T.int64() - position_map_elem_offset = T.int64() pages = T.match_buffer(var_pages, (num_pages, page_size, d_qk), dtype, elem_offset=pages_elem_offset) position_map = T.match_buffer(var_position_map, (seqlen,), "int32", elem_offset=position_map_elem_offset) compressed_kv_with_k_pe_data = T.match_buffer(var_compressed_kv_with_k_pe_data, (num_hidden_layers, seqlen, d_qk), dtype) @@ -170,11 +174,12 @@ def tir_kv_cache_debug_get_kv_mla( def _copy_single_page(num_heads, page_size, head_dim, dtype, target: Target): tx = get_max_num_threads_per_block(target) + num_pages = T.dynamic("num_pages", "int32") + pages_elem_offset = T.dynamic("pages_elem_offset") + @Ts.prim_func def copy_single_page(var_pages: T.handle, src_page_id: T.int64, tgt_page_id: T.int64, copy_length: T.int64): T.func_attr({"tirx.is_scheduled": True}) - num_pages = T.int32() - pages_elem_offset = T.int64() pages = T.match_buffer(var_pages, (num_pages, 2, num_heads, page_size, head_dim), dtype, elem_offset=pages_elem_offset) for b in T.thread_binding((copy_length * num_heads * head_dim + tx - 1) // tx, thread="blockIdx.x"): @@ -193,11 +198,12 @@ def copy_single_page(var_pages: T.handle, src_page_id: T.int64, tgt_page_id: T.i def _copy_single_page_mla(page_size, head_dim, dtype, target: Target): tx = get_max_num_threads_per_block(target) + num_pages = T.dynamic("num_pages", "int32") + pages_elem_offset = T.dynamic("pages_elem_offset") + @Ts.prim_func def copy_single_page_mla(var_pages: T.handle, src_page_id: T.int64, tgt_page_id: T.int64, copy_length: T.int64): T.func_attr({"tirx.is_scheduled": True}) - num_pages = T.int32() - pages_elem_offset = T.int64() pages = T.match_buffer(var_pages, (num_pages, page_size, head_dim), dtype, elem_offset=pages_elem_offset) for b in T.thread_binding((copy_length * head_dim + tx - 1) // tx, thread="blockIdx.x"): @@ -214,10 +220,11 @@ def copy_single_page_mla(var_pages: T.handle, src_page_id: T.int64, tgt_page_id: def _copy_single_page_cpu(num_heads, page_size, head_dim, dtype): tx = 1 + num_pages = T.dynamic("num_pages", "int32") + @Ts.prim_func def copy_single_page_cpu(var_pages: T.handle, src_page_id: T.int64, tgt_page_id: T.int64, copy_length: T.int64): T.func_attr({"tirx.is_scheduled": True}) - num_pages = T.int32() pages = T.match_buffer(var_pages, (num_pages, 2, num_heads, page_size, head_dim), dtype) for b in T.serial((copy_length * num_heads * head_dim + tx - 1) // tx): @@ -236,14 +243,15 @@ def copy_single_page_cpu(var_pages: T.handle, src_page_id: T.int64, tgt_page_id: def _compact_kv_copy(num_heads, head_dim, dtype, target: Target, page_size: int = 16): tx = get_max_num_threads_per_block(target) + num_pages = T.dynamic("num_pages", "int32") + total_copy_length = T.dynamic("total_copy_length", "int32") + copy_length_indptr_elem_offset = T.dynamic("copy_length_indptr_elem_offset", "int32") + copy_src_dst_pos_elem_offset = T.dynamic("copy_src_dst_pos_elem_offset", "int32") + pages_elem_offset = T.dynamic("pages_elem_offset") + @Ts.prim_func def compact_kv_copy(var_pages: T.handle, var_copy_length_indptr: T.handle, var_copy_src_dst_pos: T.handle, batch_size: T.int32): T.func_attr({"tirx.is_scheduled": True}) - num_pages = T.int32() - total_copy_length = T.int32() - copy_length_indptr_elem_offset = T.int32() - copy_src_dst_pos_elem_offset = T.int32() - pages_elem_offset = T.int64() pages = T.match_buffer(var_pages, (num_pages, 2, num_heads, page_size, head_dim), dtype, elem_offset=pages_elem_offset) copy_length_indptr = T.match_buffer(var_copy_length_indptr, (batch_size + 1,), "int32", elem_offset=copy_length_indptr_elem_offset) copy_src_dst_pos = T.match_buffer(var_copy_src_dst_pos, (2, total_copy_length), "int32", elem_offset=copy_src_dst_pos_elem_offset) @@ -267,13 +275,14 @@ def compact_kv_copy(var_pages: T.handle, var_copy_length_indptr: T.handle, var_c def _compact_kv_copy_cpu(num_heads, head_dim, dtype, page_size: int = 16): tx = 8 + num_pages = T.dynamic("num_pages", "int32") + total_copy_length = T.dynamic("total_copy_length", "int32") + copy_length_indptr_elem_offset = T.dynamic("copy_length_indptr_elem_offset", "int32") + copy_src_dst_pos_elem_offset = T.dynamic("copy_src_dst_pos_elem_offset", "int32") + @Ts.prim_func def compact_kv_copy_cpu(var_pages: T.handle, var_copy_length_indptr: T.handle, var_copy_src_dst_pos: T.handle, batch_size: T.int32): T.func_attr({"tirx.is_scheduled": True}) - num_pages = T.int32() - total_copy_length = T.int32() - copy_length_indptr_elem_offset = T.int32() - copy_src_dst_pos_elem_offset = T.int32() pages = T.match_buffer(var_pages, (num_pages, 2, num_heads, page_size, head_dim), dtype) copy_length_indptr = T.match_buffer(var_copy_length_indptr, (batch_size + 1,), "int32", elem_offset=copy_length_indptr_elem_offset) copy_src_dst_pos = T.match_buffer(var_copy_src_dst_pos, (2, total_copy_length), "int32", elem_offset=copy_src_dst_pos_elem_offset) diff --git a/python/tvm/relax/frontend/nn/llm/_prefill_kernels.py b/python/tvm/relax/frontend/nn/llm/_prefill_kernels.py index 2cfee6b1f41b..d4f55b037c4f 100644 --- a/python/tvm/relax/frontend/nn/llm/_prefill_kernels.py +++ b/python/tvm/relax/frontend/nn/llm/_prefill_kernels.py @@ -69,6 +69,17 @@ def _attention_prefill_cpu( group_size = h_q // h_kv # pylint: disable=too-many-branches + batch_size = T.dynamic("batch_size", "int32") + total_len = T.dynamic("total_len", "int32") + nnz_pages = T.dynamic("nnz_pages", "int32") + max_num_pages = T.dynamic("max_num_pages", "int32") + q_indptr_elem_offset = T.dynamic("q_indptr_elem_offset", "int32") + page_indptr_elem_offset = T.dynamic("page_indptr_elem_offset", "int32") + page_values_elem_offset = T.dynamic("page_values_elem_offset", "int32") + k_rope_pos_offset_elem_offset = T.dynamic("k_rope_pos_offset_elem_offset", "int32") + q_rope_position_elem_offset = T.dynamic("q_rope_position_elem_offset", "int32") + length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") + @Ts.prim_func def batch_prefill_paged_kv_cpu( var_q: T.handle, # [total_len, h_q, d] @@ -88,16 +99,6 @@ def batch_prefill_paged_kv_cpu( sm_scale: T.float32, ): T.func_attr({"global_symbol": global_symbol}) - batch_size = T.int32() - total_len = T.int32() - nnz_pages = T.int32() - max_num_pages = T.int32() - q_indptr_elem_offset = T.int32() - page_indptr_elem_offset = T.int32() - page_values_elem_offset = T.int32() - k_rope_pos_offset_elem_offset = T.int32() - q_rope_position_elem_offset = T.int32() - length_info_elem_offset = T.int32() q = T.match_buffer(var_q, (total_len, h_q, d), dtype) q_indptr = T.match_buffer(var_q_indptr, (batch_size + 1,), "int32", elem_offset=q_indptr_elem_offset) @@ -235,6 +236,18 @@ def _attention_prefill( init_states, compute_s_gemm, softmax_update_causal, compute_o_gemm, _, advance_tile_batch, paged_store_output_lse, *_ = _make_prefill_macros(tile_x, tile_y, tile_z, tile_y, bdx, num_warps, group_size) # pylint: disable=too-many-branches + batch_size = T.dynamic("batch_size", "int32") + total_len = T.dynamic("total_len", "int32") + nnz_pages = T.dynamic("nnz_pages", "int32") + max_num_pages = T.dynamic("max_num_pages", "int32") + pages_elem_offset = T.dynamic("pages_elem_offset") + q_indptr_elem_offset = T.dynamic("q_indptr_elem_offset", "int32") + page_indptr_elem_offset = T.dynamic("page_indptr_elem_offset", "int32") + page_values_elem_offset = T.dynamic("page_values_elem_offset", "int32") + k_rope_pos_offset_elem_offset = T.dynamic("k_rope_pos_offset_elem_offset", "int32") + q_rope_position_elem_offset = T.dynamic("q_rope_position_elem_offset", "int32") + length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") + @Ts.prim_func def batch_prefill_paged_kv( var_q: T.handle, # [total_len, h_q, d] @@ -254,17 +267,6 @@ def batch_prefill_paged_kv( sm_scale: T.float32, ): T.func_attr({"global_symbol": global_symbol}) - batch_size = T.int32() - total_len = T.int32() - nnz_pages = T.int32() - max_num_pages = T.int32() - pages_elem_offset = T.int64() - q_indptr_elem_offset = T.int32() - page_indptr_elem_offset = T.int32() - page_values_elem_offset = T.int32() - k_rope_pos_offset_elem_offset = T.int32() - q_rope_position_elem_offset = T.int32() - length_info_elem_offset = T.int32() q = T.match_buffer(var_q, (total_len, h_q, d), dtype) q_indptr = T.match_buffer(var_q_indptr, (batch_size + 1,), "int32", elem_offset=q_indptr_elem_offset) @@ -397,6 +399,10 @@ def _attention_sequence_prefill(h_kv, h_q, d, dtype, target: Target, causal=0, s _, LOAD_VEC, group_size, bdx, num_warps, tile_x, tile_y, tile_z = _get_prefill_kernel_config(h_kv, h_q, d, dtype, target) init_states, compute_s_gemm, softmax_update_causal, compute_o_gemm, *_ = _make_prefill_macros(tile_x, tile_y, tile_z, tile_y, bdx, num_warps, group_size) + batch_size = T.dynamic("batch_size", "int32") + qo_len = T.dynamic("qo_len", "int32") + kv_len = T.dynamic("kv_len", "int32") + @Ts.prim_func def batch_sequence_prefill_kv( # pylint: disable=too-many-branches var_q: T.handle, # [total_len, h_q, d] @@ -405,9 +411,6 @@ def batch_sequence_prefill_kv( # pylint: disable=too-many-branches var_output: T.handle, # [total_len, h_q, d] var_lse: T.handle # [total_len, h_q] ): - batch_size = T.int32() - qo_len = T.int32() - kv_len = T.int32() q = T.match_buffer(var_q, (batch_size, qo_len, h_q, d), dtype) k = T.match_buffer(var_k, (batch_size, kv_len, h_kv, d), dtype) v = T.match_buffer(var_v, (batch_size, kv_len, h_kv, d), dtype) @@ -564,6 +567,10 @@ def _kv_col_valid(col, valid_len, kv_len): pad = kv_len - valid_len return tirx.And(col < kv_len, col >= pad) + batch_size = T.dynamic("batch_size", "int32") + qo_len = T.dynamic("qo_len", "int32") + kv_len = T.dynamic("kv_len", "int32") + @Ts.prim_func def batch_sequence_prefill_kv_masked( # pylint: disable=too-many-branches var_q: T.handle, # [batch_size, qo_len, h_q, d] @@ -573,9 +580,6 @@ def batch_sequence_prefill_kv_masked( # pylint: disable=too-many-branches var_output: T.handle, # [batch_size, qo_len, h_q, d] var_lse: T.handle # [batch_size, qo_len, h_q] ): - batch_size = T.int32() - qo_len = T.int32() - kv_len = T.int32() q = T.match_buffer(var_q, (batch_size, qo_len, h_q, d), dtype) k = T.match_buffer(var_k, (batch_size, kv_len, h_kv, d), dtype) v = T.match_buffer(var_v, (batch_size, kv_len, h_kv, d), dtype) @@ -678,6 +682,14 @@ def batch_sequence_prefill_kv_masked( # pylint: disable=too-many-branches def _attention_prefill_ragged_cpu(h_kv, h_q, d_qk, d_v, dtype, rope_scaling: dict[str, Any]): group_size = h_q // h_kv + batch_size = T.dynamic("batch_size", "int32") + qo_len = T.dynamic("qo_len", "int32") + kv_len = T.dynamic("kv_len", "int32") + q_indptr_elem_offset = T.dynamic("q_indptr_elem_offset", "int32") + kv_indptr_elem_offset = T.dynamic("kv_indptr_elem_offset", "int32") + q_rope_position_elem_offset = T.dynamic("q_rope_position_elem_offset", "int32") + k_rope_pos_offset_elem_offset = T.dynamic("k_rope_pos_offset_elem_offset", "int32") + @Ts.prim_func def batch_prefill_ragged_kv( # pylint: disable=too-many-branches var_q: T.handle, # [total_len, h_q, d_qk] @@ -695,13 +707,6 @@ def batch_prefill_ragged_kv( # pylint: disable=too-many-branches rope_theta: T.float32, sm_scale: T.float32, ): - batch_size = T.int32() - qo_len = T.int32() - kv_len = T.int32() - q_indptr_elem_offset = T.int32() - kv_indptr_elem_offset = T.int32() - q_rope_position_elem_offset = T.int32() - k_rope_pos_offset_elem_offset = T.int32() q = T.match_buffer(var_q, (qo_len, h_q, d_qk), dtype) q_indptr = T.match_buffer(var_q_indptr, (batch_size + 1,), "int32", elem_offset=q_indptr_elem_offset) @@ -797,6 +802,14 @@ def _attention_prefill_ragged(h_kv, h_q, d_qk, d_v, dtype, rope_scaling: dict[st NUM_BLKS, LOAD_VEC, group_size, bdx, num_warps, tile_x, tile_y, tile_z = _get_prefill_kernel_config(h_kv, h_q, d_qk, dtype, target, d_v=d_v) init_states, compute_s_gemm, softmax_update_causal, compute_o_gemm, _, advance_tile_batch, paged_store_output_lse, *_ = _make_prefill_macros(tile_x, tile_y, tile_z, d_v, bdx, num_warps, group_size) + batch_size = T.dynamic("batch_size", "int32") + qo_len = T.dynamic("qo_len", "int32") + kv_len = T.dynamic("kv_len", "int32") + q_indptr_elem_offset = T.dynamic("q_indptr_elem_offset", "int32") + kv_indptr_elem_offset = T.dynamic("kv_indptr_elem_offset", "int32") + q_rope_position_elem_offset = T.dynamic("q_rope_position_elem_offset", "int32") + k_rope_pos_offset_elem_offset = T.dynamic("k_rope_pos_offset_elem_offset", "int32") + @Ts.prim_func def batch_prefill_ragged_kv( # pylint: disable=too-many-branches var_q: T.handle, # [total_len, h_q, d_qk] @@ -814,13 +827,6 @@ def batch_prefill_ragged_kv( # pylint: disable=too-many-branches rope_theta: T.float32, sm_scale: T.float32 ): - batch_size = T.int32() - qo_len = T.int32() - kv_len = T.int32() - q_indptr_elem_offset = T.int32() - kv_indptr_elem_offset = T.int32() - q_rope_position_elem_offset = T.int32() - k_rope_pos_offset_elem_offset = T.int32() q = T.match_buffer(var_q, (qo_len, h_q, d_qk), dtype) q_indptr = T.match_buffer(var_q_indptr, (batch_size + 1,), "int32", elem_offset=q_indptr_elem_offset) @@ -935,6 +941,16 @@ def _attention_prefill_mla(h_q, d_latent, d_rope, dtype, sliding_window: bool, t global_symbol += "_sliding_window" # pylint: disable=too-many-branches + batch_size = T.dynamic("batch_size", "int32") + total_len = T.dynamic("total_len", "int32") + nnz_pages = T.dynamic("nnz_pages", "int32") + max_num_pages = T.dynamic("max_num_pages", "int32") + pages_elem_offset = T.dynamic("pages_elem_offset") + q_indptr_elem_offset = T.dynamic("q_indptr_elem_offset", "int32") + page_indptr_elem_offset = T.dynamic("page_indptr_elem_offset", "int32") + page_values_elem_offset = T.dynamic("page_values_elem_offset", "int32") + length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") + @Ts.prim_func def batch_prefill_paged_kv_mla( var_q: T.handle, # [total_len, h_q, d_qk] @@ -949,15 +965,6 @@ def batch_prefill_paged_kv_mla( sm_scale: T.float32, ): T.func_attr({"global_symbol": global_symbol}) - batch_size = T.int32() - total_len = T.int32() - nnz_pages = T.int32() - max_num_pages = T.int32() - pages_elem_offset = T.int64() - q_indptr_elem_offset = T.int32() - page_indptr_elem_offset = T.int32() - page_values_elem_offset = T.int32() - length_info_elem_offset = T.int32() q = T.match_buffer(var_q, (total_len, h_q, d_qk), dtype) q_indptr = T.match_buffer(var_q_indptr, (batch_size + 1,), "int32", elem_offset=q_indptr_elem_offset) diff --git a/python/tvm/relax/frontend/nn/llm/position_embedding.py b/python/tvm/relax/frontend/nn/llm/position_embedding.py index b057cde24215..791eaac0edb7 100644 --- a/python/tvm/relax/frontend/nn/llm/position_embedding.py +++ b/python/tvm/relax/frontend/nn/llm/position_embedding.py @@ -391,6 +391,9 @@ def _rope( # pylint: disable=too-many-arguments expr = tirx.Let(var, value, expr) return expr + batch_size = T.dynamic("batch_size") + seq_len = T.dynamic("seq_len") + @Ts.prim_func(private=True) def fused_rope( # pylint: disable=too-many-locals var_qkv: T.handle, @@ -405,8 +408,6 @@ def fused_rope( # pylint: disable=too-many-locals "tirx.noalias": True, } ) - batch_size = T.int64() - seq_len = T.int64() qkv = T.match_buffer(var_qkv, (batch_size, seq_len, fused_heads, head_dim), dtype) q = T.match_buffer(var_q, (batch_size, seq_len, num_q_heads, head_dim), dtype) k = T.match_buffer(var_k, (batch_size, seq_len, num_kv_heads, head_dim), dtype) @@ -523,6 +524,9 @@ def _rope( # pylint: disable=too-many-arguments expr = tirx.Let(var, value, expr) return expr + seq_len = T.dynamic("seq_len", "int32") + position_map_elem_offset = T.dynamic("position_map_elem_offset", "int32") + @Ts.prim_func def fused_rope( # pylint: disable=too-many-locals var_qkv: T.handle, @@ -538,8 +542,6 @@ def fused_rope( # pylint: disable=too-many-locals "tirx.noalias": True, } ) - seq_len = T.int32() - position_map_elem_offset = T.int32() qkv = T.match_buffer(var_qkv, (seq_len, fused_heads, head_dim), dtype) q = T.match_buffer(var_q, (seq_len, num_q_heads, head_dim), dtype) k = T.match_buffer(var_k, (seq_len, num_kv_heads, head_dim), dtype) @@ -565,6 +567,9 @@ def fused_rope( # pylint: disable=too-many-locals else: v[s, h - (num_q_heads + num_kv_heads), d] = qkv[s, h, d] + seq_len = T.dynamic("seq_len") + position_map_elem_offset = T.dynamic("position_map_elem_offset") + @Ts.prim_func def fused_rope_longrope_scaling( # pylint: disable=too-many-locals var_qkv: T.handle, @@ -580,8 +585,6 @@ def fused_rope_longrope_scaling( # pylint: disable=too-many-locals "tirx.noalias": True, } ) - seq_len = T.int64() - position_map_elem_offset = T.int64() qkv = T.match_buffer(var_qkv, (seq_len, fused_heads, head_dim), dtype) q = T.match_buffer(var_q, (seq_len, num_q_heads, head_dim), dtype) k = T.match_buffer(var_k, (seq_len, num_kv_heads, head_dim), dtype) @@ -750,6 +753,9 @@ def _rope( # pylint: disable=too-many-arguments expr = tirx.Let(var, value, expr) return expr + seq_len = T.dynamic("seq_len", "int32") + position_map_elem_offset = T.dynamic("position_map_elem_offset", "int32") + @Ts.prim_func(private=True) def fused_rope( # pylint: disable=too-many-locals var_qkv: T.handle, @@ -765,8 +771,6 @@ def fused_rope( # pylint: disable=too-many-locals "tirx.noalias": True, } ) - seq_len = T.int32() - position_map_elem_offset = T.int32() qkv = T.match_buffer(var_qkv, (seq_len, fused_heads, head_dim), dtype) q = T.match_buffer(var_q, (seq_len, num_q_heads, head_dim), dtype) k = T.match_buffer(var_k, (seq_len, num_kv_heads, head_dim), dtype) @@ -792,6 +796,9 @@ def fused_rope( # pylint: disable=too-many-locals else: v[s, h - (num_q_heads + num_kv_heads), d] = qkv[s, h, d] + seq_len = T.dynamic("seq_len") + position_map_elem_offset = T.dynamic("position_map_elem_offset") + @Ts.prim_func def fused_rope_longrope_scaling( # pylint: disable=too-many-locals var_qkv: T.handle, @@ -807,8 +814,6 @@ def fused_rope_longrope_scaling( # pylint: disable=too-many-locals "tirx.noalias": True, } ) - seq_len = T.int64() - position_map_elem_offset = T.int64() qkv = T.match_buffer(var_qkv, (seq_len, fused_heads, head_dim), dtype) q = T.match_buffer(var_q, (seq_len, num_q_heads, head_dim), dtype) k = T.match_buffer(var_k, (seq_len, num_kv_heads, head_dim), dtype) diff --git a/python/tvm/relax/frontend/nn/llm/tree_attn.py b/python/tvm/relax/frontend/nn/llm/tree_attn.py index 98ae8b110c8e..bacac97e1e6b 100644 --- a/python/tvm/relax/frontend/nn/llm/tree_attn.py +++ b/python/tvm/relax/frontend/nn/llm/tree_attn.py @@ -89,6 +89,16 @@ def tree_attn_cpu(h_kv, h_q, d, dtype, rope_scaling: dict[str, Any]): group_size = h_q // h_kv # fmt: off + qo_len = T.dynamic("qo_len", "int32") + kv_len = T.dynamic("kv_len", "int32") + q_indptr_elem_offset = T.dynamic("q_indptr_elem_offset", "int32") + kv_indptr_elem_offset = T.dynamic("kv_indptr_elem_offset", "int32") + q_rope_position_elem_offset = T.dynamic("q_rope_position_elem_offset", "int32") + mn_indptr_elem_offset = T.dynamic("mn_indptr_elem_offset", "int32") + mask_elem_offset = T.dynamic("mask_elem_offset", "int32") + tree_size = T.dynamic("tree_size", "int32") + batch_size_plus_1 = T.dynamic("batch_size_plus_1", "int32") + @Ts.prim_func def batch_tree_attn( # pylint: disable=too-many-branches,line-too-long var_q: T.handle, # [total_len, h_q, d] @@ -106,15 +116,6 @@ def batch_tree_attn( # pylint: disable=too-many-branches,line-too-long rope_theta: T.float32, sm_scale: T.float32, ): - qo_len = T.int32() - kv_len = T.int32() - q_indptr_elem_offset = T.int32() - kv_indptr_elem_offset = T.int32() - q_rope_position_elem_offset = T.int32() - mn_indptr_elem_offset = T.int32() - mask_elem_offset = T.int32() - tree_size = T.int32() - batch_size_plus_1 = T.int32() q = T.match_buffer(var_q, (qo_len, h_q, d), dtype) q_indptr = T.match_buffer( @@ -288,6 +289,16 @@ def tree_attn(h_kv, h_q, d, dtype, rope_scaling: dict[str, Any], target: Target) ) # fmt: off + qo_len = T.dynamic("qo_len", "int32") + kv_len = T.dynamic("kv_len", "int32") + q_indptr_elem_offset = T.dynamic("q_indptr_elem_offset", "int32") + kv_indptr_elem_offset = T.dynamic("kv_indptr_elem_offset", "int32") + q_rope_position_elem_offset = T.dynamic("q_rope_position_elem_offset", "int32") + mn_indptr_elem_offset = T.dynamic("mn_indptr_elem_offset", "int32") + mask_elem_offset = T.dynamic("mask_elem_offset", "int32") + tree_size = T.dynamic("tree_size", "int32") + batch_size_plus_1 = T.dynamic("batch_size_plus_1", "int32") + @Ts.prim_func def batch_tree_attn( # pylint: disable=too-many-branches var_q: T.handle, # [total_len, h_q, d] @@ -305,15 +316,6 @@ def batch_tree_attn( # pylint: disable=too-many-branches rope_theta: T.float32, sm_scale: T.float32, ): - qo_len = T.int32() - kv_len = T.int32() - q_indptr_elem_offset = T.int32() - kv_indptr_elem_offset = T.int32() - q_rope_position_elem_offset = T.int32() - mn_indptr_elem_offset = T.int32() - mask_elem_offset = T.int32() - tree_size = T.int32() - batch_size_plus_1 = T.int32() q = T.match_buffer(var_q, (qo_len, h_q, d), dtype) q_indptr = T.match_buffer(var_q_indptr, (batch_size_plus_1,), "int32", elem_offset=q_indptr_elem_offset) @@ -608,6 +610,20 @@ def tree_attn_with_paged_kv_cache_cpu(h_kv, h_q, d, dtype, rope_scaling: dict[st # pylint: disable=line-too-long,too-many-branches # fmt: off + batch_size = T.dynamic("batch_size", "int32") + total_len = T.dynamic("total_len", "int32") + nnz_pages = T.dynamic("nnz_pages", "int32") + max_num_pages = T.dynamic("max_num_pages", "int32") + q_indptr_elem_offset = T.dynamic("q_indptr_elem_offset", "int32") + page_indptr_elem_offset = T.dynamic("page_indptr_elem_offset", "int32") + page_values_elem_offset = T.dynamic("page_values_elem_offset", "int32") + k_rope_pos_offset_elem_offset = T.dynamic("k_rope_pos_offset_elem_offset", "int32") + q_rope_position_elem_offset = T.dynamic("q_rope_position_elem_offset", "int32") + length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") + tree_order_elem_offset = T.dynamic("tree_order_elem_offset", "int32") + tree_order_indptr_elem_offset = T.dynamic("tree_order_indptr_elem_offset", "int32") + total_tree_order_len = T.dynamic("total_tree_order_len", "int32") + @Ts.prim_func def tree_attn_paged_kv_cpu( var_q: T.handle, # [total_len, h_q, d] @@ -628,18 +644,6 @@ def tree_attn_paged_kv_cpu( tree_order_handle: T.handle, # [total_len, 2] ): T.func_attr({"global_symbol": global_symbol}) - batch_size = T.int32() - total_len = T.int32() - nnz_pages = T.int32() - max_num_pages = T.int32() - q_indptr_elem_offset = T.int32() - page_indptr_elem_offset = T.int32() - page_values_elem_offset = T.int32() - k_rope_pos_offset_elem_offset = T.int32() - q_rope_position_elem_offset = T.int32() - length_info_elem_offset = T.int32() - tree_order_elem_offset = T.int32() - tree_order_indptr_elem_offset = T.int32() q = T.match_buffer(var_q, (total_len, h_q, d), dtype) q_indptr = T.match_buffer(var_q_indptr, (batch_size + 1,), "int32", elem_offset=q_indptr_elem_offset) @@ -656,7 +660,6 @@ def tree_attn_paged_kv_cpu( "int32", elem_offset=tree_order_indptr_elem_offset, ) - total_tree_order_len = T.int32() tree_order = T.match_buffer( tree_order_handle, (total_tree_order_len, 2), @@ -805,6 +808,20 @@ def tree_attn_with_paged_kv_cache( sliding_window = False # Sliding window is not supported in this kernel. # fmt: off + batch_size = T.dynamic("batch_size", "int32") + total_len = T.dynamic("total_len", "int32") + nnz_pages = T.dynamic("nnz_pages", "int32") + max_num_pages = T.dynamic("max_num_pages", "int32") + q_indptr_elem_offset = T.dynamic("q_indptr_elem_offset", "int32") + k_rope_pos_offset_elem_offset = T.dynamic("k_rope_pos_offset_elem_offset", "int32") + q_rope_position_elem_offset = T.dynamic("q_rope_position_elem_offset", "int32") + page_indptr_elem_offset = T.dynamic("page_indptr_elem_offset", "int32") + page_values_elem_offset = T.dynamic("page_values_elem_offset", "int32") + length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") + tree_order_elem_offset = T.dynamic("tree_order_elem_offset", "int32") + tree_order_indptr_elem_offset = T.dynamic("tree_order_indptr_elem_offset", "int32") + total_tree_order_len = T.dynamic("total_tree_order_len", "int32") + @Ts.prim_func def tree_attn_paged_kv( var_q: T.handle, # [total_len, h_q, d] @@ -826,18 +843,6 @@ def tree_attn_paged_kv( ): # pylint: disable=unused-variable, too-many-branches T.func_attr({"global_symbol": global_symbol}) - batch_size = T.int32() - total_len = T.int32() - nnz_pages = T.int32() - max_num_pages = T.int32() - q_indptr_elem_offset = T.int32() - k_rope_pos_offset_elem_offset = T.int32() - q_rope_position_elem_offset = T.int32() - page_indptr_elem_offset = T.int32() - page_values_elem_offset = T.int32() - length_info_elem_offset = T.int32() - tree_order_elem_offset = T.int32() - tree_order_indptr_elem_offset = T.int32() q = T.match_buffer(var_q, (total_len, h_q, d), dtype) q_indptr = T.match_buffer( @@ -866,7 +871,6 @@ def tree_attn_paged_kv( "int32", elem_offset=tree_order_indptr_elem_offset, ) - total_tree_order_len = T.int32() tree_order = T.match_buffer( tree_order_handle, (total_tree_order_len, 2), diff --git a/python/tvm/relax/frontend/nn/op.py b/python/tvm/relax/frontend/nn/op.py index dd6376b35570..f7e008895ec5 100644 --- a/python/tvm/relax/frontend/nn/op.py +++ b/python/tvm/relax/frontend/nn/op.py @@ -2797,9 +2797,11 @@ def sample_top_p_top_k_from_sorted_prob( def _cumsum_mask(cumsum_sorted, top_p, top_k, i, j): return _tir.all(cumsum_sorted[i, j] < top_p[i, 0], j + 1 < top_k[i, 0]) + batch = T.dynamic("batch") + vocab_size = T.dynamic("vocab_size") + @Ts.prim_func(private=True) def _get_renorm_prob(A: T.handle, B: T.handle, C: T.handle, D: T.handle): - batch, vocab_size = T.int64(), T.int64() cumsum_sorted = T.match_buffer(A, (batch, vocab_size), prob_dtype) top_p = T.match_buffer(B, (batch, 1), prob_dtype) top_k = T.match_buffer(C, (batch, 1), index_dtype) @@ -2815,12 +2817,14 @@ def _get_renorm_prob(A: T.handle, B: T.handle, C: T.handle, D: T.handle): elif not _cumsum_mask(cumsum_sorted, top_p, top_k, v_ax0, v_ax1 + 1): renorm_prob[v_ax0, 0] = cumsum_sorted[v_ax0, v_ax1 + 1] + batch = T.dynamic("batch") + vocab_size = T.dynamic("vocab_size") + out_batch = T.dynamic("out_batch") + @Ts.prim_func(private=True) def _get_index_from_sorted( A: T.handle, B: T.handle, C: T.handle, D: T.handle, E: T.handle, F: T.handle ): - batch, vocab_size = T.int64(), T.int64() - out_batch = T.int64() cumsum_sorted = T.match_buffer(A, (batch, vocab_size), prob_dtype) indices = T.match_buffer(B, (batch, vocab_size), index_dtype) renorm_prob = T.match_buffer(C, (batch, 1), prob_dtype) @@ -2903,9 +2907,11 @@ def renormalize_top_p_top_k_prob(prob, sorted_prob, top_p, top_k): def _cumsum_mask(cumsum_sorted, top_p, top_k, i, j): return _tir.all(cumsum_sorted[i, j] < top_p[i, 0], j + 1 < top_k[i, 0]) + batch = T.dynamic("batch") + vocab_size = T.dynamic("vocab_size") + @Ts.prim_func(private=True) def _get_renorm_cutoff(A: T.handle, B: T.handle, C: T.handle, D: T.handle, E: T.handle): - batch, vocab_size = T.int64(), T.int64() sorted_prob = T.match_buffer(A, (batch, vocab_size), prob_dtype) cumsum_sorted = T.match_buffer(B, (batch, vocab_size), prob_dtype) top_p = T.match_buffer(C, (batch, 1), prob_dtype) diff --git a/python/tvm/s_tir/tensor_intrin/cuda.py b/python/tvm/s_tir/tensor_intrin/cuda.py index 629916b53de0..57e2e4ab1d66 100644 --- a/python/tvm/s_tir/tensor_intrin/cuda.py +++ b/python/tvm/s_tir/tensor_intrin/cuda.py @@ -182,10 +182,11 @@ def ldmatrix_desc(warp_handle: T.handle, shared_handle: T.handle) -> None: Ts.writes(warp[warp_indices[0], warp_indices[1]]) warp[warp_indices[0], warp_indices[1]] = shared[v0, v1] + s0 = T.dynamic("s0", "int32") + s1 = T.dynamic("s1", "int32") + @Ts.prim_func def ldmatrix_impl(warp_handle: T.handle, shared_handle: T.handle) -> None: - s0 = T.int32() - s1 = T.int32() shared = T.match_buffer( shared_handle, (smem_tile_row, smem_tile_col), @@ -620,12 +621,11 @@ def mma_store_desc(a: T.handle, c: T.handle) -> None: C[v0, v1] = C_warp[warp_indices[0], warp_indices[1]] if use_mma_store_intrinic: + s0 = T.dynamic("s0", "int32") + s1 = T.dynamic("s1", "int32") @Ts.prim_func def mma_store_impl(a: T.handle, c: T.handle) -> None: - s0 = T.int32() - s1 = T.int32() - C_warp = T.match_buffer( a, [WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1 ) @@ -651,12 +651,11 @@ def mma_store_impl(a: T.handle, c: T.handle) -> None: ) else: + s0 = T.dynamic("s0", "int32") + s1 = T.dynamic("s1", "int32") @Ts.prim_func def mma_store_impl(a: T.handle, c: T.handle) -> None: - s0 = T.int32() - s1 = T.int32() - C_warp = T.match_buffer( a, [WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1 ) @@ -857,12 +856,13 @@ def wmma_load_desc(a: T.handle, c: T.handle) -> None: vii, vjj = Ts.axis.remap("SS", [i, j]) C[vii, vjj] = A[vii, vjj] + s1 = T.dynamic("s1", "int32") + s0 = T.dynamic("s0", "int32") + d1 = T.dynamic("d1", "int32") + d0 = T.dynamic("d0", "int32") + @Ts.prim_func def wmma_load_impl(a: T.handle, c: T.handle) -> None: - s1 = T.int32() - s0 = T.int32() - d1 = T.int32() - d0 = T.int32() A = T.match_buffer( a, (frag_m, frag_n), @@ -926,10 +926,11 @@ def wmma_fill_desc(c: T.handle) -> None: vii, vjj = Ts.axis.remap("SS", [i, j]) C[vii, vjj] = zero + d1 = T.dynamic("d1", "int32") + d0 = T.dynamic("d0", "int32") + @Ts.prim_func def wmma_fill_impl(c: T.handle) -> None: - d1 = T.int32() - d0 = T.int32() C = T.match_buffer( c, (m_dim, n_dim), @@ -984,12 +985,13 @@ def wmma_store_desc(a: T.handle, c: T.handle) -> None: vii, vjj = Ts.axis.remap("SS", [i, j]) C[vii, vjj] = A[vii, vjj] + s1 = T.dynamic("s1", "int32") + s0 = T.dynamic("s0", "int32") + d1 = T.dynamic("d1", "int32") + d0 = T.dynamic("d0", "int32") + @Ts.prim_func def wmma_store_impl(a: T.handle, c: T.handle) -> None: - s1 = T.int32() - s0 = T.int32() - d1 = T.int32() - d0 = T.int32() A = T.match_buffer( a, (m_dim, n_dim), @@ -1087,15 +1089,15 @@ def wmma_sync_desc(a: T.handle, b: T.handle, c: T.handle) -> None: B[b_indices[0], b_indices[1]] ) + a1 = T.dynamic("a1", "int32") + a0 = T.dynamic("a0", "int32") + b1 = T.dynamic("b1", "int32") + b0 = T.dynamic("b0", "int32") + c1 = T.dynamic("c1", "int32") + c0 = T.dynamic("c0", "int32") + @Ts.prim_func def wmma_sync_impl(a: T.handle, b: T.handle, c: T.handle) -> None: - a1 = T.int32() - a0 = T.int32() - b1 = T.int32() - b0 = T.int32() - c1 = T.int32() - c0 = T.int32() - A = T.match_buffer( a, (m_dim, k_dim), @@ -1553,10 +1555,13 @@ def mma_load_desc(a: T.handle, c: T.handle) -> None: vi, vj = Ts.axis.remap("SS", [i, j]) dst[vi, vj] = src[vi, vj] + s0 = T.dynamic("s0", "int32") + s1 = T.dynamic("s1", "int32") + d0 = T.dynamic("d0", "int32") + d1 = T.dynamic("d1", "int32") + @Ts.prim_func def mma_load_impl(a: T.handle, c: T.handle) -> None: - s0 = T.int32() - s1 = T.int32() src = T.match_buffer( a, (frag_m, frag_n), @@ -1566,8 +1571,6 @@ def mma_load_impl(a: T.handle, c: T.handle) -> None: scope=shared_scope, strides=[s0, s1], ) - d0 = T.int32() - d1 = T.int32() dst = T.match_buffer( c, (frag_m, frag_n), @@ -1639,10 +1642,15 @@ def mma_sync_desc(a: T.handle, b: T.handle, c: T.handle) -> None: B[b_indices[0], b_indices[1]] ) + a0 = T.dynamic("a0", "int32") + a1 = T.dynamic("a1", "int32") + b0 = T.dynamic("b0", "int32") + b1 = T.dynamic("b1", "int32") + c0 = T.dynamic("c0", "int32") + c1 = T.dynamic("c1", "int32") + @Ts.prim_func def mma_sync_impl(a: T.handle, b: T.handle, c: T.handle) -> None: - a0 = T.int32() - a1 = T.int32() A = T.match_buffer( a, (m_dim, k_dim), @@ -1652,8 +1660,6 @@ def mma_sync_impl(a: T.handle, b: T.handle, c: T.handle) -> None: scope="m16n8k8.matrixA", strides=[a0, a1], ) - b0 = T.int32() - b1 = T.int32() B = T.match_buffer( b, (B_shape_0, B_shape_1), @@ -1663,8 +1669,6 @@ def mma_sync_impl(a: T.handle, b: T.handle, c: T.handle) -> None: scope="m16n8k8.matrixB", strides=[b0, b1], ) - c0 = T.int32() - c1 = T.int32() C = T.match_buffer( c, (m_dim, n_dim), diff --git a/python/tvm/s_tir/tensor_intrin/metal.py b/python/tvm/s_tir/tensor_intrin/metal.py index 5eb328fb1978..bb18cead6e2c 100644 --- a/python/tvm/s_tir/tensor_intrin/metal.py +++ b/python/tvm/s_tir/tensor_intrin/metal.py @@ -53,9 +53,11 @@ def desc(a: T.handle) -> None: vi, vj = Ts.axis.remap("SS", [i, j]) A[vi, vj] = T.float32(0) + d0 = T.dynamic("d0", "int32") + d1 = T.dynamic("d1", "int32") + @Ts.prim_func def impl(a: T.handle) -> None: - d0, d1 = T.int32(), T.int32() A = T.match_buffer( a, (col, row), dtype, scope="metal.simdgroup", strides=[d1, d0], offset_factor=1 ) @@ -100,9 +102,13 @@ def desc(a: T.handle, c: T.handle) -> None: else: C[vii, vjj] = A[vii, vjj] + s0 = T.dynamic("s0", "int32") + s1 = T.dynamic("s1", "int32") + d0 = T.dynamic("d0", "int32") + d1 = T.dynamic("d1", "int32") + @Ts.prim_func def impl(a: T.handle, c: T.handle) -> None: - s0, s1, d0, d1 = T.int32(), T.int32(), T.int32(), T.int32() A = T.match_buffer( a, (col, row), @@ -163,9 +169,13 @@ def desc(a: T.handle, c: T.handle) -> None: else: C[vii, vjj] = A[vii, vjj] + s0 = T.dynamic("s0", "int32") + s1 = T.dynamic("s1", "int32") + d0 = T.dynamic("d0", "int32") + d1 = T.dynamic("d1", "int32") + @Ts.prim_func def impl(a: T.handle, c: T.handle) -> None: - s0, s1, d0, d1 = T.int32(), T.int32(), T.int32(), T.int32() A = T.match_buffer( a, (col, row), @@ -210,9 +220,15 @@ def desc(a: T.handle, b: T.handle, c: T.handle) -> None: vii, vjj, vkk = Ts.axis.remap("SSR", [i, j, k]) C[vii, vjj] += A[vii, vkk] * B[vkk, vjj] + a0 = T.dynamic("a0", "int32") + a1 = T.dynamic("a1", "int32") + b0 = T.dynamic("b0", "int32") + b1 = T.dynamic("b1", "int32") + c0 = T.dynamic("c0", "int32") + c1 = T.dynamic("c1", "int32") + @Ts.prim_func def impl(a: T.handle, b: T.handle, c: T.handle) -> None: - a0, a1, b0, b1, c0, c1 = T.int32(), T.int32(), T.int32(), T.int32(), T.int32(), T.int32() A = T.match_buffer( a, (m_dim, k_dim), dtype, scope="metal.simdgroup", strides=[a1, a0], offset_factor=1 ) diff --git a/python/tvm/s_tir/tensor_intrin/rocm.py b/python/tvm/s_tir/tensor_intrin/rocm.py index 93e244b27d80..4d85d3b12f9b 100644 --- a/python/tvm/s_tir/tensor_intrin/rocm.py +++ b/python/tvm/s_tir/tensor_intrin/rocm.py @@ -226,11 +226,11 @@ def mfma_load_desc(reg_handle: T.handle, memory_handle: T.handle) -> None: Ts.writes(reg[warp_indices[0], warp_indices[1]]) reg[warp_indices[0], warp_indices[1]] = memory[v0, v1] + s0 = T.dynamic("s0", "int32") + s1 = T.dynamic("s1", "int32") + @Ts.prim_func def mfma_load_impl(reg_handle: T.handle, memory_handle: T.handle) -> None: - s0 = T.int32() - s1 = T.int32() - memory = T.match_buffer( memory_handle, memory_shape, @@ -407,11 +407,11 @@ def mfma_store_desc(a: T.handle, c: T.handle) -> None: Ts.writes(C[v0, v1]) C[v0, v1] = C_warp[warp_indices[0], warp_indices[1]] + s0 = T.dynamic("s0", "int32") + s1 = T.dynamic("s1", "int32") + @Ts.prim_func def mfma_store_impl(a: T.handle, c: T.handle) -> None: - s0 = T.int32() - s1 = T.int32() - C_warp = T.match_buffer( a, [WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1 ) diff --git a/tests/python/codegen/test_codegen_error_handling.py b/tests/python/codegen/test_codegen_error_handling.py index b9461b761dad..7b914760ea28 100644 --- a/tests/python/codegen/test_codegen_error_handling.py +++ b/tests/python/codegen/test_codegen_error_handling.py @@ -41,9 +41,10 @@ def test_wrong_argument_count_error(codegen_target): """Wrong argument count produces TypeError with function signature.""" + n0 = T.dynamic("n0") + @T.prim_func def func(a: T.handle, b: T.handle): - n0 = T.int64() A = T.match_buffer(a, (n0,), "float32") B = T.match_buffer(b, (n0,), "float32") for i in range(n0): @@ -70,9 +71,10 @@ def func(a: T.handle, b: T.handle): def test_type_mismatch_non_tensor(codegen_target): """Passing a non-tensor where a tensor is expected raises TypeError.""" + n0 = T.dynamic("n0") + @T.prim_func def func(a: T.handle, b: T.handle): - n0 = T.int64() A = T.match_buffer(a, (n0,), "float32") B = T.match_buffer(b, (n0,), "float32") for i in range(n0): @@ -100,9 +102,10 @@ def func(a: T.handle, b: T.handle): def test_shape_mismatch_shared_variable(codegen_target): """b has different shape than a when they share symbolic variable n0.""" + n0 = T.dynamic("n0") + @T.prim_func def func(a: T.handle, b: T.handle): - n0 = T.int64() A = T.match_buffer(a, (n0,), "float32") B = T.match_buffer(b, (n0,), "float32") for i in range(n0): @@ -393,9 +396,10 @@ def test_forward_reference_symbolic_shape(codegen_target): message uses rendered access paths (e.g. "B.shape[0] + 1") for shape checks. """ + batch_size = T.dynamic("batch_size") + @T.prim_func def func(a: T.handle, b: T.handle): - batch_size = T.int64() A = T.match_buffer(a, (batch_size + 1,), "int32") B = T.match_buffer(b, (batch_size,), "int32") for i in range(batch_size): diff --git a/tests/python/codegen/test_target_codegen_aarch64.py b/tests/python/codegen/test_target_codegen_aarch64.py index 90ec2617a6c4..bae6d704e2d2 100644 --- a/tests/python/codegen/test_target_codegen_aarch64.py +++ b/tests/python/codegen/test_target_codegen_aarch64.py @@ -39,12 +39,13 @@ def test_mul(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), dtype=dtype) @@ -74,12 +75,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_add(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), dtype=dtype) @@ -109,12 +111,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_sub(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), dtype=dtype) @@ -144,12 +147,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_muladd(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle, var_D: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), dtype=dtype) @@ -190,12 +194,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle, var_D: T.handle): def test_max(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), dtype=dtype) @@ -229,12 +234,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_min(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), dtype=dtype) @@ -268,12 +274,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_div(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), dtype=dtype) @@ -302,12 +309,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_mod(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), dtype=dtype) @@ -337,12 +345,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_eq(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), "bool") @@ -375,12 +384,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_neq(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), "bool") @@ -412,12 +422,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_or(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), dtype=dtype) @@ -446,12 +457,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_and(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), dtype=dtype) @@ -480,12 +492,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_not(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) C = T.match_buffer(var_C, (m,), dtype=dtype) for i in range(m): @@ -517,12 +530,13 @@ def main(var_A: T.handle, var_C: T.handle): def test_memcpy(dtype): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": ["+sve"]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,), dtype=dtype) B = T.match_buffer(var_B, (m,), "int32") C = T.match_buffer(var_C, (m,), dtype=dtype) @@ -557,12 +571,13 @@ def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): def test_vscale_range_function_attribute(mattr, expect_attr): target = {"kind": "llvm", "mtriple": "aarch64-linux-gnu", "mattr": [mattr]} + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (m,)) C = T.match_buffer(var_C, (m,)) for i in range(m): diff --git a/tests/python/codegen/test_target_codegen_arm.py b/tests/python/codegen/test_target_codegen_arm.py index ce3b00e64968..b11d76c694dd 100644 --- a/tests/python/codegen/test_target_codegen_arm.py +++ b/tests/python/codegen/test_target_codegen_arm.py @@ -62,12 +62,13 @@ def test_vmlal_s16(): } def check_correct_assembly(N): + K = T.dynamic("K", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, C: T.Buffer((N,), "int32")): T.func_attr({"tirx.noalias": True}) - K = T.int32() A = T.match_buffer(var_A, (K, N), "int8") B = T.match_buffer(var_B, (K, N), "int8") for n in T.vectorized(N): @@ -88,12 +89,13 @@ def main(var_A: T.handle, var_B: T.handle, C: T.Buffer((N,), "int32")): check_correct_assembly(64) def check_broadcast_correct_assembly(N): + K = T.dynamic("K", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, C: T.Buffer((N,), "int32")): T.func_attr({"tirx.noalias": True}) - K = T.int32() A = T.match_buffer(var_A, (K, N), "int8") B = T.match_buffer(var_B, (K,), "int8") for n in T.vectorized(N): diff --git a/tests/python/codegen/test_target_codegen_cuda.py b/tests/python/codegen/test_target_codegen_cuda.py index 5847042cf0d4..bca4ba00f4ba 100644 --- a/tests/python/codegen/test_target_codegen_cuda.py +++ b/tests/python/codegen/test_target_codegen_cuda.py @@ -379,12 +379,14 @@ def test_crossthread_reduction1(target): pytest.skip(f"{target} not enabled") def sched(nthd): + n = T.dynamic("n", "int32") + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle): T.func_attr({"tirx.noalias": True}) - n, m = T.int32(), T.int32() A = T.match_buffer(var_A, (n, m)) B = T.match_buffer(var_B, (n,)) for i in T.thread_binding(n, thread="blockIdx.x"): @@ -439,12 +441,15 @@ def test_crossthread_reduction2(target): pytest.skip(f"{target} not enabled") def sched(nthdx, nthdy): + n = T.dynamic("n", "int32") + k0 = T.dynamic("k0", "int32") + k1 = T.dynamic("k1", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle): T.func_attr({"tirx.noalias": True}) - n, k0, k1 = T.int32(), T.int32(), T.int32() A = T.match_buffer(var_A, (n, k0, k1)) B = T.match_buffer(var_B, (n,)) for i in T.thread_binding(n, thread="blockIdx.x"): diff --git a/tests/python/codegen/test_target_codegen_device.py b/tests/python/codegen/test_target_codegen_device.py index 4b3ed37ff97f..38e5f2e33a2d 100644 --- a/tests/python/codegen/test_target_codegen_device.py +++ b/tests/python/codegen/test_target_codegen_device.py @@ -60,12 +60,13 @@ def run_and_check(): @pytest.mark.gpu @pytest.mark.skipif(not env.has_gpu(), reason="need gpu") def test_add_pipeline(): + n = T.dynamic("n", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, B: T.Buffer((), "float32"), var_D: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() A = T.match_buffer(var_A, (n,)) D = T.match_buffer(var_D, (n,)) C = T.alloc_buffer((n,)) diff --git a/tests/python/codegen/test_target_codegen_llvm.py b/tests/python/codegen/test_target_codegen_llvm.py index ecf63314d553..9a8c610ff577 100644 --- a/tests/python/codegen/test_target_codegen_llvm.py +++ b/tests/python/codegen/test_target_codegen_llvm.py @@ -205,12 +205,13 @@ def main(A: T.Buffer((nn + base,), "float32"), C: T.Buffer((nn,), "float32")): @pytest.mark.skipif(not env.has_llvm(), reason="need llvm") def test_llvm_vadd_pipeline(): + n = T.dynamic("n", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() A = T.match_buffer(var_A, (n,)) B = T.match_buffer(var_B, (n,)) C = T.match_buffer(var_C, (n,)) @@ -284,26 +285,27 @@ def main(A: T.Buffer((1024,), "float32"), C: T.Buffer((1024,), "float32")): @pytest.mark.skipif(not env.has_llvm(), reason="need llvm") def test_multiple_func(): + fadd1_n = T.dynamic("n", "int32") + fadd2_n = T.dynamic("n", "int32") + @I.ir_module class Module: @T.prim_func def fadd1(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() - A = T.match_buffer(var_A, (n,)) - B = T.match_buffer(var_B, (n,)) - C = T.match_buffer(var_C, (n,)) - for i in range(n): + A = T.match_buffer(var_A, (fadd1_n,)) + B = T.match_buffer(var_B, (fadd1_n,)) + C = T.match_buffer(var_C, (fadd1_n,)) + for i in range(fadd1_n): C[i] = A[i] + B[i] @T.prim_func def fadd2(var_A: T.handle, var_B: T.handle, var_C: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() - A = T.match_buffer(var_A, (n,)) - B = T.match_buffer(var_B, (n,)) - C = T.match_buffer(var_C, (n,)) - for i in range(n): + A = T.match_buffer(var_A, (fadd2_n,)) + B = T.match_buffer(var_B, (fadd2_n,)) + C = T.match_buffer(var_C, (fadd2_n,)) + for i in range(fadd2_n): C[i] = A[i] + B[i] f = tvm.compile(Module, target="llvm") @@ -645,12 +647,13 @@ def _show_info(): @pytest.mark.skipif(not env.has_llvm(), reason="need llvm") def test_llvm_fp_math(): + n = T.dynamic("n", "int32") + @I.ir_module class RecipModule: @T.prim_func def main(var_A: T.handle, var_B: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() A = T.match_buffer(var_A, (n,)) B = T.match_buffer(var_B, (n,)) for i in range(n): @@ -664,12 +667,13 @@ def main(var_A: T.handle, var_B: T.handle): f_recip(a, b) tvm.testing.assert_allclose(b.numpy(), np.zeros((n,), "float32")) + n = T.dynamic("n", "int32") + @I.ir_module class SigmoidModule: @T.prim_func def main(var_A: T.handle, var_B: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() A = T.match_buffer(var_A, (n,)) B = T.match_buffer(var_B, (n,)) for i in range(n): diff --git a/tests/python/codegen/test_target_codegen_static_init.py b/tests/python/codegen/test_target_codegen_static_init.py index 008c601cf240..c7637adda97d 100644 --- a/tests/python/codegen/test_target_codegen_static_init.py +++ b/tests/python/codegen/test_target_codegen_static_init.py @@ -29,12 +29,13 @@ def test_cb(sh, A): assert isinstance(sh, ctypes.c_void_p) return sh + n = T.dynamic("n") + @I.ir_module class Module: @T.prim_func def ramp(A: T.handle): T.func_attr({"global_symbol": "ramp"}) - n = T.int64() Ab = T.match_buffer(A, (n,), "int64") T.call_packed( "test_static_callback", diff --git a/tests/python/codegen/test_target_codegen_vulkan.py b/tests/python/codegen/test_target_codegen_vulkan.py index 0294070c604e..f8664e1f7fad 100644 --- a/tests/python/codegen/test_target_codegen_vulkan.py +++ b/tests/python/codegen/test_target_codegen_vulkan.py @@ -459,6 +459,27 @@ def test_cooperative_matrix(out_dtype): M, N, K = 16, 16, 32 # fmt: off + C_0_s0 = T.dynamic("C_0_s0") + C_0_s1 = T.dynamic("C_0_s1") + A_3_s0 = T.dynamic("A_3_s0") + A_3_s1 = T.dynamic("A_3_s1") + C_4_s0 = T.dynamic("C_4_s0") + C_4_s1 = T.dynamic("C_4_s1") + A_2_s0 = T.dynamic("A_2_s0") + A_2_s1 = T.dynamic("A_2_s1") + B_0_s0 = T.dynamic("B_0_s0") + B_0_s1 = T.dynamic("B_0_s1") + C_3_s0 = T.dynamic("C_3_s0") + C_3_s1 = T.dynamic("C_3_s1") + A_0_s0 = T.dynamic("A_0_s0") + A_0_s1 = T.dynamic("A_0_s1") + C_1_s0 = T.dynamic("C_1_s0") + C_1_s1 = T.dynamic("C_1_s1") + A_1_s0 = T.dynamic("A_1_s0") + A_1_s1 = T.dynamic("A_1_s1") + C_2_s0 = T.dynamic("C_2_s0") + C_2_s1 = T.dynamic("C_2_s1") + @I.ir_module class Module: @T.prim_func @@ -581,11 +602,12 @@ def test_unary(): def run_test(tvm_intrin, np_func): n = 16 + m = T.dynamic("m", "int32") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle): - m = T.int32() A = T.match_buffer(var_A, (m,), "float32") B = T.match_buffer(var_B, (m,), "float32") for i_0 in T.thread_binding((m + 63) // 64, thread="blockIdx.x"): diff --git a/tests/python/contrib/test_tir_triton_integration.py b/tests/python/contrib/test_tir_triton_integration.py index cdc150b5a740..d3ce38861fed 100644 --- a/tests/python/contrib/test_tir_triton_integration.py +++ b/tests/python/contrib/test_tir_triton_integration.py @@ -64,34 +64,35 @@ def add_kernel( BLOCK_SIZE = 64 + add_m = T.dynamic("m") + main_m = T.dynamic("m") + @I.ir_module class Module: @Ts.prim_func def add(x_handle: T.handle, y_handle: T.handle, output_handle: T.handle) -> None: T.func_attr({"global_symbol": "add"}) - m = T.int64() - x = T.match_buffer(x_handle, (m,), "float32") - y = T.match_buffer(y_handle, (m,), "float32") - output = T.match_buffer(output_handle, (m,), "float32") + x = T.match_buffer(x_handle, (add_m,), "float32") + y = T.match_buffer(y_handle, (add_m,), "float32") + output = T.match_buffer(output_handle, (add_m,), "float32") with Ts.sblock("root"): - Ts.reads(x[0:m], y[0:m]) - Ts.writes(output[0:m]) + Ts.reads(x[0:add_m], y[0:add_m]) + Ts.writes(output[0:add_m]) T.call_kernel( add_kernel, - (T.ceildiv(m, BLOCK_SIZE),), + (T.ceildiv(add_m, BLOCK_SIZE),), x.data, y.data, output.data, - m, + add_m, BLOCK_SIZE, num_warps=8, ) @R.function - def main(x: R.Tensor(("m",), "float32"), y: R.Tensor(("m",), "float32")): - m = T.int64() + def main(x: R.Tensor((main_m,), "float32"), y: R.Tensor((main_m,), "float32")): with R.dataflow(): - output = R.call_tir(Module.add, [x, y], relax.TensorType((m,), "float32")) + output = R.call_tir(Module.add, [x, y], relax.TensorType((main_m,), "float32")) R.output(output) return output @@ -102,11 +103,12 @@ def main(x: R.Tensor(("m",), "float32"), y: R.Tensor(("m",), "float32")): scratch_args.append(tvm.tirx.reinterpret("handle", tvm.tirx.IntImm("uint64", 0))) # The thread extent is 256 because the kernel is compiled with num_warps=8. + m = T.dynamic("m") + @I.ir_module class Parsed: @Ts.prim_func def add(x_handle: T.handle, y_handle: T.handle, output_handle: T.handle): - m = T.int64() x = T.match_buffer(x_handle, (m,)) y = T.match_buffer(y_handle, (m,)) output = T.match_buffer(output_handle, (m,)) diff --git a/tests/python/relax/backend/adreno/mod_utils.py b/tests/python/relax/backend/adreno/mod_utils.py index 89650720dd47..4d749205a504 100644 --- a/tests/python/relax/backend/adreno/mod_utils.py +++ b/tests/python/relax/backend/adreno/mod_utils.py @@ -727,15 +727,16 @@ def get_global_maxpool_expected_codegen(input_shape, pool_size, stride, padding, def get_dequant_matmul_module(K, N): + seq_len = T.dynamic("seq_len") + @I.ir_module class DequantMatmul: @R.function def main( - input: R.Tensor((1, "seq_len", K), dtype="float16"), + input: R.Tensor((1, seq_len, K), dtype="float16"), weight: R.Tensor((K // 8, N), dtype="uint32"), scale: R.Tensor((K // 32, N), dtype="float16"), ): - seq_len = T.int64() cls = DequantMatmul with R.dataflow(): lv2 = relax.call_tir( @@ -785,23 +786,25 @@ def dequantize(weight: T.handle, scale: T.handle, var_dequantize: T.handle): def get_dequant_vec_matmul_module(K, N): + vocab_size_main = T.dynamic("vocab_size") + vocab_size_dequantize = T.dynamic("vocab_size") + @I.ir_module class DequantVecMatmul: @R.function def main( input: R.Tensor((1, 1, K), dtype="float16"), - weight: R.Tensor((K // 8, "vocab_size"), dtype="uint32"), - scale: R.Tensor((K // 32, "vocab_size"), dtype="float16"), + weight: R.Tensor((K // 8, vocab_size_main), dtype="uint32"), + scale: R.Tensor((K // 32, vocab_size_main), dtype="float16"), ): - vocab_size = T.int64() cls = DequantVecMatmul with R.dataflow(): lv2 = relax.call_tir( cls.dequantize, (weight, scale), - out_ty=R.Tensor((K, vocab_size), dtype="float16"), + out_ty=R.Tensor((K, vocab_size_main), dtype="float16"), ) - gv: R.Tensor((1, 1, vocab_size), dtype="float16") = relax.op.matmul( + gv: R.Tensor((1, 1, vocab_size_main), dtype="float16") = relax.op.matmul( input, lv2, out_dtype="float16" ) R.output(gv) @@ -810,13 +813,18 @@ def main( @Ts.prim_func def dequantize(weight: T.handle, scale: T.handle, var_dequantize: T.handle): T.func_attr({"tirx.noalias": T.bool(True)}) - vocab_size = T.int64() - lm_head_q_weight1 = T.match_buffer(weight, (T.int64(K // 8), vocab_size), "uint32") - lm_head_q_scale1 = T.match_buffer(scale, (T.int64(K // 32), vocab_size), "float16") - dequantize = T.match_buffer(var_dequantize, (T.int64(K), vocab_size), "float16") + lm_head_q_weight1 = T.match_buffer( + weight, (T.int64(K // 8), vocab_size_dequantize), "uint32" + ) + lm_head_q_scale1 = T.match_buffer( + scale, (T.int64(K // 32), vocab_size_dequantize), "float16" + ) + dequantize = T.match_buffer( + var_dequantize, (T.int64(K), vocab_size_dequantize), "float16" + ) # with Ts.sblock("root"): - compute = T.alloc_buffer((T.int64(K), vocab_size), "float16") - for i0, i1 in T.grid(T.int64(K), vocab_size): + compute = T.alloc_buffer((T.int64(K), vocab_size_dequantize), "float16") + for i0, i1 in T.grid(T.int64(K), vocab_size_dequantize): with Ts.sblock("compute"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(lm_head_q_weight1[v_i0 // T.int64(8), v_i1]) @@ -831,7 +839,7 @@ def dequantize(weight: T.handle, scale: T.handle, var_dequantize: T.handle): T.uint32(15), ), ) - for i0, i1 in T.grid(T.int64(K), vocab_size): + for i0, i1 in T.grid(T.int64(K), vocab_size_dequantize): with Ts.sblock("dequantize"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(compute[v_i0, v_i1], lm_head_q_scale1[v_i0 // T.int64(32), v_i1]) diff --git a/tests/python/relax/backend/adreno/test_transform_annotate_custom_scope.py b/tests/python/relax/backend/adreno/test_transform_annotate_custom_scope.py index 5b5d8bf12c95..f40fcdbe911a 100644 --- a/tests/python/relax/backend/adreno/test_transform_annotate_custom_scope.py +++ b/tests/python/relax/backend/adreno/test_transform_annotate_custom_scope.py @@ -195,6 +195,15 @@ def main( def _test_conv2d_symbolic_sub_indexed(): + N, H, W = T.dynamic("N"), T.dynamic("H"), T.dynamic("W") + Hw, Ww = T.dynamic("Hw"), T.dynamic("Ww") + + N = T.dynamic("N") + H = T.dynamic("H") + W = T.dynamic("W") + Hw = T.dynamic("Hw") + Ww = T.dynamic("Ww") + @I.ir_module class Input: @R.function @@ -202,13 +211,9 @@ def main(x: R.Tensor("float32", ndim=4), w: R.Tensor("float32", ndim=4)) -> R.Te "float32", ndim=4 ): with R.dataflow(): - N, C, H, W = T.int64(), I.meta_var(T.int64(16)), T.int64(), T.int64() - Nw, Cw, Hw, Ww = ( - I.meta_var(T.int64(4)), - I.meta_var(T.int64(16)), - T.int64(), - T.int64(), - ) + C = I.meta_var(T.int64(16)) + Nw = I.meta_var(T.int64(4)) + Cw = I.meta_var(T.int64(16)) lv0 = R.match_cast(x, R.Tensor((N, C, H, W), "float32")) lv1 = R.match_cast(w, R.Tensor((Nw, Cw, Hw, Ww), "float32")) gv: R.Tensor( diff --git a/tests/python/relax/distributed/test_distributed_transform_propagate_sharding.py b/tests/python/relax/distributed/test_distributed_transform_propagate_sharding.py index 3e2a3ac26c8b..35c243f2da13 100644 --- a/tests/python/relax/distributed/test_distributed_transform_propagate_sharding.py +++ b/tests/python/relax/distributed/test_distributed_transform_propagate_sharding.py @@ -235,6 +235,11 @@ def foo( def test_mlp_dynamic_shape(): + m = T.dynamic("m") + n = T.dynamic("n") + k0 = T.dynamic("k0") + k1 = T.dynamic("k1") + @I.ir_module class MLPDynamicShape: I.module_attrs({"device_num": 10}) @@ -249,16 +254,21 @@ class MLPDynamicShape: @R.function def foo( - x: R.Tensor(("m", "k0"), "float32"), - weight1: R.Tensor(("k0", "k1"), "float32"), - weight2: R.Tensor(("k1", "n"), "float32"), - ) -> R.Tensor(("m", "n"), "float32"): + x: R.Tensor((m, k0), "float32"), + weight1: R.Tensor((k0, k1), "float32"), + weight2: R.Tensor((k1, n), "float32"), + ) -> R.Tensor((m, n), "float32"): lv0 = R.matmul(x, weight1) lv1 = R.nn.gelu(lv0) lv2 = R.dist.annotate_sharding(lv1, device_mesh="mesh[0]", placement="S[1]") lv3 = R.matmul(lv2, weight2) return lv3 + m = T.dynamic("m") + n = T.dynamic("n") + k0 = T.dynamic("k0") + k1 = T.dynamic("k1") + @I.ir_module class ShardedMLPDynamicShape: I.module_attrs({"device_num": 10}) @@ -268,14 +278,10 @@ class ShardedMLPDynamicShape: @R.function def foo( - x: 'R.DTensor(("m", "k0"), "float32", "mesh[0]", "R")', - weight1: 'R.DTensor(("k0", "k1"), "float32", "mesh[0]", "S[1]")', - weight2: 'R.DTensor(("k1", "n"), "float32", "mesh[0]", "S[0]")', - ) -> 'R.DTensor(("m", "n"), "float32", "mesh[0]", "R")': - m = T.int64() - n = T.int64() - k0 = T.int64() - k1 = T.int64() + x: 'R.DTensor((m, k0), "float32", "mesh[0]", "R")', + weight1: 'R.DTensor((k0, k1), "float32", "mesh[0]", "S[1]")', + weight2: 'R.DTensor((k1, n), "float32", "mesh[0]", "S[0]")', + ) -> 'R.DTensor((m, n), "float32", "mesh[0]", "R")': lv0: R.DTensor((m, k1), "float32", "mesh[0]", "S[1]") = R.matmul(x, weight1) lv1: R.DTensor((m, k1), "float32", "mesh[0]", "S[1]") = R.nn.gelu(lv0) lv3: R.DTensor((m, n), "float32", "mesh[0]", "R") = R.matmul(lv1, weight2) @@ -1542,6 +1548,11 @@ def foo( def test_decoder_layer_dynamic_shape(): + rms_norm_n = T.dynamic("n") + rotary_embedding_n = T.dynamic("n") + foo_n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class LlamaAttentionLayerDynamicShape: I.module_attrs({"device_num": 10}) @@ -1559,12 +1570,13 @@ def rms_norm( var_A: T.handle, B: T.Buffer((T.int64(4096),), "float16"), var_rms_norm: T.handle ): T.func_attr({"tirx.noalias": True}) - n = T.int64() - A = T.match_buffer(var_A, (T.int64(1), n, T.int64(4096)), "float16") - rms_norm_1 = T.match_buffer(var_rms_norm, (T.int64(1), n, T.int64(4096)), "float16") + A = T.match_buffer(var_A, (T.int64(1), rms_norm_n, T.int64(4096)), "float16") + rms_norm_1 = T.match_buffer( + var_rms_norm, (T.int64(1), rms_norm_n, T.int64(4096)), "float16" + ) # with Ts.sblock("root"): - Ared_temp = Ts.sblock_alloc_buffer((T.int64(1), n)) - for bsz, i, k in T.grid(T.int64(1), n, T.int64(4096)): + Ared_temp = Ts.sblock_alloc_buffer((T.int64(1), rms_norm_n)) + for bsz, i, k in T.grid(T.int64(1), rms_norm_n, T.int64(4096)): with Ts.sblock("Ared_temp"): v_bsz, v_i, v_k = Ts.axis.remap("SSR", [bsz, i, k]) Ts.reads(A[v_bsz, v_i, v_k]) @@ -1574,7 +1586,7 @@ def rms_norm( Ared_temp[v_bsz, v_i] = Ared_temp[v_bsz, v_i] + T.Cast( "float32", A[v_bsz, v_i, v_k] ) * T.Cast("float32", A[v_bsz, v_i, v_k]) - for bsz, i, k in T.grid(T.int64(1), n, T.int64(4096)): + for bsz, i, k in T.grid(T.int64(1), rms_norm_n, T.int64(4096)): with Ts.sblock("rms_norm"): v_bsz, v_i, v_k = Ts.axis.remap("SSS", [bsz, i, k]) Ts.reads(B[v_k], A[v_bsz, v_i, v_k], Ared_temp[v_bsz, v_i]) @@ -1600,24 +1612,25 @@ def rotary_embedding( var_rotary: T.handle, ): T.func_attr({"tirx.noalias": True}) - n = T.int64() - A = T.match_buffer(var_A, (T.int64(1), n, T.int64(32), T.int64(128)), "float16") + A = T.match_buffer( + var_A, (T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)), "float16" + ) rotary = T.match_buffer( - var_rotary, (T.int64(1), n, T.int64(32), T.int64(128)), "float16" + var_rotary, (T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)), "float16" ) # with Ts.sblock("root"): - for i0, i1, i2, i3 in T.grid(T.int64(1), n, T.int64(32), T.int64(128)): + for i0, i1, i2, i3 in T.grid(T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)): with Ts.sblock("rotary"): v_i0, v_i1, v_i2, v_i3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads( - B[m + v_i1 - n, v_i3], + B[m + v_i1 - rotary_embedding_n, v_i3], A[v_i0, v_i1, v_i2, v_i3 - T.int64(64) : v_i3 - T.int64(64) + T.int64(129)], - C[m + v_i1 - n, v_i3], + C[m + v_i1 - rotary_embedding_n, v_i3], ) Ts.writes(rotary[v_i0, v_i1, v_i2, v_i3]) - rotary[v_i0, v_i1, v_i2, v_i3] = B[m + v_i1 - n, v_i3] * A[ + rotary[v_i0, v_i1, v_i2, v_i3] = B[m + v_i1 - rotary_embedding_n, v_i3] * A[ v_i0, v_i1, v_i2, v_i3 - ] + C[m + v_i1 - n, v_i3] * T.Select( + ] + C[m + v_i1 - rotary_embedding_n, v_i3] * T.Select( T.int64(64) <= v_i3, A[v_i0, v_i1, v_i2, v_i3 - T.int64(64)], A[v_i0, v_i1, v_i2, v_i3 + T.int64(64)] * T.float16(-1), @@ -1625,8 +1638,8 @@ def rotary_embedding( @R.function(pure=False) def foo( - input_tokens: R.Tensor((1, "n", 4096), dtype="float16"), - mask: R.Tensor((1, 1, "n", "m"), dtype="float16"), + input_tokens: R.Tensor((1, foo_n, 4096), dtype="float16"), + mask: R.Tensor((1, 1, foo_n, m), dtype="float16"), kv_cache: R.Tuple(R.Any, R.Any), linear_weight: R.Tensor((4096, 4096), dtype="float16"), linear_weight1: R.Tensor((4096, 4096), dtype="float16"), @@ -1636,21 +1649,19 @@ def foo( cos_cached: R.Tensor((2048, 128), dtype="float16"), sin_cached: R.Tensor((2048, 128), dtype="float16"), ): - n = T.int64() - m = T.int64() cls = LlamaAttentionLayerDynamicShape lv6 = R.call_tir( cls.rms_norm, (input_tokens, rms_norm_weight), - out_ty=R.Tensor((1, n, 4096), dtype="float16"), + out_ty=R.Tensor((1, foo_n, 4096), dtype="float16"), ) lv7: R.Tensor((4096, 4096), dtype="float16") = R.permute_dims(linear_weight, axes=None) lv7_copy: R.Tensor((4096, 4096), dtype="float16") = R.dist.annotate_sharding( lv7, "mesh[0]", "S[1]" ) - lv8: R.Tensor((1, n, 4096), dtype="float16") = R.matmul(lv6, lv7_copy) - lv9: R.Tensor((1, n, 32, 128), dtype="float16") = R.reshape( - lv8, R.shape([1, n, 32, 128]) + lv8: R.Tensor((1, foo_n, 4096), dtype="float16") = R.matmul(lv6, lv7_copy) + lv9: R.Tensor((1, foo_n, 32, 128), dtype="float16") = R.reshape( + lv8, R.shape([1, foo_n, 32, 128]) ) lv10: R.Tensor((4096, 4096), dtype="float16") = R.permute_dims( linear_weight1, axes=None @@ -1658,9 +1669,9 @@ def foo( lv10_copy: R.Tensor((4096, 4096), dtype="float16") = R.dist.annotate_sharding( lv10, "mesh[0]", "S[1]" ) - lv11: R.Tensor((1, n, 4096), dtype="float16") = R.matmul(lv6, lv10_copy) - lv12: R.Tensor((1, n, 32, 128), dtype="float16") = R.reshape( - lv11, R.shape([1, n, 32, 128]) + lv11: R.Tensor((1, foo_n, 4096), dtype="float16") = R.matmul(lv6, lv10_copy) + lv12: R.Tensor((1, foo_n, 32, 128), dtype="float16") = R.reshape( + lv11, R.shape([1, foo_n, 32, 128]) ) lv13: R.Tensor((4096, 4096), dtype="float16") = R.permute_dims( linear_weight2, axes=None @@ -1668,22 +1679,26 @@ def foo( lv13_copy: R.Tensor((4096, 4096), dtype="float16") = R.dist.annotate_sharding( lv13, "mesh[0]", "S[1]" ) - lv14: R.Tensor((1, n, 4096), dtype="float16") = R.matmul(lv6, lv13_copy) - lv15: R.Tensor((1, n, 32, 128), dtype="float16") = R.reshape( - lv14, R.shape([1, n, 32, 128]) + lv14: R.Tensor((1, foo_n, 4096), dtype="float16") = R.matmul(lv6, lv13_copy) + lv15: R.Tensor((1, foo_n, 32, 128), dtype="float16") = R.reshape( + lv14, R.shape([1, foo_n, 32, 128]) ) lv16 = R.call_tir( cls.rotary_embedding, (lv9, cos_cached, sin_cached, m), - out_ty=R.Tensor((1, n, 32, 128), dtype="float16"), + out_ty=R.Tensor((1, foo_n, 32, 128), dtype="float16"), ) lv17 = R.call_tir( cls.rotary_embedding, (lv12, cos_cached, sin_cached, m), - out_ty=R.Tensor((1, n, 32, 128), dtype="float16"), + out_ty=R.Tensor((1, foo_n, 32, 128), dtype="float16"), + ) + lv18: R.Tensor((foo_n, 32, 128), dtype="float16") = R.reshape( + lv17, R.shape([foo_n, 32, 128]) + ) + lv19: R.Tensor((foo_n, 32, 128), dtype="float16") = R.reshape( + lv15, R.shape([foo_n, 32, 128]) ) - lv18: R.Tensor((n, 32, 128), dtype="float16") = R.reshape(lv17, R.shape([n, 32, 128])) - lv19: R.Tensor((n, 32, 128), dtype="float16") = R.reshape(lv15, R.shape([n, 32, 128])) lv20: R.Any = kv_cache[0] lv21: R.Any = R.call_packed( "vm.builtin.attention_kv_cache_append", lv20, lv18, ty_args=(R.Any,) @@ -1710,7 +1725,7 @@ def foo( lv27: R.Tensor((1, m, 32, 128), dtype="float16") = R.reshape( lv25, R.shape([1, m, 32, 128]) ) - lv28: R.Tensor((1, 32, n, 128), dtype="float16") = R.permute_dims( + lv28: R.Tensor((1, 32, foo_n, 128), dtype="float16") = R.permute_dims( lv16, axes=[0, 2, 1, 3] ) lv29: R.Tensor((1, 32, m, 128), dtype="float16") = R.permute_dims( @@ -1722,31 +1737,38 @@ def foo( lv31: R.Tensor((1, 32, 128, m), dtype="float16") = R.permute_dims( lv29, axes=[0, 1, 3, 2] ) - lv32: R.Tensor((1, 32, n, m), dtype="float16") = R.matmul(lv28, lv31) - lv33: R.Tensor((1, 32, n, m), dtype="float16") = R.divide( + lv32: R.Tensor((1, 32, foo_n, m), dtype="float16") = R.matmul(lv28, lv31) + lv33: R.Tensor((1, 32, foo_n, m), dtype="float16") = R.divide( lv32, R.const(8, dtype="float16") ) # just choose some random value - lv34: R.Tensor((1, 32, n, m), dtype="float16") = R.maximum( + lv34: R.Tensor((1, 32, foo_n, m), dtype="float16") = R.maximum( lv33, R.const(1, dtype="float16") ) # just choose some random value - lv35: R.Tensor((1, 32, n, m), dtype="float16") = R.minimum(lv34, mask) + lv35: R.Tensor((1, 32, foo_n, m), dtype="float16") = R.minimum(lv34, mask) # lv36: R.Tensor((1, 32, n, m), dtype="float32") = R.astype(lv35, dtype="float32") - lv37: R.Tensor((1, 32, n, m), dtype="float16") = R.nn.softmax(lv35, axis=-1) + lv37: R.Tensor((1, 32, foo_n, m), dtype="float16") = R.nn.softmax(lv35, axis=-1) # lv38: R.Tensor((1, 32, n, m), dtype="float16") = R.astype(lv37, dtype="float16") - lv39: R.Tensor((1, 32, n, 128), dtype="float16") = R.matmul(lv37, lv30) - lv40: R.Tensor((1, n, 32, 128), dtype="float16") = R.permute_dims( + lv39: R.Tensor((1, 32, foo_n, 128), dtype="float16") = R.matmul(lv37, lv30) + lv40: R.Tensor((1, foo_n, 32, 128), dtype="float16") = R.permute_dims( lv39, axes=[0, 2, 1, 3] ) - lv41: R.Tensor((1, n, 4096), dtype="float16") = R.reshape(lv40, R.shape([1, n, 4096])) + lv41: R.Tensor((1, foo_n, 4096), dtype="float16") = R.reshape( + lv40, R.shape([1, foo_n, 4096]) + ) lv42: R.Tensor((4096, 4096), dtype="float16") = R.permute_dims( linear_weight3, axes=None ) - lv43: R.Tensor((1, n, 4096), dtype="float16") = R.matmul(lv41, lv42) - lv44: R.Tensor((1, n, 4096), dtype="float16") = R.add(input_tokens, lv43) + lv43: R.Tensor((1, foo_n, 4096), dtype="float16") = R.matmul(lv41, lv42) + lv44: R.Tensor((1, foo_n, 4096), dtype="float16") = R.add(input_tokens, lv43) gv = lv44 return gv + rms_norm_n = T.dynamic("n") + rotary_embedding_n = T.dynamic("n") + foo_n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class ShardedLlamaAttentionLayerDynamicShape: I.module_attrs({"device_num": 10}) @@ -1759,12 +1781,13 @@ def rms_norm( var_A: T.handle, B: T.Buffer((T.int64(4096),), "float16"), var_rms_norm: T.handle ): T.func_attr({"tirx.noalias": True}) - n = T.int64() - A = T.match_buffer(var_A, (T.int64(1), n, T.int64(4096)), "float16") - rms_norm_1 = T.match_buffer(var_rms_norm, (T.int64(1), n, T.int64(4096)), "float16") + A = T.match_buffer(var_A, (T.int64(1), rms_norm_n, T.int64(4096)), "float16") + rms_norm_1 = T.match_buffer( + var_rms_norm, (T.int64(1), rms_norm_n, T.int64(4096)), "float16" + ) # with Ts.sblock("root"): - Ared_temp = Ts.sblock_alloc_buffer((T.int64(1), n)) - for bsz, i, k in T.grid(T.int64(1), n, T.int64(4096)): + Ared_temp = Ts.sblock_alloc_buffer((T.int64(1), rms_norm_n)) + for bsz, i, k in T.grid(T.int64(1), rms_norm_n, T.int64(4096)): with Ts.sblock("Ared_temp"): v_bsz, v_i, v_k = Ts.axis.remap("SSR", [bsz, i, k]) Ts.reads(A[v_bsz, v_i, v_k]) @@ -1774,7 +1797,7 @@ def rms_norm( Ared_temp[v_bsz, v_i] = Ared_temp[v_bsz, v_i] + T.Cast( "float32", A[v_bsz, v_i, v_k] ) * T.Cast("float32", A[v_bsz, v_i, v_k]) - for bsz, i, k in T.grid(T.int64(1), n, T.int64(4096)): + for bsz, i, k in T.grid(T.int64(1), rms_norm_n, T.int64(4096)): with Ts.sblock("rms_norm"): v_bsz, v_i, v_k = Ts.axis.remap("SSS", [bsz, i, k]) Ts.reads(B[v_k], A[v_bsz, v_i, v_k], Ared_temp[v_bsz, v_i]) @@ -1800,24 +1823,25 @@ def rotary_embedding( var_rotary: T.handle, ): T.func_attr({"tirx.noalias": True}) - n = T.int64() - A = T.match_buffer(var_A, (T.int64(1), n, T.int64(32), T.int64(128)), "float16") + A = T.match_buffer( + var_A, (T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)), "float16" + ) rotary = T.match_buffer( - var_rotary, (T.int64(1), n, T.int64(32), T.int64(128)), "float16" + var_rotary, (T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)), "float16" ) # with Ts.sblock("root"): - for i0, i1, i2, i3 in T.grid(T.int64(1), n, T.int64(32), T.int64(128)): + for i0, i1, i2, i3 in T.grid(T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)): with Ts.sblock("rotary"): v_i0, v_i1, v_i2, v_i3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads( - B[m + v_i1 - n, v_i3], + B[m + v_i1 - rotary_embedding_n, v_i3], A[v_i0, v_i1, v_i2, v_i3 - T.int64(64) : v_i3 - T.int64(64) + T.int64(129)], - C[m + v_i1 - n, v_i3], + C[m + v_i1 - rotary_embedding_n, v_i3], ) Ts.writes(rotary[v_i0, v_i1, v_i2, v_i3]) - rotary[v_i0, v_i1, v_i2, v_i3] = B[m + v_i1 - n, v_i3] * A[ + rotary[v_i0, v_i1, v_i2, v_i3] = B[m + v_i1 - rotary_embedding_n, v_i3] * A[ v_i0, v_i1, v_i2, v_i3 - ] + C[m + v_i1 - n, v_i3] * T.Select( + ] + C[m + v_i1 - rotary_embedding_n, v_i3] * T.Select( T.int64(64) <= v_i3, A[v_i0, v_i1, v_i2, v_i3 - T.int64(64)], A[v_i0, v_i1, v_i2, v_i3 + T.int64(64)] * T.float16(-1), @@ -1825,8 +1849,8 @@ def rotary_embedding( @R.function(pure=False) def foo( - input_tokens: 'R.DTensor((1, "n", 4096), "float16", "mesh[0]", "R")', - mask: 'R.DTensor((1, 1, "n", "m"), "float16", "mesh[0]", "R")', + input_tokens: 'R.DTensor((1, foo_n, 4096), "float16", "mesh[0]", "R")', + mask: 'R.DTensor((1, 1, foo_n, m), "float16", "mesh[0]", "R")', kv_cache: R.Tuple(R.Any, R.Any), linear_weight: 'R.DTensor((4096, 4096), "float16", "mesh[0]", "S[0]")', linear_weight1: 'R.DTensor((4096, 4096), "float16", "mesh[0]", "S[0]")', @@ -1835,51 +1859,49 @@ def foo( rms_norm_weight: 'R.DTensor((4096,), "float16", "mesh[0]", "R")', cos_cached: 'R.DTensor((2048, 128), "float16", "mesh[0]", "R")', sin_cached: 'R.DTensor((2048, 128), "float16", "mesh[0]", "R")', - ) -> 'R.DTensor((1, "n", 4096), "float16", "mesh[0]", "R")': - n = T.int64() - m = T.int64() + ) -> 'R.DTensor((1, foo_n, 4096), "float16", "mesh[0]", "R")': cls = ShardedLlamaAttentionLayerDynamicShape lv6 = R.dist.call_tir( cls.rms_norm, (input_tokens, rms_norm_weight), - out_ty=R.DTensor((1, n, 4096), "float16", "mesh[0]", "R"), + out_ty=R.DTensor((1, foo_n, 4096), "float16", "mesh[0]", "R"), ) lv7: R.DTensor((4096, 4096), "float16", "mesh[0]", "S[1]") = R.permute_dims( linear_weight, axes=None ) - lv8: R.DTensor((1, n, 4096), "float16", "mesh[0]", "S[2]") = R.matmul(lv6, lv7) - lv9: R.DTensor((1, n, 32, 128), "float16", "mesh[0]", "S[2]") = R.reshape( - lv8, R.shape([1, n, 32, 128]) + lv8: R.DTensor((1, foo_n, 4096), "float16", "mesh[0]", "S[2]") = R.matmul(lv6, lv7) + lv9: R.DTensor((1, foo_n, 32, 128), "float16", "mesh[0]", "S[2]") = R.reshape( + lv8, R.shape([1, foo_n, 32, 128]) ) lv10: R.DTensor((4096, 4096), "float16", "mesh[0]", "S[1]") = R.permute_dims( linear_weight1, axes=None ) - lv11: R.DTensor((1, n, 4096), "float16", "mesh[0]", "S[2]") = R.matmul(lv6, lv10) - lv12: R.DTensor((1, n, 32, 128), "float16", "mesh[0]", "S[2]") = R.reshape( - lv11, R.shape([1, n, 32, 128]) + lv11: R.DTensor((1, foo_n, 4096), "float16", "mesh[0]", "S[2]") = R.matmul(lv6, lv10) + lv12: R.DTensor((1, foo_n, 32, 128), "float16", "mesh[0]", "S[2]") = R.reshape( + lv11, R.shape([1, foo_n, 32, 128]) ) lv13: R.DTensor((4096, 4096), "float16", "mesh[0]", "S[1]") = R.permute_dims( linear_weight2, axes=None ) - lv14: R.DTensor((1, n, 4096), "float16", "mesh[0]", "S[2]") = R.matmul(lv6, lv13) - lv15: R.DTensor((1, n, 32, 128), "float16", "mesh[0]", "S[2]") = R.reshape( - lv14, R.shape([1, n, 32, 128]) + lv14: R.DTensor((1, foo_n, 4096), "float16", "mesh[0]", "S[2]") = R.matmul(lv6, lv13) + lv15: R.DTensor((1, foo_n, 32, 128), "float16", "mesh[0]", "S[2]") = R.reshape( + lv14, R.shape([1, foo_n, 32, 128]) ) lv16 = R.dist.call_tir( cls.rotary_embedding, (lv9, cos_cached, sin_cached, m), - out_ty=R.DTensor((1, n, 32, 128), "float16", "mesh[0]", "S[2]"), + out_ty=R.DTensor((1, foo_n, 32, 128), "float16", "mesh[0]", "S[2]"), ) lv17 = R.dist.call_tir( cls.rotary_embedding, (lv12, cos_cached, sin_cached, m), - out_ty=R.DTensor((1, n, 32, 128), "float16", "mesh[0]", "S[2]"), + out_ty=R.DTensor((1, foo_n, 32, 128), "float16", "mesh[0]", "S[2]"), ) - lv18: R.DTensor((n, 32, 128), "float16", "mesh[0]", "S[1]") = R.reshape( - lv17, R.shape([n, 32, 128]) + lv18: R.DTensor((foo_n, 32, 128), "float16", "mesh[0]", "S[1]") = R.reshape( + lv17, R.shape([foo_n, 32, 128]) ) - lv19: R.DTensor((n, 32, 128), "float16", "mesh[0]", "S[1]") = R.reshape( - lv15, R.shape([n, 32, 128]) + lv19: R.DTensor((foo_n, 32, 128), "float16", "mesh[0]", "S[1]") = R.reshape( + lv15, R.shape([foo_n, 32, 128]) ) lv20: R.Any = kv_cache[0] lv21: R.Any = R.call_packed( @@ -1913,7 +1935,7 @@ def foo( lv27: R.DTensor((1, m, 32, 128), "float16", "mesh[0]", "S[2]") = R.reshape( lv25, R.shape([1, m, 32, 128]) ) - lv28: R.DTensor((1, 32, n, 128), "float16", "mesh[0]", "S[1]") = R.permute_dims( + lv28: R.DTensor((1, 32, foo_n, 128), "float16", "mesh[0]", "S[1]") = R.permute_dims( lv16, axes=[0, 2, 1, 3] ) lv29: R.DTensor((1, 32, m, 128), "float16", "mesh[0]", "S[1]") = R.permute_dims( @@ -1925,30 +1947,32 @@ def foo( lv31: R.DTensor((1, 32, 128, m), "float16", "mesh[0]", "S[1]") = R.permute_dims( lv29, axes=[0, 1, 3, 2] ) - lv32: R.DTensor((1, 32, n, m), "float16", "mesh[0]", "S[1]") = R.matmul(lv28, lv31) - lv33: R.DTensor((1, 32, n, m), "float16", "mesh[0]", "S[1]") = R.divide( + lv32: R.DTensor((1, 32, foo_n, m), "float16", "mesh[0]", "S[1]") = R.matmul(lv28, lv31) + lv33: R.DTensor((1, 32, foo_n, m), "float16", "mesh[0]", "S[1]") = R.divide( lv32, R.dist.const(8, R.DTensor((), "float16", "mesh[0]", "R")) ) - lv34: R.DTensor((1, 32, n, m), "float16", "mesh[0]", "S[1]") = R.maximum( + lv34: R.DTensor((1, 32, foo_n, m), "float16", "mesh[0]", "S[1]") = R.maximum( lv33, R.dist.const(1, R.DTensor((), "float16", "mesh[0]", "R")) ) - lv35: R.DTensor((1, 32, n, m), "float16", "mesh[0]", "S[1]") = R.minimum(lv34, mask) - lv37: R.DTensor((1, 32, n, m), "float16", "mesh[0]", "S[1]") = R.nn.softmax( + lv35: R.DTensor((1, 32, foo_n, m), "float16", "mesh[0]", "S[1]") = R.minimum(lv34, mask) + lv37: R.DTensor((1, 32, foo_n, m), "float16", "mesh[0]", "S[1]") = R.nn.softmax( lv35, axis=-1 ) - lv39: R.DTensor((1, 32, n, 128), "float16", "mesh[0]", "S[1]") = R.matmul(lv37, lv30) - lv40: R.DTensor((1, n, 32, 128), "float16", "mesh[0]", "S[2]") = R.permute_dims( + lv39: R.DTensor((1, 32, foo_n, 128), "float16", "mesh[0]", "S[1]") = R.matmul( + lv37, lv30 + ) + lv40: R.DTensor((1, foo_n, 32, 128), "float16", "mesh[0]", "S[2]") = R.permute_dims( lv39, axes=[0, 2, 1, 3] ) - lv41: R.DTensor((1, n, 4096), "float16", "mesh[0]", "S[2]") = R.reshape( - lv40, R.shape([1, n, 4096]) + lv41: R.DTensor((1, foo_n, 4096), "float16", "mesh[0]", "S[2]") = R.reshape( + lv40, R.shape([1, foo_n, 4096]) ) lv42: R.DTensor((4096, 4096), "float16", "mesh[0]", "S[0]") = R.permute_dims( linear_weight3, axes=None ) - lv43: R.DTensor((1, n, 4096), "float16", "mesh[0]", "R") = R.matmul(lv41, lv42) - lv44: R.DTensor((1, n, 4096), "float16", "mesh[0]", "R") = R.add(input_tokens, lv43) - gv: R.DTensor((1, n, 4096), "float16", "mesh[0]", "R") = lv44 + lv43: R.DTensor((1, foo_n, 4096), "float16", "mesh[0]", "R") = R.matmul(lv41, lv42) + lv44: R.DTensor((1, foo_n, 4096), "float16", "mesh[0]", "R") = R.add(input_tokens, lv43) + gv: R.DTensor((1, foo_n, 4096), "float16", "mesh[0]", "R") = lv44 return gv after = relax.distributed.transform.PropagateSharding()(LlamaAttentionLayerDynamicShape) diff --git a/tests/python/relax/test_analysis.py b/tests/python/relax/test_analysis.py index 48200548a06b..35207ed64557 100644 --- a/tests/python/relax/test_analysis.py +++ b/tests/python/relax/test_analysis.py @@ -22,7 +22,6 @@ import tvm import tvm.testing from tvm import relax as rx -from tvm import tirx from tvm.relax.analysis import ( all_global_vars, all_vars, @@ -45,8 +44,8 @@ def var_name_set(vars: list[rx.Var | rx.GlobalVar]) -> set[str]: def test_use_def(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = rx.Var("x", R.Tensor([m, n], "float16")) y = rx.Var("y", R.Tensor([n], "float16")) ib = rx.BlockBuilder() @@ -75,8 +74,8 @@ def test_use_def(): ids=["binary_op", "self_reference", "tuple"], ) def test_used_vars(expr_fn, expected_var_names): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = rx.Var("x", R.Tensor([m, n], "float16")) y = rx.Var("y", R.Tensor([n], "float16")) z = rx.Var("z", R.Tensor([m], "float16")) @@ -293,13 +292,14 @@ def main(x: R.Tensor((32, 32), "float32")) -> R.Tensor((32, 32), "float32"): def test_edge_binding_block_fake_unused_remove_all_unused2(): + m = T.dynamic("m") + n = T.dynamic("n") + k = T.dynamic("k") + @tvm.script.ir_module class IdentityUnused: @R.function def main(x: R.Tensor((3,), dtype="int64")) -> R.Tensor(dtype="int32", ndim=3): - m = T.int64() - n = T.int64() - k = T.int64() with R.dataflow(): lv: R.Shape(ndim=3) = R.call_pure_packed( "vm.builtin.tensor_to_shape", x, ty_args=(R.Shape(ndim=3),) @@ -380,6 +380,8 @@ def expected(x: R.Tensor((32, 32), "float32")) -> R.Tensor: def test_retain_calls_to_impure_builtin_ops(): + n = T.dynamic("n") + @I.ir_module class Module: @Ts.prim_func(private=True) @@ -387,9 +389,8 @@ def my_tir(A: T.handle, B: T.handle, n: T.int64): T.evaluate(0) @R.function(pure=False) - def main(x: R.Tensor(("n",), "float32")): + def main(x: R.Tensor((n,), "float32")): cls = Module - n = T.int64() storage = R.memory.alloc_storage((n * 4,), 0, "global", "float32") alloc = R.memory.alloc_tensor(storage, R.prim_value(0), R.shape([n]), "float32") # "call_tir_dyn" is impure which shouldn't be removed. @@ -522,7 +523,7 @@ def test_free_vars(): @pytest.mark.parametrize("definition_site", ["parameter", "match_cast"]) def test_free_vars_primitive_definition_sites(definition_site): - n = tirx.Var("n", "int64") + n = T.dynamic("n", "int64") if definition_site == "parameter": y = rx.Var("y", rx.TensorType([n], "float32")) else: @@ -650,9 +651,10 @@ def expand_dims( def test_reshape_pattern_dyn_1(): + n = T.dynamic("n") + @Ts.prim_func def reshape(var_A: T.handle, var_T_reshape: T.handle): - n = T.int64() A = T.match_buffer(var_A, (n, T.int64(32), T.int64(128)), "float16") T_reshape = T.match_buffer( var_T_reshape, (T.int64(1), n, T.int64(32), T.int64(128)), "float16" @@ -678,9 +680,10 @@ def reshape(var_A: T.handle, var_T_reshape: T.handle): def test_reshape_pattern_dyn_2(): + n = T.dynamic("n") + @Ts.prim_func def reshape(var_A: T.handle, var_T_reshape: T.handle): - n = T.int64() A = T.match_buffer(var_A, (T.int64(1), n), "int32") T_reshape = T.match_buffer(var_T_reshape, (n,), "int32") for ax0 in range(n): @@ -694,10 +697,11 @@ def reshape(var_A: T.handle, var_T_reshape: T.handle): def test_reshape_pattern_dyn_3(): + n = T.dynamic("n") + @Ts.prim_func def reshape(var_A: T.handle, var_T_reshape: T.handle): T.func_attr({"op_pattern": 8, "tirx.noalias": True}) - n = T.int64() A = T.match_buffer(var_A, (n, T.int64(4096)), "float16") T_reshape = T.match_buffer(var_T_reshape, (T.int64(1), n, T.int64(4096)), "float16") for ax0, ax1, ax2 in T.grid(T.int64(1), n, T.int64(4096)): @@ -713,10 +717,11 @@ def reshape(var_A: T.handle, var_T_reshape: T.handle): def test_reshape_pattern_dyn_4(): + n = T.dynamic("n") + @Ts.prim_func def reshape(var_A: T.handle, var_T_reshape: T.handle): T.func_attr({"op_pattern": 8, "tirx.noalias": True}) - n = T.int64() A = T.match_buffer(var_A, (T.int64(1), n, T.int64(4096)), "float16") T_reshape = T.match_buffer( var_T_reshape, (T.int64(1), n, T.int64(32), T.int64(128)), "float16" @@ -742,10 +747,11 @@ def reshape(var_A: T.handle, var_T_reshape: T.handle): def test_reshape_pattern_dyn_5(): + n = T.dynamic("n") + @Ts.prim_func def reshape(var_A: T.handle, var_T_reshape: T.handle): T.func_attr({"op_pattern": 8, "tirx.noalias": True}) - n = T.int64() A = T.match_buffer(var_A, (T.int64(1), n, T.int64(32), T.int64(128)), "float16") T_reshape = T.match_buffer(var_T_reshape, (T.int64(1), n, T.int64(4096)), "float16") # with Ts.sblock("root"): diff --git a/tests/python/relax/test_analysis_computable_at_compile_time.py b/tests/python/relax/test_analysis_computable_at_compile_time.py index a438a8dbbcc2..c28c50d8beec 100644 --- a/tests/python/relax/test_analysis_computable_at_compile_time.py +++ b/tests/python/relax/test_analysis_computable_at_compile_time.py @@ -134,10 +134,11 @@ def func(A: R.Tensor([16], "int32"), B: R.Tensor([16], "int32")): def test_compile_time_symbolic_shape(): """Compile-time bindings may contain symbolic shapes""" + n = T.dynamic("n") + @R.function - def func(A: R.Tensor([1], "int32"), B: R.Tensor(["n"], "int32")): + def func(A: R.Tensor([1], "int32"), B: R.Tensor([n], "int32")): R.func_attr({"num_input": 1}) - n = T.int64() C: R.Tensor([n], "int32") = R.add(B, B) D: R.Tensor([], "int32") = R.max(C, axis=0) @@ -150,11 +151,12 @@ def func(A: R.Tensor([1], "int32"), B: R.Tensor(["n"], "int32")): def test_symbolic_variables_from_match_binding(): """Symbolic vars may be inferred from compile-time bindings""" + n = T.dynamic("n") + m = T.dynamic("m") + @R.function def func(A: R.Tensor(ndim=1, dtype="int32"), B: R.Tensor(ndim=1, dtype="int32")): R.func_attr({"num_input": 1}) - n = T.int64() - m = T.int64() A2 = R.match_cast(A, R.Tensor([n], "int32")) B2 = R.match_cast(B, R.Tensor([m], "int32")) @@ -177,11 +179,12 @@ def test_compile_time_expressions_may_not_use_runtime_symbolic_variables(): first knowing `A`, and is therefore unknown at compile-time. """ + n = T.dynamic("n") + m = T.dynamic("m") + @R.function - def func(A: R.Tensor(["n"], "int32"), B: R.Tensor(["m"], "int32")): + def func(A: R.Tensor([n], "int32"), B: R.Tensor([m], "int32")): R.func_attr({"num_input": 1}) - n = T.int64() - m = T.int64() C = R.ones([m], "int32") D = R.ones([n], "int32") @@ -200,10 +203,11 @@ def test_compile_time_expressions_may_infer_same_variable_as_run_time(): can also be inferred from the compile-time parameter `B`. """ + n = T.dynamic("n") + @R.function - def func(A: R.Tensor(["n"], "int32"), B: R.Tensor(["n"], "int32")): + def func(A: R.Tensor([n], "int32"), B: R.Tensor([n], "int32")): R.func_attr({"num_input": 1}) - n = T.int64() C = R.ones([n], "int32") D = R.ones([n], "int32") @@ -223,11 +227,12 @@ def test_compile_time_expressions_may_use_variables_from_match_cast(): first knowing `A`, and is therefore unknown at compile-time. """ + n = T.dynamic("n") + m = T.dynamic("m") + @R.function - def func(A: R.Tensor(["n"], "int32"), B: R.Tensor(ndim=1, dtype="int32")): + def func(A: R.Tensor([n], "int32"), B: R.Tensor(ndim=1, dtype="int32")): R.func_attr({"num_input": 1}) - n = T.int64() - m = T.int64() B2 = R.match_cast(B, R.Tensor([m], "int32")) diff --git a/tests/python/relax/test_analysis_suggest_layout_transforms.py b/tests/python/relax/test_analysis_suggest_layout_transforms.py index 8a61f40144cd..dcf7a4672bf4 100644 --- a/tests/python/relax/test_analysis_suggest_layout_transforms.py +++ b/tests/python/relax/test_analysis_suggest_layout_transforms.py @@ -258,12 +258,13 @@ def expected( def test_op_elemwise_symbolic(): + N = T.dynamic("N") + C = T.dynamic("C") + H = T.dynamic("H") + W = T.dynamic("W") + @Ts.prim_func(private=True) def before(arg: T.handle, relu: T.handle): - N = T.int64() - C = T.int64() - H = T.int64() - W = T.int64() Arg = T.match_buffer(arg, (N, C, H, W)) Relu = T.match_buffer(relu, (N, C, H, W)) for i0, i1, i2, i3 in T.grid(N, C, H, W): @@ -273,12 +274,13 @@ def before(arg: T.handle, relu: T.handle): Ts.writes(Relu[v_i0, v_i1, v_i2, v_i3]) Relu[v_i0, v_i1, v_i2, v_i3] = T.max(Arg[v_i0, v_i1, v_i2, v_i3], T.float32(0)) + N = T.dynamic("N") + C = T.dynamic("C") + H = T.dynamic("H") + W = T.dynamic("W") + @Ts.prim_func(private=True) def expected(arg: T.handle, relu: T.handle): - N = T.int64() - C = T.int64() - H = T.int64() - W = T.int64() Arg = T.match_buffer(arg, (N, H, W, C)) Relu = T.match_buffer(relu, (N, H, W, C)) # with Ts.sblock("root"): diff --git a/tests/python/relax/test_analysis_type_analysis.py b/tests/python/relax/test_analysis_type_analysis.py index 29e4728e7fab..de64b619a2b3 100644 --- a/tests/python/relax/test_analysis_type_analysis.py +++ b/tests/python/relax/test_analysis_type_analysis.py @@ -755,11 +755,14 @@ def test_collect_symbolic_var_from_non_tensor_params(param_type, param_order): def test_collect_nonnegative_expressions(): + M = T.dynamic("M") + N = T.dynamic("N") + @R.function def func( - A: R.Tensor([1024, "M", "N-2"]), - B: R.Tensor([128, "N", "M+2"]), - C: R.Shape(["M", "N"]), + A: R.Tensor([1024, M, N - 2]), + B: R.Tensor([128, N, M + 2]), + C: R.Shape([M, N]), D: T.int64, ): return R.tuple() diff --git a/tests/python/relax/test_analysis_well_formed.py b/tests/python/relax/test_analysis_well_formed.py index 003ca272993f..8d928c3fef68 100644 --- a/tests/python/relax/test_analysis_well_formed.py +++ b/tests/python/relax/test_analysis_well_formed.py @@ -27,8 +27,8 @@ from tvm.script import s_tir as Ts from tvm.script import tirx as T -m = tirx.Var("m", "int64") -n = tirx.Var("n", "int64") +m = T.dynamic("m", "int64") +n = T.dynamic("n", "int64") x = rx.Var("x", R.Tensor([m, n], "float32")) cond = rx.Var("cond", R.Tensor([], "bool")) @@ -542,10 +542,11 @@ def test_ty_args_tir_var_used_before_define_call_tir(): def test_ty_erase_to_well_formed(): # Error: The return ty contains undefined symbolic vars """ + m, n = T.dynamic("m"), T.dynamic("n") + m1, n1 = T.dynamic("m1"), T.dynamic("n1") + @R.function - def foo(x: R.Tensor(("m", "n"), dtype="float32")) -> R.Tensor(("m1", "n1"), dtype="float32"): - m = T.int64() - n = T.int64() + def foo(x: R.Tensor((m, n), dtype="float32")) -> R.Tensor((m1, n1), dtype="float32"): gv = R.call_dps_packed("my_func", (x,), out_ty=R.Tensor((m, n), dtype="float32")) return gv """ @@ -565,7 +566,7 @@ def test_func_ty_well_formed(): @R.function def foo(): @R.function - def local(x: R.Tensor(["m", "n"], "float32")): + def local(x: R.Tensor([m, n], "float32")): return x return local @@ -981,6 +982,9 @@ def test_call_tir_with_correct_dynamic_output_shape(): """ + M = T.dynamic("M") + N = T.dynamic("N") + @I.ir_module class Module: @R.function @@ -990,8 +994,6 @@ def main(A: R.Tensor([16], "float16")): @Ts.prim_func def reshape(A: T.Buffer(16, "float16"), B_handle: T.handle): - M = T.int64() - N = T.int64() B = T.match_buffer(B_handle, [M, N], dtype="float16") for i, j in T.grid(M, N): @@ -1014,6 +1016,9 @@ def test_call_tir_with_incorrect_dynamic_output_shape(): """ + M = T.dynamic("M") + N = T.dynamic("N") + @I.ir_module(check_well_formed=False) class Module: @R.function @@ -1023,8 +1028,6 @@ def main(A: R.Tensor([16], "float16")): @Ts.prim_func def reshape(A: T.Buffer(16, "float16"), B_handle: T.handle): - M = T.int64() - N = T.int64() B = T.match_buffer(B_handle, [M, N], dtype="float16") for i, j in T.grid(M, N): @@ -1049,6 +1052,9 @@ def test_call_tir_incorrect_dimensionality_of_output_shape(): """ + M = T.dynamic("M") + N = T.dynamic("N") + @I.ir_module(check_well_formed=False) class Module: @R.function @@ -1058,8 +1064,6 @@ def main(A: R.Tensor([16], "float16")): @Ts.prim_func def reshape(A: T.Buffer(16, "float16"), B_handle: T.handle): - M = T.int64() - N = T.int64() B = T.match_buffer(B_handle, [M, N], dtype="float16") for i, j in T.grid(M, N): @@ -1087,6 +1091,9 @@ def test_call_tir_output_shape_with_mixed_static_and_dynamic(): """ + M = T.dynamic("M") + N = T.dynamic("N") + @I.ir_module(check_well_formed=False) class Module: @R.function @@ -1096,8 +1103,6 @@ def main(A: R.Tensor([256], "float16")): @Ts.prim_func def reshape(A: T.Buffer(256, "float16"), B_handle: T.handle): - M = T.int64() - N = T.int64() B = T.match_buffer(B_handle, [16, M, N], dtype="float16") for i, j, k in T.grid(16, M, N): @@ -1119,6 +1124,9 @@ def test_call_tir_with_correct_inferred_dynamic_output_shape(): """ + M = T.dynamic("M") + N = T.dynamic("N") + @I.ir_module class Module: @R.function @@ -1128,8 +1136,6 @@ def main(A: R.Tensor([8, 4], "float16")): @Ts.prim_func def flatten(A_handle: T.handle, B_handle: T.handle): - M = T.int64() - N = T.int64() A = T.match_buffer(A_handle, [M, N], dtype="float16") B = T.match_buffer(B_handle, [M * N], dtype="float16") @@ -1157,6 +1163,9 @@ def test_call_tir_with_incorrect_inferred_dynamic_output_shape(): """ + M = T.dynamic("M") + N = T.dynamic("N") + @I.ir_module(check_well_formed=False) class Module: @R.function @@ -1166,8 +1175,6 @@ def main(A: R.Tensor([8, 4], "float16")): @Ts.prim_func def flatten(A_handle: T.handle, B_handle: T.handle): - M = T.int64() - N = T.int64() A = T.match_buffer(A_handle, [M, N], dtype="float16") B = T.match_buffer(B_handle, [M * N], dtype="float16") @@ -1191,6 +1198,9 @@ def test_call_tir_with_dtensor_arguments(): # from tvm.script.parser import relax as R + M = T.dynamic("M") + N = T.dynamic("N") + @I.ir_module class Module: I.module_attrs({"device_num": 4}) @@ -1205,8 +1215,6 @@ def main(A: 'R.dist.DTensor([8, 4], "float16", "mesh[0]", "S[0]")'): @Ts.prim_func def flatten(A_handle: T.handle, B_handle: T.handle): - M = T.int64() - N = T.int64() A = T.match_buffer(A_handle, [M, N], dtype="float16") B = T.match_buffer(B_handle, [M * N], dtype="float16") diff --git a/tests/python/relax/test_ast_printer.py b/tests/python/relax/test_ast_printer.py index 108467e670d1..9399624bcc93 100644 --- a/tests/python/relax/test_ast_printer.py +++ b/tests/python/relax/test_ast_printer.py @@ -105,8 +105,8 @@ def test_dataflow_var() -> None: def test_match_cast() -> None: # match_cast([16, 8], [m, n]) - m = tirx.Var("m", ty="int64") - n = tirx.Var("n", ty="int64") + m = T.dynamic("m", dtype="int64") + n = T.dynamic("n", dtype="int64") shape = rx.const([16, 8], "int32") var = rx.Var("v0", R.Shape()) b0 = rx.MatchCast(var, shape, R.Tensor([m, n], "int32")) @@ -142,8 +142,8 @@ def test_var_binding() -> None: def test_binding_block() -> None: - m = tirx.Var("m", ty="int64") - n = tirx.Var("n", ty="int64") + m = T.dynamic("m", dtype="int64") + n = T.dynamic("n", dtype="int64") shape = rx.const([16, 8], "int32") b0 = rx.MatchCast(rx.Var("v0"), shape, R.Tensor([m, n], "int32")) @@ -161,8 +161,8 @@ def test_binding_block() -> None: def test_dataflow_block() -> None: - m = tirx.Var("m", ty="int64") - n = tirx.Var("n", ty="int64") + m = T.dynamic("m", dtype="int64") + n = T.dynamic("n", dtype="int64") shape = rx.const([16, 8], "int32") b0 = rx.MatchCast(rx.Var("v0"), shape, R.Tensor([m, n], "int32")) @@ -196,8 +196,8 @@ def test_seq_expr() -> None: def test_shape_expr() -> None: - m = tirx.Var("m", ty="int32") - n = tirx.Var("n", ty="int32") + m = T.dynamic("m", dtype="int32") + n = T.dynamic("n", dtype="int32") s = rx.ShapeExpr([m, n]) s_str = dump_ast(s) assert s_str.startswith("ShapeExpr(") @@ -350,13 +350,14 @@ def test_ty(): def test_call_packed(): # test case from test_parser + m = T.dynamic("m") + @R.function(pure=False) def f( - x: R.Tensor((32, "m"), "float32"), - y: R.Tensor(("m",), "float32"), + x: R.Tensor((32, m), "float32"), + y: R.Tensor((m,), "float32"), r: R.Tensor(dtype="int64"), ) -> R.Any: - 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) @@ -438,24 +439,26 @@ def test_op_attrs(): def test_call_tir(): # also from test_parser + m_addone = T.dynamic("m") + n_addone = T.dynamic("n") + m_foo = T.dynamic("m") + n_foo = T.dynamic("n") + @tvm.script.ir_module class TestCallTIR: @Ts.prim_func def addone(A_handle: T.handle, B_handle: T.handle) -> None: - m = T.int64() - n = T.int64() - A = T.match_buffer(A_handle, (m, n), "float32") - B = T.match_buffer(B_handle, (m, n), "float32") + A = T.match_buffer(A_handle, (m_addone, n_addone), "float32") + B = T.match_buffer(B_handle, (m_addone, n_addone), "float32") T.func_attr({"global_symbol": "addone"}) - for i, j in T.grid(m, n): + for i, j in T.grid(m_addone, n_addone): with Ts.sblock("addone"): vi, vj = Ts.axis.remap("SS", [i, j]) B[vi, vj] = A[vi, vj] + T.int32(1) @R.function - def foo(x: R.Tensor(("m", "n"), "float32")): - m, n = T.int64(), T.int64() - gv0 = R.call_tir(TestCallTIR.addone, (x,), R.Tensor((m, n), dtype="float32")) + def foo(x: R.Tensor((m_foo, n_foo), "float32")): + gv0 = R.call_tir(TestCallTIR.addone, (x,), R.Tensor((m_foo, n_foo), dtype="float32")) return gv0 mod = TestCallTIR @@ -504,9 +507,11 @@ def foo(x: R.Tensor(("m", "n"), "float32")): def test_call_dps_packed(): + m = T.dynamic("m") + n = T.dynamic("n") + @R.function - def foo(x: R.Tensor(("m", "n"), "float32")): - m, n = T.int64(), T.int64() + def foo(x: R.Tensor((m, n), "float32")): gv0 = R.call_dps_packed("test.op.identity", (x,), R.Tensor((m, n), dtype="float32")) return gv0 diff --git a/tests/python/relax/test_backend_dispatch_sampling.py b/tests/python/relax/test_backend_dispatch_sampling.py index bf191e6b2d32..f122329d8717 100644 --- a/tests/python/relax/test_backend_dispatch_sampling.py +++ b/tests/python/relax/test_backend_dispatch_sampling.py @@ -45,13 +45,15 @@ def foo( def test_dispatch_multinomial_from_uniform_generic(): # fmt: off + batch = T.dynamic("batch") + vocab_size = T.dynamic("vocab_size") + out_batch = T.dynamic("out_batch") + @I.ir_module class Expected: @Ts.prim_func(private=True) def get_sample_index(A: T.handle, B: T.handle, C: T.handle, D: T.handle): - batch, vocab_size = T.int64(), T.int64() prob = T.match_buffer(A, (batch, vocab_size)) - out_batch = T.int64() usample = T.match_buffer(B, (out_batch, 1)) sample_indices = T.match_buffer(C, (out_batch, 1), "int64") output_index = T.match_buffer(D, (out_batch, 1), "int64") @@ -84,14 +86,16 @@ def foo(prob: R.Tensor((3, 5), dtype="float32"), uniform_sample: R.Tensor((6, 1) def test_dispatch_multinomial_from_uniform_gpu(): # fmt: off + n = T.dynamic("n") + vocab_size = T.dynamic("vocab_size") + batch_size = T.dynamic("batch_size") + @I.ir_module class Expected: @Ts.prim_func def parallel_sampling_from_prob(var_prob: T.handle, var_uniform_samples: T.handle, var_row_indices: T.handle, var_sampled_token_ids: T.handle): T.func_attr({"tirx.is_scheduled": True}) - n, vocab_size = T.int64(), T.int64() prob = T.match_buffer(var_prob, (n, vocab_size)) - batch_size = T.int64() uniform_samples = T.match_buffer(var_uniform_samples, (batch_size, 1)) row_indices = T.match_buffer(var_row_indices, (batch_size, 1), "int64") token_ids = T.match_buffer(var_sampled_token_ids, (batch_size, 1), "int64") diff --git a/tests/python/relax/test_backend_dispatch_sort_scan.py b/tests/python/relax/test_backend_dispatch_sort_scan.py index 0884ca42c7e7..b5b380db3f4d 100644 --- a/tests/python/relax/test_backend_dispatch_sort_scan.py +++ b/tests/python/relax/test_backend_dispatch_sort_scan.py @@ -14,7 +14,6 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -# ruff: noqa: F841 import numpy as np import pytest @@ -73,12 +72,14 @@ def test_dispatch_scanop_cuda(): lowered to the packed func `"gpu_2d_continuous_cumsum"`. """ + m = T.dynamic("m") + @I.ir_module class Before: I.module_global_infos({"vdevice": [R.vdevice("cuda", 0)]}) @R.function - def main(x: 'R.Tensor(("m", 3), "float32", "cuda")'): + def main(x: 'R.Tensor((m, 3), "float32", "cuda")'): with R.dataflow(): lv0 = R.cumsum(x, axis=1, exclusive=True) lv1 = R.cumprod(lv0, axis=1) @@ -89,7 +90,7 @@ def main(x: 'R.Tensor(("m", 3), "float32", "cuda")'): target = tvm.target.Target("cuda", host="llvm") vdevices = [R.vdevice("cuda", 0)] - m = tirx.Var("m", "int64") + m = T.dynamic("m", "int64") x = relax.Var("x", R.Tensor((m, 3), "float32", vdevices[0])) bb = relax.BlockBuilder() with target: @@ -119,13 +120,14 @@ def main(x: 'R.Tensor(("m", 3), "float32", "cuda")'): def test_dispatch_sort(): + m = T.dynamic("m") + @I.ir_module class Before: I.module_global_infos({"vdevice": [R.vdevice("llvm", 0)]}) @R.function - def foo(x: 'R.Tensor(("m", 3), "float32", "llvm")'): - m = T.int64() + def foo(x: 'R.Tensor((m, 3), "float32", "llvm")'): with R.dataflow(): lv = R.sort(x, axis=1, descending=False) gv = lv @@ -133,7 +135,7 @@ def foo(x: 'R.Tensor(("m", 3), "float32", "llvm")'): return gv vdevices = [R.vdevice("llvm", 0)] - m = tirx.Var("m", "int64") + m = T.dynamic("m", "int64") x = relax.Var("x", R.Tensor((m, 3), "float32", vdevices[0])) bb = relax.BlockBuilder() @@ -216,13 +218,14 @@ def foo2(y: R.Tensor((2, 3), "float32")): def test_dispatch_argsort(): + m = T.dynamic("m") + @I.ir_module class Before: I.module_global_infos({"vdevice": [R.vdevice("llvm", 0)]}) @R.function - def foo(x: 'R.Tensor(("m", 3), "float32", "llvm")'): - m = T.int64() + def foo(x: 'R.Tensor((m, 3), "float32", "llvm")'): with R.dataflow(): lv = R.argsort(x, axis=1, descending=False, dtype="int32") gv = lv @@ -230,7 +233,7 @@ def foo(x: 'R.Tensor(("m", 3), "float32", "llvm")'): return gv vdevices = [R.vdevice("llvm", 0)] - m = tirx.Var("m", "int64") + m = T.dynamic("m", "int64") x = relax.Var("x", R.Tensor((m, 3), "float32", vdevices[0])) bb = relax.BlockBuilder() @@ -309,13 +312,14 @@ def foo2(y: R.Tensor((2, 3), "float32")): def test_dispatch_topk(): + m = T.dynamic("m") + @I.ir_module class Before: I.module_global_infos({"vdevice": [R.vdevice("llvm", 0)]}) @R.function - def foo(x: 'R.Tensor(("m", 3), "float32", "llvm")'): - m = T.int64() + def foo(x: 'R.Tensor((m, 3), "float32", "llvm")'): with R.dataflow(): lv = R.topk(x, k=2, axis=1, largest=True) gv = lv @@ -323,7 +327,7 @@ def foo(x: 'R.Tensor(("m", 3), "float32", "llvm")'): return gv vdevices = [R.vdevice("llvm", 0)] - m = tirx.Var("m", "int64") + m = T.dynamic("m", "int64") x = relax.Var("x", R.Tensor((m, 3), "float32", vdevices[0])) bb = relax.BlockBuilder() @@ -418,10 +422,13 @@ def test_dispatch_sort_cuda_large_batch(size): if not tvm.testing.device_enabled(target): pytest.skip(f"{target} not enabled") + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Module: @R.function - def main(x: R.Tensor(("m", "n"), "float32")): + def main(x: R.Tensor((m, n), "float32")): with R.dataflow(): gv = R.sort(x, axis=-1, descending=False) R.output(gv) @@ -451,10 +458,13 @@ def test_dispatch_topk_cuda_large_batch(): if not tvm.testing.device_enabled(target): pytest.skip(f"{target} not enabled") + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Module: @R.function - def main(x: R.Tensor(("m", "n"), "float32")): + def main(x: R.Tensor((m, n), "float32")): with R.dataflow(): gv = R.topk(x, k=1, axis=-1, ret_type="values", largest=True) R.output(gv) @@ -515,10 +525,13 @@ def test_dispatch_cumsum_gpu(target, index_bits): if not tvm.testing.device_enabled(target): pytest.skip(f"{target} not enabled") + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Module: @R.function - def main(x: R.Tensor(("m", "n"), "int32")): + def main(x: R.Tensor((m, n), "int32")): with R.dataflow(): gv = R.cumsum(x, axis=-1, exclusive=False) R.output(gv) @@ -549,10 +562,13 @@ def test_dispatch_cumsum_index_width(target_kind, index_bits): """Respect the caller's index budget without restricting Metal's default.""" from tvm.relax.backend.gpu_generic import gpu_2d_continuous_cumsum + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Module: @R.function - def main(x: R.Tensor(("m", "n"), "float32")): + def main(x: R.Tensor((m, n), "float32")): gv = R.cumsum(x, axis=-1) return gv @@ -588,10 +604,13 @@ def test_dispatch_cumprod_cuda_large_batch(): if not tvm.testing.device_enabled(target): pytest.skip(f"{target} not enabled") + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Module: @R.function - def main(x: R.Tensor(("m", "n"), "float32")): + def main(x: R.Tensor((m, n), "float32")): with R.dataflow(): gv = R.cumprod(x, axis=1) R.output(gv) @@ -691,12 +710,14 @@ def collect_floor_divisors(node): def test_dispatch_cumsum_webgpu_symbolic_non_contiguous_axis(): """The serial WebGPU fallback accepts a symbolic scan extent.""" + n = T.dynamic("n") + @I.ir_module class Symbolic: I.module_global_infos({"vdevice": [R.vdevice("webgpu", 0)]}) @R.function - def main(x: 'R.Tensor((1, "n", 9), "float32", "webgpu")'): + def main(x: 'R.Tensor((1, n, 9), "float32", "webgpu")'): return R.cumsum(x, axis=1) target = tvm.target.Target("webgpu", host="llvm") diff --git a/tests/python/relax/test_backend_transform_shape_lower.py b/tests/python/relax/test_backend_transform_shape_lower.py index c8e830880779..7748187dd06f 100644 --- a/tests/python/relax/test_backend_transform_shape_lower.py +++ b/tests/python/relax/test_backend_transform_shape_lower.py @@ -118,10 +118,13 @@ def main(f: R.Callable([R.Any], R.Any), y: R.Shape([1, 2])): def test_simple_symbolic_shape(): MS = MatchShapeCode + n = T.dynamic("n") + m = T.dynamic("m") + @tvm.script.ir_module class Before: @R.function - def main(x: R.Tensor(["n", 2, "m"], "float32")): + def main(x: R.Tensor([n, 2, m], "float32")): R.func_attr({"relax.force_pure": True}) return x @@ -130,10 +133,13 @@ def main(x: R.Tensor(["n", 2, "m"], "float32")): "m": 1, } + n = T.dynamic("n") + m = T.dynamic("m") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(["n", 2, "m"], "float32")): + def main(x: R.Tensor([n, 2, m], "float32")): R.func_attr({"relax.force_pure": True}) shape_heap = R.call_builtin_with_ctx( "vm.builtin.alloc_shape_heap", @@ -174,15 +180,17 @@ def test_symbolic_compute(): MS = MatchShapeCode MK = MakeShapeCode + m = T.dynamic("m") + k = T.dynamic("k") + n = T.dynamic("n") + @tvm.script.ir_module class Before: @R.function - def main(x: R.Tensor(["n", "m"], "float32"), y: R.Tensor(ndim=3, dtype=None)) -> R.Shape( + def main(x: R.Tensor([n, m], "float32"), y: R.Tensor(ndim=3, dtype=None)) -> R.Shape( ndim=3 ): R.func_attr({"relax.force_pure": True}) - m = T.int64() - k = T.int64() z = R.match_cast(y, R.Tensor([k, m, k + 1], dtype=None)) return R.shape([k + 1, m, 2]) @@ -190,6 +198,10 @@ def main(x: R.Tensor(["n", "m"], "float32"), y: R.Tensor(ndim=3, dtype=None)) -> # 0: n, 1: m, 2:k, 3: k+1 sindex = {"n": 0, "m": 1, "k": 2, "k+1": 3} + m = T.dynamic("m") + k = T.dynamic("k") + n = T.dynamic("n") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) @@ -199,12 +211,10 @@ def shape_func(H: T.Buffer(T.int64(4), "int64")): H[T.int64(sindex["k+1"])] = H[T.int64(sindex["k"])] + T.int64(1) @R.function - def main(x: R.Tensor(["n", "m"], "float32"), y: R.Tensor(ndim=3, dtype=None)) -> R.Shape( + def main(x: R.Tensor([n, m], "float32"), y: R.Tensor(ndim=3, dtype=None)) -> R.Shape( ndim=3 ): R.func_attr({"relax.force_pure": True}) - m = T.int64() - k = T.int64() cls = Expected shape_heap = R.call_builtin_with_ctx( "vm.builtin.alloc_shape_heap", @@ -288,13 +298,15 @@ def main(x: R.Tensor(["n", "m"], "float32"), y: R.Tensor(ndim=3, dtype=None)) -> def test_tuple_handling(): MS = MatchShapeCode + n = T.dynamic("n") + m = T.dynamic("m") + k = T.dynamic("k") + @tvm.script.ir_module class Before: @R.function def main( - x: R.Tuple( - R.Tensor(["n", "m"], "float32"), R.Tuple(R.Shape, R.Tensor(["n", "k"], "int32")) - ), + x: R.Tuple(R.Tensor([n, m], "float32"), R.Tuple(R.Shape, R.Tensor([n, k], "int32"))), ): R.func_attr({"relax.force_pure": True}) return x @@ -302,13 +314,15 @@ def main( # slot assignment: sindex = {"n": 0, "m": 1, "k": 2} + n = T.dynamic("n") + m = T.dynamic("m") + k = T.dynamic("k") + @tvm.script.ir_module class Expected: @R.function def main( - x: R.Tuple( - R.Tensor(["n", "m"], "float32"), R.Tuple(R.Shape, R.Tensor(["n", "k"], "int32")) - ), + x: R.Tuple(R.Tensor([n, m], "float32"), R.Tuple(R.Shape, R.Tensor([n, k], "int32"))), ): R.func_attr({"relax.force_pure": True}) shape_heap = R.call_builtin_with_ctx( @@ -377,12 +391,13 @@ def test_return_match_check(): """Test when return body is not same as ret_ty, runtime match check needed.""" MS = MatchShapeCode + n = T.dynamic("n") + m = T.dynamic("m") + @tvm.script.ir_module class Before: @R.function - def main(x: R.Tensor(["n", "m"], "float32"), y: R.Any) -> R.Tuple( - R.Tensor(["n", "m"], "float32") - ): + def main(x: R.Tensor([n, m], "float32"), y: R.Any) -> R.Tuple(R.Tensor([n, m], "float32")): R.func_attr({"relax.force_pure": True}) return y @@ -392,12 +407,13 @@ def main(x: R.Tensor(["n", "m"], "float32"), y: R.Any) -> R.Tuple( "m": 1, } + n = T.dynamic("n") + m = T.dynamic("m") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(["n", "m"], "float32"), y: R.Any) -> R.Tuple( - R.Tensor(["n", "m"], "float32") - ): + def main(x: R.Tensor([n, m], "float32"), y: R.Any) -> R.Tuple(R.Tensor([n, m], "float32")): R.func_attr({"relax.force_pure": True}) shape_heap = R.call_builtin_with_ctx( "vm.builtin.alloc_shape_heap", @@ -462,10 +478,12 @@ def test_return_match_check_with_new_expr(): """ MS = MatchShapeCode + n = T.dynamic("n") + @tvm.script.ir_module class Before: @R.function - def main(x: R.Tensor(["n", "n"], "float32")) -> R.Tensor(["n * n"], "float32"): + def main(x: R.Tensor([n, n], "float32")) -> R.Tensor([n * n], "float32"): R.func_attr({"relax.force_pure": True}) out = R.call_packed("flatten_matrix", x, ty_args=R.Any) return out @@ -476,10 +494,12 @@ def main(x: R.Tensor(["n", "n"], "float32")) -> R.Tensor(["n * n"], "float32"): "n * n": 1, } + n = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(["n", "n"], "float32")) -> R.Tensor(["n * n"], "float32"): + def main(x: R.Tensor([n, n], "float32")) -> R.Tensor([n * n], "float32"): R.func_attr({"relax.force_pure": True}) shape_heap = R.call_builtin_with_ctx( "vm.builtin.alloc_shape_heap", @@ -541,20 +561,21 @@ def test_symbolic_shape_multiple_function(): MS = MatchShapeCode MK = MakeShapeCode + m_fn1 = T.dynamic("m") + n_fn1 = T.dynamic("n") + n_fn2 = T.dynamic("n") + m_fn2 = T.dynamic("m") + @I.ir_module class Before: @R.function - def fn1(A: R.Tensor(("m", "n"), dtype="float32")): + def fn1(A: R.Tensor((m_fn1, n_fn1), dtype="float32")): R.func_attr({"relax.force_pure": True}) - m = T.int64() - n = T.int64() return A @R.function - def fn2(A: R.Tensor(("n", "m"), dtype="float32")): + def fn2(A: R.Tensor((n_fn2, m_fn2), dtype="float32")): R.func_attr({"relax.force_pure": True}) - n = T.int64() - m = T.int64() return A # slot assignment: @@ -567,13 +588,18 @@ def fn2(A: R.Tensor(("n", "m"), dtype="float32")): "m": 1, } + m_fn1 = T.dynamic("m") + n_fn1 = T.dynamic("n") + n_fn2 = T.dynamic("n") + m_fn2 = T.dynamic("m") + @I.ir_module class Expected: @R.function - def fn1(A: R.Tensor(("m", "n"), dtype="float32")) -> R.Tensor(("m", "n"), dtype="float32"): + def fn1(A: R.Tensor((m_fn1, n_fn1), dtype="float32")) -> R.Tensor( + (m_fn1, n_fn1), dtype="float32" + ): R.func_attr({"relax.force_pure": True}) - m = T.int64() - n = T.int64() shape_heap: R.Tensor(dtype="int64", ndim=1) = R.call_builtin_with_ctx( "vm.builtin.alloc_shape_heap", (R.prim_value(2),), @@ -602,10 +628,10 @@ def fn1(A: R.Tensor(("m", "n"), dtype="float32")) -> R.Tensor(("m", "n"), dtype= return A @R.function - def fn2(A: R.Tensor(("n", "m"), dtype="float32")) -> R.Tensor(("n", "m"), dtype="float32"): + def fn2(A: R.Tensor((n_fn2, m_fn2), dtype="float32")) -> R.Tensor( + (n_fn2, m_fn2), dtype="float32" + ): R.func_attr({"relax.force_pure": True}) - 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),), @@ -731,27 +757,29 @@ def main( def test_check_weights_with_dynamic_shape(): MS = MatchShapeCode + n = T.dynamic("n") + @I.ir_module class Before: @R.function def main( x: R.Tensor((16, 16), "float32"), - params: R.Tuple(R.Tensor((16, 16), dtype="float32"), R.Tensor(("n",), "float32")), + params: R.Tuple(R.Tensor((16, 16), dtype="float32"), R.Tensor((n,), "float32")), ): R.func_attr({"relax.force_pure": True, "num_input": 1}) - n = T.int64() param_0 = params[0] param_1 = params[1] return (x, param_0, param_1) + n = T.dynamic("n") + @I.ir_module class Expected: @R.function def main( x: R.Tensor((16, 16), "float32"), - params: R.Tuple(R.Tensor((16, 16), dtype="float32"), R.Tensor(("n",), "float32")), + params: R.Tuple(R.Tensor((16, 16), dtype="float32"), R.Tensor((n,), "float32")), ): - n = T.int64() R.func_attr({"num_input": 1, "relax.force_pure": True}) shape_heap: R.Tensor(dtype="int64", ndim=1) = R.call_builtin_with_ctx( "vm.builtin.alloc_shape_heap", diff --git a/tests/python/relax/test_base_py_module_symbolic_shape.py b/tests/python/relax/test_base_py_module_symbolic_shape.py index bd3c26d6b006..cbd677911234 100644 --- a/tests/python/relax/test_base_py_module_symbolic_shape.py +++ b/tests/python/relax/test_base_py_module_symbolic_shape.py @@ -63,23 +63,26 @@ def test_infer_concrete_shape_error_when_uninferrable(): bpm._infer_concrete_shape_from_args([k, 8], in_args=[]) +n_add_tir = T.dynamic("n") +n_main_relax = T.dynamic("n") + + @R.py_module class AddModuleSymbolic(BasePyModule): @Ts.prim_func def add_tir(var_x: T.handle, var_y: T.handle, var_out: T.handle): T.func_attr({"global_symbol": "add_tir"}) - n = T.int64() - x = T.match_buffer(var_x, (n,), dtype="float32") - y = T.match_buffer(var_y, (n,), dtype="float32") - out = T.match_buffer(var_out, (n,), dtype="float32") + x = T.match_buffer(var_x, (n_add_tir,), dtype="float32") + y = T.match_buffer(var_y, (n_add_tir,), dtype="float32") + out = T.match_buffer(var_out, (n_add_tir,), dtype="float32") - for i in T.serial(n): + for i in T.serial(n_add_tir): out[i] = x[i] + y[i] @R.function - def main_relax(x: R.Tensor(("n",), "float32"), y: R.Tensor(("n",), "float32")) -> R.Tensor( - ("n",), "float32" - ): + def main_relax( + x: R.Tensor((n_main_relax,), "float32"), y: R.Tensor((n_main_relax,), "float32") + ) -> R.Tensor((n_main_relax,), "float32"): return R.add(x, y) @@ -193,28 +196,34 @@ def test_infer_concrete_shape_wrong_ndim(): bpm._infer_concrete_shape_from_args(sym_shape, [x]) +m_matmul_tir = T.dynamic("m") +n_matmul_tir = T.dynamic("n") +k_matmul_tir = T.dynamic("k") +m_matmul_relax = T.dynamic("m") +k_matmul_relax = T.dynamic("k") +n_matmul_relax = T.dynamic("n") + + @R.py_module class MatrixModuleSymbolic(BasePyModule): @Ts.prim_func def matmul_tir(var_a: T.handle, var_b: T.handle, var_c: T.handle): T.func_attr({"global_symbol": "matmul_tir"}) - m = T.int64() - n = T.int64() - k = T.int64() - a = T.match_buffer(var_a, (m, k), dtype="float32") - b = T.match_buffer(var_b, (k, n), dtype="float32") - c = T.match_buffer(var_c, (m, n), dtype="float32") - - for i in T.serial(m): - for j in T.serial(n): + a = T.match_buffer(var_a, (m_matmul_tir, k_matmul_tir), dtype="float32") + b = T.match_buffer(var_b, (k_matmul_tir, n_matmul_tir), dtype="float32") + c = T.match_buffer(var_c, (m_matmul_tir, n_matmul_tir), dtype="float32") + + for i in T.serial(m_matmul_tir): + for j in T.serial(n_matmul_tir): c[i, j] = 0.0 - for l in T.serial(k): + for l in T.serial(k_matmul_tir): c[i, j] = c[i, j] + a[i, l] * b[l, j] @R.function def matmul_relax( - a: R.Tensor(("m", "k"), "float32"), b: R.Tensor(("k", "n"), "float32") - ) -> R.Tensor(("m", "n"), "float32"): + a: R.Tensor((m_matmul_relax, k_matmul_relax), "float32"), + b: R.Tensor((k_matmul_relax, n_matmul_relax), "float32"), + ) -> R.Tensor((m_matmul_relax, n_matmul_relax), "float32"): return R.matmul(a, b) diff --git a/tests/python/relax/test_bind_symbolic_vars.py b/tests/python/relax/test_bind_symbolic_vars.py index f42271b65ab8..0c3e2efefc20 100644 --- a/tests/python/relax/test_bind_symbolic_vars.py +++ b/tests/python/relax/test_bind_symbolic_vars.py @@ -32,8 +32,12 @@ def test_bind_static_value(replace_by_tir_var): The replaced variables may be given either as strings, or as TIR variables """ + M = T.dynamic("M") + K = T.dynamic("K") + N = T.dynamic("N") + @R.function(private=True) - def before(A: R.Tensor(("M", "K")), B: R.Tensor(("K", "N"))) -> R.Tensor(("M", "N")): + def before(A: R.Tensor((M, K)), B: R.Tensor((K, N))) -> R.Tensor((M, N)): return R.matmul(A, B) @R.function(private=True) @@ -58,8 +62,8 @@ def test_error_with_duplicate_var_names(): variables share the same name, the replacement map may not refer to that variable by string. """ - N1 = tvm.tirx.Var("N", "int64") - N2 = tvm.tirx.Var("N", "int64") + N1 = T.dynamic("N", "int64") + N2 = T.dynamic("N", "int64") @R.function(private=True) def func(A: R.Tensor((N1, N1)), B: R.Tensor((N1, N2))) -> R.Tensor((N1, N2)): @@ -77,9 +81,9 @@ def test_string_var_when_other_var_has_duplicate_var_names(): replacing variables by name only applies to those duplicate names. Other variables may still be replaced by name. """ - N1 = tvm.tirx.Var("N", "int64") - N2 = tvm.tirx.Var("N", "int64") - BatchSize = tvm.tirx.Var("BatchSize", "int64") + N1 = T.dynamic("N", "int64") + N2 = T.dynamic("N", "int64") + BatchSize = T.dynamic("BatchSize", "int64") @R.function(private=True) def before(A: R.Tensor((BatchSize, N1, N1)), B: R.Tensor((N1, N2))) -> R.Tensor( @@ -100,8 +104,11 @@ def expected(A: R.Tensor((16, N1, N1)), B: R.Tensor((N1, N2))) -> R.Tensor((16, def test_error_with_nonexisting_var_name(): """A string name of a symbolic var must be used by the function""" + M = T.dynamic("M") + N = T.dynamic("N") + @R.function(private=True) - def func(A: R.Tensor(("M", "N"))): + def func(A: R.Tensor((M, N))): return A with pytest.raises(RuntimeError): @@ -111,8 +118,11 @@ def func(A: R.Tensor(("M", "N"))): def test_error_with_nonexisting_tir_var(): """A TIR symbolic var must be a symbolic var of the function""" + M = T.dynamic("M") + N = T.dynamic("N") + @R.function(private=True) - def func(A: R.Tensor(["M", "N"])): + def func(A: R.Tensor([M, N])): return A with pytest.raises(RuntimeError): @@ -122,8 +132,11 @@ def func(A: R.Tensor(["M", "N"])): def test_error_with_multiple_definitions(): """The string/TIR var syntaxes may not define the same variable""" + M = T.dynamic("M") + N = T.dynamic("N") + @R.function(private=True) - def func(A: R.Tensor(["M", "N"])): + def func(A: R.Tensor([M, N])): return A tir_var = func.params[0].ty.shape[0] @@ -136,11 +149,14 @@ def func(A: R.Tensor(["M", "N"])): def test_error_if_output_has_undefined(): """The replacements may not introduce undefined symbolic vars""" + M = T.dynamic("M") + N = T.dynamic("N") + @R.function(private=True) - def func(A: R.Tensor(["M", "N"])): + def func(A: R.Tensor([M, N])): return A - outside_var = tvm.tirx.Var("outside_var", "int64") + outside_var = T.dynamic("outside_var", "int64") with pytest.raises(RuntimeError): func.bind_symbolic_vars({"M": outside_var * 2}) @@ -149,15 +165,20 @@ def func(A: R.Tensor(["M", "N"])): def test_replacements_may_produce_new_symbolic_vars(): """The output may introduce symbolic vars, but they must be bound""" + M = T.dynamic("M") + N = T.dynamic("N") + @R.function(private=True) - def before(A: R.Tensor(["M", "N"])): + def before(A: R.Tensor([M, N])): return A + outside_var = T.dynamic("outside_var") + @R.function(private=True) - def expected(A: R.Tensor(["outside_var * 2", "outside_var"])): + def expected(A: R.Tensor([outside_var * 2, outside_var])): return A - outside_var = tvm.tirx.Var("outside_var", "int64") + outside_var = T.dynamic("outside_var", "int64") after = before.bind_symbolic_vars({"M": outside_var * 2, "N": outside_var}) tvm.ir.assert_structural_equal(expected, after) @@ -166,16 +187,18 @@ def expected(A: R.Tensor(["outside_var * 2", "outside_var"])): def test_bind_symbolic_vars_in_tensor_shape(): """The bound variable should be replaced when appearing in type""" + M = T.dynamic("M") + N = T.dynamic("N") + @R.function(private=True) - def before(A: R.Tensor(["M", "N"])): - M = T.int64() - N = T.int64() + def before(A: R.Tensor([M, N])): B = R.call_dps_packed("dummy_func", [A], out_ty=R.Tensor([2 * M * N])) return B + M = T.dynamic("M") + @R.function(private=True) - def expected(A: R.Tensor(["M", 16])): - M = T.int64() + def expected(A: R.Tensor([M, 16])): B = R.call_dps_packed("dummy_func", [A], out_ty=R.Tensor([M * 32])) return B @@ -186,16 +209,18 @@ def expected(A: R.Tensor(["M", 16])): def test_bind_symbolic_vars_in_shape_expr(): """The bound variable should be replaced when appearing in R.Shape""" + M = T.dynamic("M") + N = T.dynamic("N") + @R.function(private=True) - def before(A: R.Tensor(["M * N"]), x: R.Shape(["M", "N"])): - M = T.int64() - N = T.int64() + def before(A: R.Tensor([M * N]), x: R.Shape([M, N])): B = R.call_dps_packed("dummy_func", [A], out_ty=R.Tensor([2 * M * N])) return B + M = T.dynamic("M") + @R.function(private=True) - def expected(A: R.Tensor(["M * 16"]), x: R.Shape(["M", 16])): - M = T.int64() + def expected(A: R.Tensor([M * 16]), x: R.Shape([M, 16])): B = R.call_dps_packed("dummy_func", [A], out_ty=R.Tensor([M * 32])) return B @@ -206,14 +231,18 @@ def expected(A: R.Tensor(["M * 16"]), x: R.Shape(["M", 16])): def test_bind_strided_slice(): """relax.op.strided_slice stores Expr attributes""" + N = T.dynamic("N") + M = T.dynamic("M") + @R.function(private=True) - def before(A: R.Tensor(["M", "N"])): - N = T.int64() + def before(A: R.Tensor([M, N])): B = R.strided_slice(A, [1], [0], [N // 4]) return B + M = T.dynamic("M") + @R.function(private=True) - def expected(A: R.Tensor(["M", 32])): + def expected(A: R.Tensor([M, 32])): # Binding substitutes runtime primitive arguments without applying # shape-only analyzer simplification to them. B = R.strided_slice(A, [1], [0], [T.FloorDiv(T.int64(32), T.int64(4))]) @@ -226,17 +255,19 @@ def expected(A: R.Tensor(["M", 32])): def test_bind_inside_match_cast(): """Symbolic variables may occur within R.match_cast""" + M = T.dynamic("M") + N = T.dynamic("N") + @R.function(private=True) - def before(A: R.Tensor(["M", "N"]), B: R.Tensor(ndim=2)): - M = T.int64() - N = T.int64() + def before(A: R.Tensor([M, N]), B: R.Tensor(ndim=2)): C = R.match_cast(B, R.Tensor([M, N])) D = R.add(A, C) return D + M = T.dynamic("M") + @R.function(private=True) - def expected(A: R.Tensor(["M", 32]), B: R.Tensor(ndim=2)): - M = T.int64() + def expected(A: R.Tensor([M, 32]), B: R.Tensor(ndim=2)): C = R.match_cast(B, R.Tensor([M, 32])) D = R.add(A, C) return D diff --git a/tests/python/relax/test_blockbuilder_core.py b/tests/python/relax/test_blockbuilder_core.py index 988317ee3dc9..b7a7d27911db 100644 --- a/tests/python/relax/test_blockbuilder_core.py +++ b/tests/python/relax/test_blockbuilder_core.py @@ -41,8 +41,8 @@ def nop(): def test_block_builder(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = rx.Var("x", rx.TensorType([m, n], "float16")) y = rx.Var("y", rx.TensorType([n], "float16")) bb = rx.BlockBuilder() @@ -66,8 +66,8 @@ def test_block_builder(): def test_emit_with_name(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = rx.Var("x", rx.TensorType([m, n], "float16")) y = rx.Var("y", rx.TensorType([n], "float16")) bb = rx.BlockBuilder() @@ -82,8 +82,8 @@ def test_emit_with_name(): def test_function_single_block(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = rx.Var("x", rx.TensorType([m, n], "float16")) y = rx.Var("y", rx.TensorType([n], "float16")) bb = rx.BlockBuilder() @@ -108,8 +108,8 @@ def test_function_single_block(): def test_function_multi_blocks(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = rx.Var("x", rx.TensorType([m, n], "float16")) y = rx.Var("y", rx.TensorType([n], "float16")) bb = rx.BlockBuilder() @@ -143,8 +143,8 @@ def test_function_multi_blocks(): def test_multi_functions(): bb = rx.BlockBuilder() - m_1 = tirx.Var("m", "int64") - n_1 = tirx.Var("n", "int64") + m_1 = T.dynamic("m", "int64") + n_1 = T.dynamic("n", "int64") x_1 = rx.Var("x", rx.TensorType([m_1, n_1], "float16")) y_1 = rx.Var("y", rx.TensorType([n_1], "float16")) @@ -180,8 +180,8 @@ def test_multi_functions(): def test_binary_shape_type_deduction(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") k = tirx.Var("k", "int64") x = rx.Var("x", rx.TensorType([m, 1], "float16")) y = rx.Var("y", rx.TensorType([n], "float16")) @@ -216,8 +216,8 @@ def test_binary_shape_type_deduction(): def test_emit_match_cast(): - m = tirx.Var("m", ty="int64") - n = tirx.Var("n", ty="int64") + m = T.dynamic("m", dtype="int64") + n = T.dynamic("n", dtype="int64") x = rx.Var("tensor_value", rx.TensorType(dtype="float32", ndim=-1)) y = rx.Var("shape_value", rx.ShapeType([16, 8])) bb = rx.BlockBuilder() @@ -256,7 +256,7 @@ def test_emit_match_cast_binding_in_dataflow_block(): bb = rx.BlockBuilder() x = rx.Var("x", rx.TensorType(dtype="float32", ndim=-1)) - m = tirx.Var("m", ty="int64") + m = T.dynamic("m", dtype="int64") gv = rx.Var("gv", rx.TensorType(dtype="float32", ndim=-1)) match_cast = rx.MatchCast(gv, x, rx.TensorType((m,), "float32")) @@ -278,8 +278,8 @@ def test_emit_match_cast_binding_in_dataflow_block(): def test_normalize(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = rx.Var("x", rx.TensorType([m, n], "float16")) y = rx.Var("y", rx.TensorType([n], "float16")) @@ -314,8 +314,8 @@ def test_normalize(): def test_tuple_indexing(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") shape_x = rx.TensorType([m, n], "float16") shape_y = rx.TensorType([n], "float16") @@ -561,8 +561,8 @@ def test_emit_te_prim_value(): def test_nested_function_fail(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = rx.Var("x", rx.TensorType([m, n], "float16")) y = rx.Var("y", rx.TensorType([n], "float16")) bb = rx.BlockBuilder() @@ -576,8 +576,8 @@ def test_nested_function_fail(): def test_emit_func_output_twice_fail(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = rx.Var("x", rx.TensorType([m, n], "float16")) y = rx.Var("y", rx.TensorType([n], "float16")) bb = rx.BlockBuilder() @@ -590,8 +590,8 @@ def test_emit_func_output_twice_fail(): def test_func_params_twice_fail(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = rx.Var("x", rx.TensorType([m, n], "float16")) y = rx.Var("y", rx.TensorType([n], "float16")) bb = rx.BlockBuilder() @@ -603,8 +603,8 @@ def test_func_params_twice_fail(): def test_no_func_params_fail(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = rx.Var("x", rx.TensorType([m, n], "float16")) y = rx.Var("y", rx.TensorType([n], "float16")) bb = rx.BlockBuilder() @@ -661,24 +661,28 @@ def make_function(emit_nested_tuple: bool): def make_expected(emit_nested_tuple: bool): if emit_nested_tuple: + n = T.dynamic("n") + m = T.dynamic("m") @R.function def func( n_1: T.int64, m_1: T.int64, - x: R.Tensor(("n", "m"), dtype="float32"), - y: R.Tensor(("m", "n"), dtype="float32"), + x: R.Tensor((n, m), dtype="float32"), + y: R.Tensor((m, n), dtype="float32"), ): return ((n_1, m_1), x, y) else: + n = T.dynamic("n") + m = T.dynamic("m") @R.function def func( n_1: T.int64, m_1: T.int64, - x: R.Tensor(("n", "m"), dtype="float32"), - y: R.Tensor(("m", "n"), dtype="float32"), + x: R.Tensor((n, m), dtype="float32"), + y: R.Tensor((m, n), dtype="float32"), ): gv = n_1, m_1 return (gv, x, y) diff --git a/tests/python/relax/test_blockbuilder_emit_te.py b/tests/python/relax/test_blockbuilder_emit_te.py index 8c5f8db1344a..1cbeea03c6fa 100644 --- a/tests/python/relax/test_blockbuilder_emit_te.py +++ b/tests/python/relax/test_blockbuilder_emit_te.py @@ -20,7 +20,7 @@ import tvm from tvm import relax as rx -from tvm import te, tirx +from tvm import te from tvm.ir.base import assert_structural_equal from tvm.script import s_tir as Ts from tvm.script.parser import ir as I @@ -30,7 +30,7 @@ def test_emit_te_with_symbolic_arg(): bb = rx.BlockBuilder() - m = tirx.Var("m", "int64") + m = T.dynamic("m", "int64") x = rx.Var("x", R.Tensor([10], "float32")) y = rx.Var("y", R.Shape([m])) @@ -43,6 +43,8 @@ def te_func(A, offset): after = bb.get() + m = T.dynamic("m") + @I.ir_module class Expected: @Ts.prim_func(private=True) @@ -59,10 +61,9 @@ def te_func( B[v_i] = A[v_i + m] @R.function - def main(x: R.Tensor((10,), dtype="float32"), y: R.Shape(["m"])) -> R.Tensor( + def main(x: R.Tensor((10,), dtype="float32"), y: R.Shape([m])) -> R.Tensor( (10,), dtype="float32" ): - m = T.int64() cls = Expected gv = R.call_tir( cls.te_func, diff --git a/tests/python/relax/test_codegen_cutlass.py b/tests/python/relax/test_codegen_cutlass.py index 5ce792ef8b04..ea2ac3effb03 100644 --- a/tests/python/relax/test_codegen_cutlass.py +++ b/tests/python/relax/test_codegen_cutlass.py @@ -1801,6 +1801,8 @@ def main( def test_fp16A_int8B_gemm_batched(): + b = T.dynamic("b") + @I.ir_module class Module: @Ts.prim_func @@ -1868,12 +1870,11 @@ def encode( @R.function def main( - x: R.Tensor(("b", 64, 64), dtype="float16"), + x: R.Tensor((b, 64, 64), dtype="float16"), y: R.Tensor((64, 64), dtype="float16"), - ) -> R.Tensor(("b", 64, 64), dtype="float16"): + ) -> R.Tensor((b, 64, 64), dtype="float16"): R.func_attr({"num_input": 1}) cls = Module - b = T.int64() with R.dataflow(): lv = R.call_tir( cls.encode, @@ -1935,6 +1936,8 @@ def run_and_check(): def test_fp16A_int8B_gemm_batched_finegrained(): + b = T.dynamic("b") + @I.ir_module class Module: @Ts.prim_func @@ -2021,12 +2024,11 @@ def encode( @R.function def main( - x: R.Tensor(("b", 128, 128), dtype="float16"), + x: R.Tensor((b, 128, 128), dtype="float16"), y: R.Tensor((128, 128), dtype="float16"), - ) -> R.Tensor(("b", 128, 128), dtype="float16"): + ) -> R.Tensor((b, 128, 128), dtype="float16"): R.func_attr({"num_input": 1}) cls = Module - b = T.int64() with R.dataflow(): lv = R.call_tir( cls.encode, @@ -2213,6 +2215,9 @@ def _test_batched_var_len_attention( def test_batched_var_len_attention(): + num_tokens = T.dynamic("num_tokens") + num_seq = T.dynamic("num_seq") + @I.ir_module class Module: I.module_global_infos( @@ -2225,15 +2230,13 @@ class Module: @R.function def main( - queries: R.Tensor(("num_tokens", 4096), dtype="float16"), - keys: R.Tensor(("num_tokens", 4096), dtype="float16"), - values: R.Tensor(("num_tokens", 4096), dtype="float16"), - seq_lens: R.Tensor(("num_seq",), dtype="int32"), - ) -> R.Tensor(("num_tokens", 4096), dtype="float16"): + queries: R.Tensor((num_tokens, 4096), dtype="float16"), + keys: R.Tensor((num_tokens, 4096), dtype="float16"), + values: R.Tensor((num_tokens, 4096), dtype="float16"), + seq_lens: R.Tensor((num_seq,), dtype="int32"), + ) -> R.Tensor((num_tokens, 4096), dtype="float16"): R.func_attr({"num_input": 4}) cls = Module - num_tokens = T.int64() - num_seq = T.int64() with R.dataflow(): # TODO(masahi): Workaround for the broken Relax cumsum op on GPU. @@ -2266,6 +2269,9 @@ def main( def test_batched_var_len_multi_query_attention(): + num_tokens = T.dynamic("num_tokens") + num_seq = T.dynamic("num_seq") + @I.ir_module class Module: I.module_global_infos( @@ -2278,15 +2284,13 @@ class Module: @R.function def main( - queries: R.Tensor(("num_tokens", 4096), dtype="float16"), - keys: R.Tensor(("num_tokens", 512), dtype="float16"), - values: R.Tensor(("num_tokens", 512), dtype="float16"), - seq_lens: R.Tensor(("num_seq",), dtype="int32"), - ) -> R.Tensor(("num_tokens", 4096), dtype="float16"): + queries: R.Tensor((num_tokens, 4096), dtype="float16"), + keys: R.Tensor((num_tokens, 512), dtype="float16"), + values: R.Tensor((num_tokens, 512), dtype="float16"), + seq_lens: R.Tensor((num_seq,), dtype="int32"), + ) -> R.Tensor((num_tokens, 4096), dtype="float16"): R.func_attr({"num_input": 4}) cls = Module - num_tokens = T.int64() - num_seq = T.int64() with R.dataflow(): # TODO(masahi): Workaround for the broken Relax cumsum op on GPU. @@ -2361,6 +2365,9 @@ def test_sliding_window(): def test_batched_var_len_sliding_window(): + num_tokens = T.dynamic("num_tokens") + num_seq = T.dynamic("num_seq") + @I.ir_module class Module: I.module_global_infos( @@ -2373,15 +2380,13 @@ class Module: @R.function def main( - queries: R.Tensor(("num_tokens", 4096), dtype="float16"), - keys: R.Tensor(("num_tokens", 4096), dtype="float16"), - values: R.Tensor(("num_tokens", 4096), dtype="float16"), - seq_lens: R.Tensor(("num_seq",), dtype="int32"), - ) -> R.Tensor(("num_tokens", 4096), dtype="float16"): + queries: R.Tensor((num_tokens, 4096), dtype="float16"), + keys: R.Tensor((num_tokens, 4096), dtype="float16"), + values: R.Tensor((num_tokens, 4096), dtype="float16"), + seq_lens: R.Tensor((num_seq,), dtype="int32"), + ) -> R.Tensor((num_tokens, 4096), dtype="float16"): R.func_attr({"num_input": 4}) cls = Module - num_tokens = T.int64() - num_seq = T.int64() with R.dataflow(): # TODO(masahi): Workaround for the broken Relax cumsum op on GPU. diff --git a/tests/python/relax/test_contrib_vllm.py b/tests/python/relax/test_contrib_vllm.py index c9070e1caf35..c10b08f719a4 100644 --- a/tests/python/relax/test_contrib_vllm.py +++ b/tests/python/relax/test_contrib_vllm.py @@ -64,6 +64,10 @@ def build_and_run(mod, inputs_np, target, legalize=True): def test_attention(): + num_seqs = T.dynamic("num_seqs") + num_blocks = T.dynamic("num_blocks") + max_num_blocks_per_seq = T.dynamic("max_num_blocks_per_seq") + @I.ir_module class ModulePagedAttentionV1: I.module_global_infos( @@ -76,12 +80,12 @@ class ModulePagedAttentionV1: @R.function def main( - query: R.Tensor(("num_seqs", 1, 64), dtype="float16"), - key_cache: R.Tensor(("num_blocks", 1, 8, 16, 8), dtype="float16"), - value_cache: R.Tensor(("num_blocks", 1, 64, 16), dtype="float16"), - block_tables: R.Tensor(("num_seqs", "max_num_blocks_per_seq"), dtype="int32"), - context_lens: R.Tensor(("num_seqs",), dtype="int32"), - ) -> R.Tensor(("num_seqs", 1, 64), dtype="float16"): + query: R.Tensor((num_seqs, 1, 64), dtype="float16"), + key_cache: R.Tensor((num_blocks, 1, 8, 16, 8), dtype="float16"), + value_cache: R.Tensor((num_blocks, 1, 64, 16), dtype="float16"), + block_tables: R.Tensor((num_seqs, max_num_blocks_per_seq), dtype="int32"), + context_lens: R.Tensor((num_seqs,), dtype="int32"), + ) -> R.Tensor((num_seqs, 1, 64), dtype="float16"): with R.dataflow(): max_len = R.to_vdevice(R.max(context_lens), "llvm:0") out = R.call_dps_packed( @@ -100,6 +104,10 @@ def main( R.output(out) return out + num_seqs = T.dynamic("num_seqs") + num_blocks = T.dynamic("num_blocks") + max_num_blocks_per_seq = T.dynamic("max_num_blocks_per_seq") + @I.ir_module class ModulePagedAttentionV2: I.module_global_infos( @@ -112,14 +120,13 @@ class ModulePagedAttentionV2: @R.function def main( - query: R.Tensor(("num_seqs", 1, 64), dtype="float16"), - key_cache: R.Tensor(("num_blocks", 1, 8, 16, 8), dtype="float16"), - value_cache: R.Tensor(("num_blocks", 1, 64, 16), dtype="float16"), - block_tables: R.Tensor(("num_seqs", "max_num_blocks_per_seq"), dtype="int32"), - context_lens: R.Tensor(("num_seqs",), dtype="int32"), - ) -> R.Tensor(("num_seqs", 1, 64), dtype="float16"): + query: R.Tensor((num_seqs, 1, 64), dtype="float16"), + key_cache: R.Tensor((num_blocks, 1, 8, 16, 8), dtype="float16"), + value_cache: R.Tensor((num_blocks, 1, 64, 16), dtype="float16"), + block_tables: R.Tensor((num_seqs, max_num_blocks_per_seq), dtype="int32"), + context_lens: R.Tensor((num_seqs,), dtype="int32"), + ) -> R.Tensor((num_seqs, 1, 64), dtype="float16"): with R.dataflow(): - num_seqs = T.int64() max_len = R.to_vdevice(R.max(context_lens), "llvm:0") # alloc workspace exp_sums = R.zeros((num_seqs, 1, 1), "float32") @@ -344,19 +351,22 @@ def main( def test_cache(): + num_tokens = T.dynamic("num_tokens") + num_blocks = T.dynamic("num_blocks") + @I.ir_module class Module: @R.function def main( - key: R.Tensor(("num_tokens", 1, 8), dtype="float16"), - value: R.Tensor(("num_tokens", 1, 8), dtype="float16"), - key_cache: R.Tensor(("num_blocks", 1, 1, 16, 8), dtype="float16"), - value_cache: R.Tensor(("num_blocks", 1, 8, 16), dtype="float16"), - slot_mapping: R.Tensor(("num_tokens",), dtype="int32"), + key: R.Tensor((num_tokens, 1, 8), dtype="float16"), + value: R.Tensor((num_tokens, 1, 8), dtype="float16"), + key_cache: R.Tensor((num_blocks, 1, 1, 16, 8), dtype="float16"), + value_cache: R.Tensor((num_blocks, 1, 8, 16), dtype="float16"), + slot_mapping: R.Tensor((num_tokens,), dtype="int32"), ) -> R.Tuple( [ - R.Tensor(("num_blocks", 1, 8, 16, 8), dtype="float16"), - R.Tensor(("num_blocks", 1, 8, 16), dtype="float16"), + R.Tensor((num_blocks, 1, 8, 16, 8), dtype="float16"), + R.Tensor((num_blocks, 1, 8, 16), dtype="float16"), ] ): with R.dataflow(): diff --git a/tests/python/relax/test_dataflow_inplace.py b/tests/python/relax/test_dataflow_inplace.py index 19565d301c36..b78efa188f19 100644 --- a/tests/python/relax/test_dataflow_inplace.py +++ b/tests/python/relax/test_dataflow_inplace.py @@ -172,17 +172,20 @@ def main(x: R.Tensor((60,), "int32")) -> R.Tensor((15,), "int32"): def test_alias_call_tir(): # call TIR can yield either a single tensor or a tuple + m_tir_id = T.dynamic("m", "int32") + n_tir_id = T.dynamic("n", "int32") + m_tir_id2 = T.dynamic("m", "int32") + n_tir_id2 = T.dynamic("n", "int32") + @I.ir_module class AliasCallTir: @Ts.prim_func def tir_id(x: T.handle, y: T.handle) -> None: T.func_attr({"global_symbol": "tir_id"}) - m = T.int32() - n = T.int32() - A = T.match_buffer(x, (m, n), "int32") - B = T.match_buffer(y, (m, n), "int32") + A = T.match_buffer(x, (m_tir_id, n_tir_id), "int32") + B = T.match_buffer(y, (m_tir_id, n_tir_id), "int32") - for i, j in T.grid(m, n): + for i, j in T.grid(m_tir_id, n_tir_id): with Ts.sblock("id"): vi, vj = Ts.axis.remap("SS", [i, j]) B[vi, vj] = A[vi, vj] @@ -190,13 +193,11 @@ def tir_id(x: T.handle, y: T.handle) -> None: @Ts.prim_func def tir_id2(x: T.handle, y: T.handle, z: T.handle) -> None: T.func_attr({"global_symbol": "tir_id"}) - m = T.int32() - n = T.int32() - A = T.match_buffer(x, (m, n), "int32") - B = T.match_buffer(y, (m, n), "int32") - C = T.match_buffer(z, (m, n), "int32") + A = T.match_buffer(x, (m_tir_id2, n_tir_id2), "int32") + B = T.match_buffer(y, (m_tir_id2, n_tir_id2), "int32") + C = T.match_buffer(z, (m_tir_id2, n_tir_id2), "int32") - for i, j in T.grid(m, n): + for i, j in T.grid(m_tir_id2, n_tir_id2): with Ts.sblock("id"): vi, vj = Ts.axis.remap("SS", [i, j]) B[vi, vj] = A[vi, vj] @@ -545,12 +546,15 @@ def main( def test_dynamic(): + a = T.dynamic("a") + b = T.dynamic("b") + @I.ir_module class DynamicTestCase: @R.function def main( - x: R.Tensor(("a", "b"), dtype="float32"), y: R.Tensor(("a", "b"), dtype="float32") - ) -> R.Tensor(("a", "b"), dtype="float32"): + x: R.Tensor((a, b), dtype="float32"), y: R.Tensor((a, b), dtype="float32") + ) -> R.Tensor((a, b), dtype="float32"): with R.dataflow(): z = R.add(x, y) # Cannot be done in-place because x and y are arguments @@ -563,15 +567,21 @@ def main( transform_pass = DataflowUseInplaceCalls() new_mod = transform_pass(DynamicTestCase) + a_add_inplace = T.dynamic("a") + b_add_inplace = T.dynamic("b") + a_subtract_inplace = T.dynamic("a") + b_subtract_inplace = T.dynamic("b") + a_main = T.dynamic("a") + b_main = T.dynamic("b") + @I.ir_module class Expected: @Ts.prim_func(private=True) def add_inplace(var_A: T.handle, var_B: T.handle): T.func_attr({"tirx.noalias": True}) - a, b = T.int64(), T.int64() - A = T.match_buffer(var_A, (a, b)) - B = T.match_buffer(var_B, (a, b)) - for ax0, ax1 in T.grid(a, b): + A = T.match_buffer(var_A, (a_add_inplace, b_add_inplace)) + B = T.match_buffer(var_B, (a_add_inplace, b_add_inplace)) + for ax0, ax1 in T.grid(a_add_inplace, b_add_inplace): with Ts.sblock("T_add"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(A[v_ax0, v_ax1], B[v_ax0, v_ax1]) @@ -581,10 +591,9 @@ def add_inplace(var_A: T.handle, var_B: T.handle): @Ts.prim_func(private=True) def subtract_inplace(var_A: T.handle, var_B: T.handle): T.func_attr({"tirx.noalias": True}) - a, b = T.int64(), T.int64() - A = T.match_buffer(var_A, (a, b)) - B = T.match_buffer(var_B, (a, b)) - for ax0, ax1 in T.grid(a, b): + A = T.match_buffer(var_A, (a_subtract_inplace, b_subtract_inplace)) + B = T.match_buffer(var_B, (a_subtract_inplace, b_subtract_inplace)) + for ax0, ax1 in T.grid(a_subtract_inplace, b_subtract_inplace): with Ts.sblock("T_subtract"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(A[v_ax0, v_ax1], B[v_ax0, v_ax1]) @@ -593,23 +602,22 @@ def subtract_inplace(var_A: T.handle, var_B: T.handle): @R.function def main( - x: R.Tensor(("a", "b"), dtype="float32"), y: R.Tensor(("a", "b"), dtype="float32") - ) -> R.Tensor(("a", "b"), dtype="float32"): - a = T.int64() - b = T.int64() + x: R.Tensor((a_main, b_main), dtype="float32"), + y: R.Tensor((a_main, b_main), dtype="float32"), + ) -> R.Tensor((a_main, b_main), dtype="float32"): cls = Expected with R.dataflow(): z = R.add(x, y) a_1 = R.call_tir_inplace( cls.add_inplace, (z, y), - out_ty=R.Tensor((a, b), dtype="float32"), + out_ty=R.Tensor((a_main, b_main), dtype="float32"), inplace_indices=[0], ) s = R.call_tir_inplace( cls.subtract_inplace, (a_1, a_1), - out_ty=R.Tensor((a, b), dtype="float32"), + out_ty=R.Tensor((a_main, b_main), dtype="float32"), inplace_indices=[1], ) R.output(s) @@ -629,12 +637,15 @@ def main( def test_dynamic_mismatch(): # cannot statically prove the shapes to be equal so the module should be unchanged + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @I.ir_module class DynamicMistmatchTestCase: @R.function - def main( - x: R.Tensor(("a", "b"), dtype="float32"), y: R.Tensor(("c", "d"), dtype="float32") - ): + def main(x: R.Tensor((a, b), dtype="float32"), y: R.Tensor((c, d), dtype="float32")): with R.dataflow(): z = R.add(x, y) # Cannot be done in-place because x and y are arguments diff --git a/tests/python/relax/test_dataflow_pattern.py b/tests/python/relax/test_dataflow_pattern.py index 940345032569..e32c077bcd57 100644 --- a/tests/python/relax/test_dataflow_pattern.py +++ b/tests/python/relax/test_dataflow_pattern.py @@ -671,16 +671,20 @@ def main( def test_self_attention(): # The example comes from. # https://developer.nvidia.com/blog/nlu-with-tensorrt-bert/ + b = T.dynamic("b") + s = T.dynamic("s") + n = T.dynamic("n") + h = T.dynamic("h") + @tvm.script.ir_module class SelfAttention: @R.function def main( - x: R.Tensor(("b", "s", "n", "h"), "float32"), - wq: R.Tensor(("h", "h"), "float32"), - wk: R.Tensor(("h", "h"), "float32"), - wv: R.Tensor(("h", "h"), "float32"), + x: R.Tensor((b, s, n, h), "float32"), + wq: R.Tensor((h, h), "float32"), + wk: R.Tensor((h, h), "float32"), + wv: R.Tensor((h, h), "float32"), ) -> R.Tensor: - b, s, n, h = T.int64(), T.int64(), T.int64(), T.int64() with R.dataflow(): fcq = R.call_dps_packed("my_fc", (x, wq), R.Tensor((b, s, n, h), dtype="float32")) tpq = R.call_dps_packed( @@ -1506,11 +1510,12 @@ def func( return out elif same_shape_func_type == "same_dynamic_shape": + n = T.dynamic("n") @R.function(private=True) def func( - a: R.Tensor(("n", 128), "float32"), - b: R.Tensor(("n", 128), "float32"), + a: R.Tensor((n, 128), "float32"), + b: R.Tensor((n, 128), "float32"), ) -> R.Tensor: with R.dataflow(): c = R.multiply(a, R.const(2.0)) @@ -1534,11 +1539,13 @@ def func( return out elif same_shape_func_type == "different_dynamic_shape": + n = T.dynamic("n") + m = T.dynamic("m") @R.function(private=True) def func( - a: R.Tensor(("n", 128), "float32"), - b: R.Tensor(("m", 128), "float32"), + a: R.Tensor((n, 128), "float32"), + b: R.Tensor((m, 128), "float32"), ) -> R.Tensor: with R.dataflow(): c = R.multiply(a, R.const(2.0)) @@ -1889,8 +1896,8 @@ def test_wildcard_ty_with_symbolic_vars(): broadcasted `R.add`. """ - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") pat_lhs = wildcard().has_ty(R.Tensor([m, n])) pat_rhs = wildcard().has_ty(R.Tensor([m, n])) diff --git a/tests/python/relax/test_dataflow_rewriter.py b/tests/python/relax/test_dataflow_rewriter.py index bcfc002ef52c..db9d46b201ba 100644 --- a/tests/python/relax/test_dataflow_rewriter.py +++ b/tests/python/relax/test_dataflow_rewriter.py @@ -1225,13 +1225,20 @@ def test_match_dynamic_shape(): """ + N1_pattern = T.dynamic("N1") + M_pattern = T.dynamic("M") + N2_pattern = T.dynamic("N2") + N1_replacement = T.dynamic("N1") + N2_replacement = T.dynamic("N2") + M_replacement = T.dynamic("M") + @R.rewriter class Rewriter: @R.function def pattern( - lhs_A: R.Tensor(["N1", "M"], "float32"), - lhs_B: R.Tensor(["N2", "M"], "float32"), - rhs: R.Tensor(["M"], "float32"), + lhs_A: R.Tensor([N1_pattern, M_pattern], "float32"), + lhs_B: R.Tensor([N2_pattern, M_pattern], "float32"), + rhs: R.Tensor([M_pattern], "float32"), ): proj_A = R.matmul(lhs_A, rhs) proj_B = R.matmul(lhs_B, rhs) @@ -1239,20 +1246,17 @@ def pattern( @R.function def replacement( - lhs_A: R.Tensor(["N1", "M"], "float32"), - lhs_B: R.Tensor(["N2", "M"], "float32"), - rhs: R.Tensor(["M"], "float32"), + lhs_A: R.Tensor([N1_replacement, M_replacement], "float32"), + lhs_B: R.Tensor([N2_replacement, M_replacement], "float32"), + rhs: R.Tensor([M_replacement], "float32"), ): - N1 = T.int64() - N2 = T.int64() - lhs = R.concat([lhs_A, lhs_B]) proj_concat = R.matmul(lhs, rhs) - proj_A: R.Tensor([N1], "float32") = R.strided_slice( - proj_concat, axes=[0], begin=[0], end=[N1] + proj_A: R.Tensor([N1_replacement], "float32") = R.strided_slice( + proj_concat, axes=[0], begin=[0], end=[N1_replacement] ) - proj_B: R.Tensor([N2], "float32") = R.strided_slice( - proj_concat, axes=[0], begin=[N1], end=[N2 + N1] + proj_B: R.Tensor([N2_replacement], "float32") = R.strided_slice( + proj_concat, axes=[0], begin=[N1_replacement], end=[N2_replacement + N1_replacement] ) return (proj_A, proj_B) @@ -1267,15 +1271,16 @@ def before( out = proj_A + proj_B return out + N1 = T.dynamic("N1") + M = T.dynamic("M") + N2 = T.dynamic("N2") + @R.function(private=True) def expected( state: R.Tensor([16], "float32"), A: R.Tensor([16, 16], "float32"), B: R.Tensor([16, 16], "float32"), ): - N1 = T.int64() - M = T.int64() - N2 = T.int64() with R.dataflow(): lhs_A = R.match_cast(A, R.Tensor([N1, M], "float32")) lhs_B = R.match_cast(B, R.Tensor([N2, M], "float32")) @@ -1298,49 +1303,53 @@ def expected( def test_match_dynamic_pattern_against_dynamic_shape(): """A dynamic pattern may match a static shape""" + M_pattern = T.dynamic("M") + N_pattern = T.dynamic("N") + M_replacement = T.dynamic("M") + N_replacement = T.dynamic("N") + @R.rewriter class Rewriter: @R.function def pattern( - A: R.Tensor(["M", "N"], "float32"), - B: R.Tensor(["N", "N"], "float32"), + A: R.Tensor([M_pattern, N_pattern], "float32"), + B: R.Tensor([N_pattern, N_pattern], "float32"), ): return R.matmul(A, B) @R.function def replacement( - A: R.Tensor(["M", "N"], "float32"), - B: R.Tensor(["N", "N"], "float32"), + A: R.Tensor([M_replacement, N_replacement], "float32"), + B: R.Tensor([N_replacement, N_replacement], "float32"), ): - M = T.int64() - N = T.int64() return R.call_pure_packed( "my_optimized_square_matmul", A, B, - ty_args=R.Tensor([M, N], "float32"), + ty_args=R.Tensor([M_replacement, N_replacement], "float32"), ) + N = T.dynamic("N") + @R.function(private=True) def before( - A: R.Tensor(["N", "N*2"], "float32"), - B: R.Tensor(["N*2", "N*2"], "float32"), - C: R.Tensor(["N", "N"], "float32"), + A: R.Tensor([N, N * 2], "float32"), + B: R.Tensor([N * 2, N * 2], "float32"), + C: R.Tensor([N, N], "float32"), ): - N = T.int64() D: R.Tensor([N, N * 2], "float32") = R.matmul(A, B) E: R.Tensor([N * 2, N], "float32") = R.permute_dims(D) F: R.Tensor([N * 2, N], "float32") = R.matmul(E, C) return F + N = T.dynamic("N") + @R.function(private=True) def expected( - A: R.Tensor(["N", "N*2"], "float32"), - B: R.Tensor(["N*2", "N*2"], "float32"), - C: R.Tensor(["N", "N"], "float32"), + A: R.Tensor([N, N * 2], "float32"), + B: R.Tensor([N * 2, N * 2], "float32"), + C: R.Tensor([N, N], "float32"), ): - N = T.int64() - D: R.Tensor([N, N * 2], "float32") = R.call_pure_packed( "my_optimized_square_matmul", A, diff --git a/tests/python/relax/test_e2e_op_dynamic.py b/tests/python/relax/test_e2e_op_dynamic.py index 56628c7270e5..6c00e985b3a0 100644 --- a/tests/python/relax/test_e2e_op_dynamic.py +++ b/tests/python/relax/test_e2e_op_dynamic.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: E501, F401, F841 +# ruff: noqa: E501, F401 import numpy as np import pytest @@ -85,12 +85,13 @@ def main(x: R.Tensor((8, 9, 10, 10), "float32"), begin: R.Tensor((4,),"int64"), ) def test_dynamic_strided_slice_symbolic(begin, end, strides): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class DynamicStridedSlice: @R.function - def main(x: R.Tensor(("m", "n", 10, 10), "float32"), begin: R.Tensor((4,),"int64"), end: R.Tensor((4,),"int64"), strides: R.Tensor((4,),"int64")) -> R.Tensor("float32", ndim=4): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n, 10, 10), "float32"), begin: R.Tensor((4,),"int64"), end: R.Tensor((4,),"int64"), strides: R.Tensor((4,),"int64")) -> R.Tensor("float32", ndim=4): gv: R.Tensor("float32", ndim=4) = R.dynamic_strided_slice(x, begin, end, strides) return gv # fmt: on diff --git a/tests/python/relax/test_frontend_from_exported_program.py b/tests/python/relax/test_frontend_from_exported_program.py index 22d5fee0d707..e4fbfacca826 100644 --- a/tests/python/relax/test_frontend_from_exported_program.py +++ b/tests/python/relax/test_frontend_from_exported_program.py @@ -5505,14 +5505,15 @@ def forward(self, x): return x[:, :seq_len] + 0.0 # +0.0 to ensure output is a new tensor # The identity slice is elided; only x + 0.0 remains. + s0 = T.dynamic("s0") + s1 = T.dynamic("s1") + @I.ir_module class ExpectedIdentity: @R.function - def main(x: R.Tensor(("s0", "s1", 4), dtype="float32")) -> R.Tuple( - R.Tensor(("s0", "s1", 4), dtype="float32") + def main(x: R.Tensor((s0, s1, 4), dtype="float32")) -> R.Tuple( + R.Tensor((s0, s1, 4), dtype="float32") ): - s0 = T.int64() - s1 = T.int64() R.func_attr({"tir_var_lower_bound": {"s27": 2, "s77": 2}}) with R.dataflow(): lv: R.Tensor((s0, s1, 4), dtype="float32") = R.add(x, R.const(0.0, "float32")) @@ -6463,6 +6464,8 @@ class MaskedSelect(Module): def forward(self, data: torch.Tensor, mask: torch.Tensor): return torch.masked_select(data, mask) + u0 = T.dynamic("u0") + @tvm.script.ir_module class Expected: @R.function @@ -6470,7 +6473,6 @@ def main( data: R.Tensor((2, 3), dtype="float32"), mask: R.Tensor((2, 3), dtype="bool") ) -> R.Tuple(R.Tensor(dtype="float32", ndim=1)): R.func_attr({"tir_var_lower_bound": {"u0": 0}, "tir_var_upper_bound": {"u0": 6}}) - u0 = T.int64() with R.dataflow(): lv: R.Tensor((6,), dtype="float32") = R.reshape(data, R.shape([6])) lv1: R.Tensor((6,), dtype="bool") = R.reshape(mask, R.shape([6])) @@ -7787,14 +7789,15 @@ class DynamicModel(torch.nn.Module): def forward(self, x1, x2): return torch.ops.aten.add.Tensor(x1, x2) + s0 = T.dynamic("s0") + @I.ir_module class Expected: @R.function def main( - lhs: R.Tensor(("s0", 4), dtype="float32"), - rhs: R.Tensor(("s0", 4), dtype="float32"), - ) -> R.Tuple(R.Tensor(("s0", 4), dtype="float32")): - s0 = T.int64() + lhs: R.Tensor((s0, 4), dtype="float32"), + rhs: R.Tensor((s0, 4), dtype="float32"), + ) -> R.Tuple(R.Tensor((s0, 4), dtype="float32")): R.func_attr({"tir_var_lower_bound": {"s24": 0}}) with R.dataflow(): lv: R.Tensor((s0, 4), dtype="float32") = R.add(lhs, rhs) @@ -8388,13 +8391,14 @@ class DynamicModel(torch.nn.Module): def forward(self, x1, x2): return torch.ops.aten.add.Tensor(x1, x2) + s0 = T.dynamic("s0") + @I.ir_module class Expected: @R.function def main( - x1: R.Tensor(("s0", 4), dtype="float32"), x2: R.Tensor(("s0", 4), dtype="float32") - ) -> R.Tuple(R.Tensor(("s0", 4), dtype="float32")): - s0 = T.int64() + x1: R.Tensor((s0, 4), dtype="float32"), x2: R.Tensor((s0, 4), dtype="float32") + ) -> R.Tuple(R.Tensor((s0, 4), dtype="float32")): R.func_attr({"tir_var_lower_bound": {"s24": 1}, "tir_var_upper_bound": {"s24": 64}}) with R.dataflow(): lv: R.Tensor((s0, 4), dtype="float32") = R.add(x1, x2) @@ -8421,13 +8425,14 @@ class ConcatModel(torch.nn.Module): def forward(self, x, y): return torch.cat([x, y], dim=0) + s0 = T.dynamic("s0") + @I.ir_module class Expected: @R.function def main( - x: R.Tensor(("s0", 4), dtype="float32"), y: R.Tensor(("1 + s0", 4), dtype="float32") - ) -> R.Tuple(R.Tensor(("s0 + (1 + s0)", 4), dtype="float32")): - s0 = T.int64() + x: R.Tensor((s0, 4), dtype="float32"), y: R.Tensor((1 + s0, 4), dtype="float32") + ) -> R.Tuple(R.Tensor((s0 + (1 + s0), 4), dtype="float32")): R.func_attr( { "tir_var_lower_bound": {"s77": 1}, @@ -8454,13 +8459,14 @@ class ConcatModel(torch.nn.Module): def forward(self, x, y): return torch.cat([x, y], dim=0) + s0 = T.dynamic("s0") + @I.ir_module class Expected: @R.function def main( - x: R.Tensor(("s0", 4), dtype="float32"), y: R.Tensor(("2 * s0", 4), dtype="float32") - ) -> R.Tuple(R.Tensor(("s0 + 2 * s0", 4), dtype="float32")): - s0 = T.int64() + x: R.Tensor((s0, 4), dtype="float32"), y: R.Tensor((2 * s0, 4), dtype="float32") + ) -> R.Tuple(R.Tensor((s0 + 2 * s0, 4), dtype="float32")): R.func_attr( { "tir_var_lower_bound": {"s77": 1}, @@ -8487,13 +8493,14 @@ class DynamicModel(torch.nn.Module): def forward(self, x): return torch.ops.aten.add.Tensor(x, x) + s0 = T.dynamic("s0") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(("s0", 4), dtype="float32")) -> R.Tuple( - R.Tensor(("s0", 4), dtype="float32") + def main(x: R.Tensor((s0, 4), dtype="float32")) -> R.Tuple( + R.Tensor((s0, 4), dtype="float32") ): - s0 = T.int64() R.func_attr({"tir_var_lower_bound": {"s77": 2}}) with R.dataflow(): lv: R.Tensor((s0, 4), dtype="float32") = R.add(x, x) @@ -8521,13 +8528,14 @@ def forward(self, x): shape_dim = torch.ops.aten.sym_size.int(x, 0) return x.reshape(shape_dim, -1) + s0 = T.dynamic("s0") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(("s0", 3, 4), dtype="float32")) -> R.Tuple( - R.Tensor(("s0", 12), dtype="float32") + def main(x: R.Tensor((s0, 3, 4), dtype="float32")) -> R.Tuple( + R.Tensor((s0, 12), dtype="float32") ): - s0 = T.int64() R.func_attr({"tir_var_lower_bound": {"s77": 0}}) with R.dataflow(): lv: R.Tensor((s0, 12), dtype="float32") = R.reshape(x, R.shape([s0, 12])) @@ -8931,40 +8939,45 @@ def false_fn(x): return torch.cond(x.shape[0] > 4, true_fn, false_fn, (x,)) + s77_cond_true_branch_0 = T.dynamic("s77") + s77_cond_false_branch_1 = T.dynamic("s77") + s77_main = T.dynamic("s77") + @tvm.script.ir_module class expected: @R.function def cond_true_branch_0( - x: R.Tensor(("s77", 4), dtype="float32"), - ) -> R.Tensor(("s77", 4), dtype="float32"): - s77 = T.int64() - gv: R.Tensor((s77, 4), dtype="float32") = R.add(x, R.const(1.0, "float32")) - gv1: R.Tensor((s77, 4), dtype="float32") = gv + x: R.Tensor((s77_cond_true_branch_0, 4), dtype="float32"), + ) -> R.Tensor((s77_cond_true_branch_0, 4), dtype="float32"): + gv: R.Tensor((s77_cond_true_branch_0, 4), dtype="float32") = R.add( + x, R.const(1.0, "float32") + ) + gv1: R.Tensor((s77_cond_true_branch_0, 4), dtype="float32") = gv return gv1 @R.function def cond_false_branch_1( - x: R.Tensor(("s77", 4), dtype="float32"), - ) -> R.Tensor(("s77", 4), dtype="float32"): - s77 = T.int64() - gv: R.Tensor((s77, 4), dtype="float32") = R.subtract(x, R.const(1.0, "float32")) - gv1: R.Tensor((s77, 4), dtype="float32") = gv + x: R.Tensor((s77_cond_false_branch_1, 4), dtype="float32"), + ) -> R.Tensor((s77_cond_false_branch_1, 4), dtype="float32"): + gv: R.Tensor((s77_cond_false_branch_1, 4), dtype="float32") = R.subtract( + x, R.const(1.0, "float32") + ) + gv1: R.Tensor((s77_cond_false_branch_1, 4), dtype="float32") = gv return gv1 @R.function def main( - x: R.Tensor(("s77", 4), dtype="float32"), - ) -> R.Tuple(R.Tensor(("s77", 4), dtype="float32")): - s77 = T.int64() + x: R.Tensor((s77_main, 4), dtype="float32"), + ) -> R.Tuple(R.Tensor((s77_main, 4), dtype="float32")): R.func_attr({"tir_var_lower_bound": {"s77": 1}}) cls = expected - gv: T.bool = s77 > 4 + gv: T.bool = s77_main > 4 if gv: - gv1: R.Tensor((s77, 4), dtype="float32") = cls.cond_true_branch_0(x) - cond_result: R.Tensor((s77, 4), dtype="float32") = gv1 + gv1: R.Tensor((s77_main, 4), dtype="float32") = cls.cond_true_branch_0(x) + cond_result: R.Tensor((s77_main, 4), dtype="float32") = gv1 else: - gv2: R.Tensor((s77, 4), dtype="float32") = cls.cond_false_branch_1(x) - cond_result: R.Tensor((s77, 4), dtype="float32") = gv2 + gv2: R.Tensor((s77_main, 4), dtype="float32") = cls.cond_false_branch_1(x) + cond_result: R.Tensor((s77_main, 4), dtype="float32") = gv2 return (cond_result,) batch = torch.export.Dim("batch", min=1) diff --git a/tests/python/relax/test_frontend_nn_exporter.py b/tests/python/relax/test_frontend_nn_exporter.py index bb7dcf4e7858..5571393c06d9 100644 --- a/tests/python/relax/test_frontend_nn_exporter.py +++ b/tests/python/relax/test_frontend_nn_exporter.py @@ -125,10 +125,12 @@ def test_dynamic_shape(): debug=False, ) + batch_size = T.dynamic("batch_size") + @I.ir_module class Expected: @R.function - def forward(x: R.Tensor(["batch_size", 8], dtype="float32")): + def forward(x: R.Tensor([batch_size, 8], dtype="float32")): R.func_attr({"num_input": 1}) with R.dataflow(): relu = R.nn.relu(x) @@ -158,10 +160,13 @@ def forward_silu(self, x: nn.Tensor): debug=False, ) + batch_size_forward_relu = T.dynamic("batch_size") + batch_size_forward_silu = T.dynamic("batch_size") + @I.ir_module class Expected: @R.function - def forward_relu(x: R.Tensor(["batch_size", 8], dtype="float32")): + def forward_relu(x: R.Tensor([batch_size_forward_relu, 8], dtype="float32")): R.func_attr({"num_input": 1}) with R.dataflow(): relu = R.nn.relu(x) @@ -170,7 +175,7 @@ def forward_relu(x: R.Tensor(["batch_size", 8], dtype="float32")): return relu @R.function - def forward_silu(x: R.Tensor(["batch_size", 8], dtype="float32")): + def forward_silu(x: R.Tensor([batch_size_forward_silu, 8], dtype="float32")): R.func_attr({"num_input": 1}) with R.dataflow(): silu = R.nn.silu(x) @@ -227,17 +232,18 @@ def forward(self, x: nn.Tensor): debug=False, ) + batch_size = T.dynamic("batch_size") + @I.ir_module class Expected: @R.function def forward( - x: R.Tensor(["batch_size", hidden_size], "float16"), + x: R.Tensor([batch_size, hidden_size], "float16"), gate_proj_weights: R.Tensor([intermediate_size, hidden_size], "float16"), up_proj_weights: R.Tensor([intermediate_size, hidden_size], "float16"), down_proj_weights: R.Tensor([hidden_size, intermediate_size], "float16"), ): R.func_attr({"num_input": 1}) - batch_size = T.int64() with R.dataflow(): gate: R.Tensor([batch_size, intermediate_size]) = R.matmul( x, R.permute_dims(gate_proj_weights) @@ -351,11 +357,13 @@ def forward(self, x: nn.Tensor): debug=False, ) + batch_size = T.dynamic("batch_size") + @I.ir_module class Expected: @R.function def forward( - x: R.Tensor(["batch_size", hidden_size], "float16"), + x: R.Tensor([batch_size, hidden_size], "float16"), # The function's parameters are defined by the # `nn.Parameter` instances, and still reference the # original `gate_proj` and `up_proj` weights. This @@ -366,7 +374,6 @@ def forward( down_proj_weights: R.Tensor([hidden_size, intermediate_size], "float16"), ): R.func_attr({"num_input": 1}) - batch_size = T.int64() with R.dataflow(): # At this stage of compilation, the concatenation is # written within the body of the function. This will @@ -389,11 +396,13 @@ def forward( assert_structural_equal(exported_mod, Expected) + batch_size = T.dynamic("batch_size") + @I.ir_module class ExpectedAfterLift: @R.function def forward( - x: R.Tensor(["batch_size", hidden_size], "float16"), + x: R.Tensor([batch_size, hidden_size], "float16"), # After `relax.transform.LiftTransformParams`, the # `gate_proj` and `up_proj` weights have been concatenated # together. @@ -403,7 +412,6 @@ def forward( down_proj_weights_transpose: R.Tensor([intermediate_size, hidden_size], "float16"), ): R.func_attr({"num_input": 1}) - batch_size = T.int64() with R.dataflow(): gate_up: R.Tensor([batch_size, intermediate_size * 2], "float16") = R.matmul( x, gate_up_proj_weights_transpose @@ -459,14 +467,15 @@ def test_linear_dynamic_shape(): Even if dynamic, the weight/bias must be the same value. """ + n = T.dynamic("n") + @R.function def forward( x: R.Tensor((1, 4), dtype="float32"), _io: R.Any, - weight: R.Tensor(("n", 4), dtype="float32"), - bias: R.Tensor(("n",), dtype="float32"), - ) -> R.Tuple(R.Tensor((1, "n"), dtype="float32"), R.Tuple(R.Any)): - n = T.int64() + weight: R.Tensor((n, 4), dtype="float32"), + bias: R.Tensor((n,), dtype="float32"), + ) -> R.Tuple(R.Tensor((1, n), dtype="float32"), R.Tuple(R.Any)): R.func_attr({"num_input": 2}) with R.dataflow(): permute_dims: R.Tensor((4, n), dtype="float32") = R.permute_dims(weight, axes=None) @@ -563,19 +572,20 @@ def forward(self, state: nn.Tensor): ) def get_expected_with_intermediate_size(): + batch_size = T.dynamic("batch_size") + hidden_size = T.dynamic("hidden_size") + intermediate_size = T.dynamic("intermediate_size") + @I.ir_module class Expected: @R.function def forward( - state: R.Tensor(["batch_size", 1024], "float32"), - embedding_weights: R.Tensor(["hidden_size", 1024], "float32"), - up_weights: R.Tensor(["intermediate_size", "hidden_size"], "float32"), - down_weights: R.Tensor(["hidden_size", "intermediate_size"], "float32"), + state: R.Tensor([batch_size, 1024], "float32"), + embedding_weights: R.Tensor([hidden_size, 1024], "float32"), + up_weights: R.Tensor([intermediate_size, hidden_size], "float32"), + down_weights: R.Tensor([hidden_size, intermediate_size], "float32"), ): R.func_attr({"num_input": 1}) - batch_size = T.int64() - hidden_size = T.int64() - intermediate_size = T.int64() with R.dataflow(): state: R.Tensor([batch_size, hidden_size], "float32") = R.matmul( state, R.permute_dims(embedding_weights) @@ -594,18 +604,19 @@ def forward( return Expected def get_expected_without_intermediate_size(): + batch_size = T.dynamic("batch_size") + hidden_size = T.dynamic("hidden_size") + @I.ir_module class Expected: @R.function def forward( - state: R.Tensor(["batch_size", 1024], "float32"), - embedding_weights: R.Tensor(["hidden_size", 1024], "float32"), - up_weights: R.Tensor(["hidden_size", "hidden_size"], "float32"), - down_weights: R.Tensor(["hidden_size", "hidden_size"], "float32"), + state: R.Tensor([batch_size, 1024], "float32"), + embedding_weights: R.Tensor([hidden_size, 1024], "float32"), + up_weights: R.Tensor([hidden_size, hidden_size], "float32"), + down_weights: R.Tensor([hidden_size, hidden_size], "float32"), ): R.func_attr({"num_input": 1}) - batch_size = T.int64() - hidden_size = T.int64() with R.dataflow(): state: R.Tensor([batch_size, hidden_size], "float32") = R.matmul( state, R.permute_dims(embedding_weights) diff --git a/tests/python/relax/test_frontend_nn_extern_module.py b/tests/python/relax/test_frontend_nn_extern_module.py index f0c0d91e904f..49d0e4348093 100644 --- a/tests/python/relax/test_frontend_nn_extern_module.py +++ b/tests/python/relax/test_frontend_nn_extern_module.py @@ -86,6 +86,10 @@ def _check_ir_equality(mod): # pylint: enable=import-outside-toplevel + x = T.dynamic("x") + y = T.dynamic("y") + z = T.dynamic("z") + @I.ir_module class ExpectedModule: @R.function @@ -103,11 +107,8 @@ def scalar_add( @R.function def test_sym( - a: R.Tensor(("x", "y", 1), dtype="float32"), b: R.Tensor(("y", "z", 5), dtype="float32") - ) -> R.Tensor(("x", "y", "z", 9), dtype="float32"): - x = T.int64() - y = T.int64() - z = T.int64() + a: R.Tensor((x, y, 1), dtype="float32"), b: R.Tensor((y, z, 5), dtype="float32") + ) -> R.Tensor((x, y, z, 9), dtype="float32"): R.func_attr({"num_input": 2}) with R.dataflow(): ext_test_sym = R.call_dps_packed( diff --git a/tests/python/relax/test_frontend_nn_modules.py b/tests/python/relax/test_frontend_nn_modules.py index 869016e7f4a3..ae29cb290f13 100644 --- a/tests/python/relax/test_frontend_nn_modules.py +++ b/tests/python/relax/test_frontend_nn_modules.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: E501, F401, F841 +# ruff: noqa: E501, F401 import numpy as np import pytest @@ -282,18 +282,19 @@ def forward( def test_conv2d_dynamic(): + n = T.dynamic("n") + h = T.dynamic("h") + w = T.dynamic("w") + c = T.dynamic("c") + in_channels = T.dynamic("in_channels") + @R.function def forward( - x: R.Tensor(("n", "c", "h", "w"), dtype="float32"), + x: R.Tensor((n, c, h, w), dtype="float32"), _io: R.Any, - weight: R.Tensor((32, "in_channels", 3, 3), dtype="float32"), + weight: R.Tensor((32, in_channels, 3, 3), dtype="float32"), bias: R.Tensor((32,), dtype="float32"), - ) -> R.Tuple(R.Tensor(("n", 32, "h - 2", "w - 2"), dtype="float32"), R.Tuple(R.Any)): - n = T.int64() - h = T.int64() - w = T.int64() - c = T.int64() - in_channels = T.int64() + ) -> R.Tuple(R.Tensor((n, 32, h - 2, w - 2), dtype="float32"), R.Tuple(R.Any)): R.func_attr({"num_input": 2}) with R.dataflow(): lv1: R.Tensor((n, 32, h - 2, w - 2), dtype="float32") = R.nn.conv2d(x, weight) diff --git a/tests/python/relax/test_frontend_nn_op.py b/tests/python/relax/test_frontend_nn_op.py index dbc55ce7721d..8babea026199 100644 --- a/tests/python/relax/test_frontend_nn_op.py +++ b/tests/python/relax/test_frontend_nn_op.py @@ -636,6 +636,9 @@ def test_tensor_ir_op(): fused_heads = num_q_heads + num_kv_heads * 2 dtype = "float16" + batch_size = T.dynamic("batch_size") + seq_len = T.dynamic("seq_len") + @Ts.prim_func(private=True) def fused_rope( # pylint: disable=too-many-locals var_qkv: T.handle, @@ -644,8 +647,6 @@ def fused_rope( # pylint: disable=too-many-locals var_k: T.handle, var_v: T.handle, ): - batch_size = T.int64() - seq_len = T.int64() qkv = T.match_buffer(var_qkv, (batch_size, seq_len, fused_heads, head_dim), dtype) q = T.match_buffer(var_q, (batch_size, seq_len, num_q_heads, head_dim), dtype) k = T.match_buffer(var_k, (batch_size, seq_len, num_kv_heads, head_dim), dtype) @@ -667,11 +668,14 @@ def test(self, qkv: Tensor, offset: tirx.Var): return tensor_expr_op_out # fmt: off + batch_size = T.dynamic("batch_size") + seq_len = T.dynamic("seq_len") + offset_1 = T.dynamic("offset_1") + @I.ir_module class Expected: @Ts.prim_func(private=True) def llama_fused_rope(var_qkv: T.handle, offset: T.int64, var_q: T.handle, var_k: T.handle, var_v: T.handle): - batch_size, seq_len = T.int64(), T.int64() qkv = T.match_buffer(var_qkv, (batch_size, seq_len, 24, 16), "float16") q = T.match_buffer(var_q, (batch_size, seq_len, 8, 16), "float16") k = T.match_buffer(var_k, (batch_size, seq_len, 8, 16), "float16") @@ -688,8 +692,7 @@ def _initialize_effect() -> R.Tuple(R.Any): return gv @R.function - def test(qkv: R.Tensor((1, 1, 24, 16), dtype="float16"), offset: R.Shape(["offset_1"]), _io: R.Any) -> R.Tuple(R.Tuple(R.Tensor((1, 1, 8, 16), dtype="float16"), R.Tensor((1, 1, 8, 16), dtype="float16"), R.Tensor((1, 1, 8, 16), dtype="float16")), R.Tuple(R.Any)): - offset_1 = T.int64() + def test(qkv: R.Tensor((1, 1, 24, 16), dtype="float16"), offset: R.Shape([offset_1]), _io: R.Any) -> R.Tuple(R.Tuple(R.Tensor((1, 1, 8, 16), dtype="float16"), R.Tensor((1, 1, 8, 16), dtype="float16"), R.Tensor((1, 1, 8, 16), dtype="float16")), R.Tuple(R.Any)): R.func_attr({"num_input": 3}) cls = Expected with R.dataflow(): @@ -716,15 +719,16 @@ def test_tensor_ir_inplace_op(): hidden_size = 4096 dtype = "float16" + vocab_size = T.dynamic("vocab_size") + seq_len = T.dynamic("seq_len") + total_seq_len = T.dynamic("total_seq_len") + @Ts.prim_func def inplace_take( var_weight: T.handle, var_pos: T.handle, var_embeddings: T.handle, offset: T.int64 ): T.func_attr({"tirx.noalias": True}) - vocab_size = T.int64() weight = T.match_buffer(var_weight, (vocab_size, hidden_size), dtype) - seq_len = T.int64() - total_seq_len = T.int64() pos = T.match_buffer(var_pos, (seq_len,), "int32") embeddings = T.match_buffer(var_embeddings, (total_seq_len, hidden_size), dtype) for ax0, ax1 in T.grid(seq_len, hidden_size): @@ -747,6 +751,14 @@ def test( ) return tensor_expr_op_out + vocab_size_inplace_take = T.dynamic("vocab_size") + seq_len_inplace_take = T.dynamic("seq_len") + total_seq_len_inplace_take = T.dynamic("total_seq_len") + total_seq_len_test = T.dynamic("total_seq_len") + offset_1 = T.dynamic("offset_1") + vocab_size_test = T.dynamic("vocab_size") + seq_len_test = T.dynamic("seq_len") + @I.ir_module class Expected: @Ts.prim_func @@ -754,13 +766,12 @@ def inplace_take( var_weight: T.handle, var_pos: T.handle, var_embeddings: T.handle, offset: T.int64 ): T.func_attr({"tirx.noalias": True}) - vocab_size = T.int64() - weight = T.match_buffer(var_weight, (vocab_size, hidden_size), dtype) - seq_len = T.int64() - total_seq_len = T.int64() - pos = T.match_buffer(var_pos, (seq_len,), "int32") - embeddings = T.match_buffer(var_embeddings, (total_seq_len, hidden_size), dtype) - for ax0, ax1 in T.grid(seq_len, hidden_size): + weight = T.match_buffer(var_weight, (vocab_size_inplace_take, hidden_size), dtype) + pos = T.match_buffer(var_pos, (seq_len_inplace_take,), "int32") + embeddings = T.match_buffer( + var_embeddings, (total_seq_len_inplace_take, hidden_size), dtype + ) + for ax0, ax1 in T.grid(seq_len_inplace_take, hidden_size): with Ts.sblock("T_take"): v0, v1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(weight[pos[v0], v1], pos[v0]) @@ -778,24 +789,22 @@ def _initialize_effect() -> R.Tuple(R.Any): @R.function def test( - embedding_table: R.Tensor(("vocab_size", hidden_size), dtype), - input_ids: R.Tensor(("seq_len",), "int32"), - embedding_dst: R.Tensor(("total_seq_len", hidden_size), dtype), - offset: R.Shape(["offset_1"]), + embedding_table: R.Tensor((vocab_size_test, hidden_size), dtype), + input_ids: R.Tensor((seq_len_test,), "int32"), + embedding_dst: R.Tensor((total_seq_len_test, hidden_size), dtype), + offset: R.Shape([offset_1]), packed_params: R.Tuple, - ) -> R.Tensor(("total_seq_len", hidden_size), dtype): - total_seq_len = T.int64() - offset_1 = T.int64() + ) -> R.Tensor((total_seq_len_test, hidden_size), dtype): R.func_attr({"num_input": 4}) cls = Expected with R.dataflow(): lv1 = R.call_tir_inplace( cls.inplace_take, (embedding_table, input_ids, embedding_dst, offset_1), - out_ty=R.Tensor((total_seq_len, hidden_size), dtype), + out_ty=R.Tensor((total_seq_len_test, hidden_size), dtype), inplace_indices=[2], ) - gv1: R.Tensor((total_seq_len, hidden_size), dtype) = lv1 + gv1: R.Tensor((total_seq_len_test, hidden_size), dtype) = lv1 R.output(gv1) return gv1 @@ -1016,25 +1025,29 @@ def foo( return z0 # fmt: off + batch_get_index_from_sorted = T.dynamic("batch") + vocab_size_get_index_from_sorted = T.dynamic("vocab_size") + out_batch = T.dynamic("out_batch") + batch_get_renorm_prob = T.dynamic("batch") + vocab_size_get_renorm_prob = T.dynamic("vocab_size") + @I.ir_module class Expected: @Ts.prim_func(private=True) def get_index_from_sorted(A: T.handle, B: T.handle, C: T.handle, D: T.handle, E: T.handle, F: T.handle): - batch, vocab_size = T.int64(), T.int64() - cumsum_sorted = T.match_buffer(A, (batch, vocab_size)) - indices = T.match_buffer(B, (batch, vocab_size), "int64") - renorm_prob = T.match_buffer(C, (batch, 1)) - out_batch = T.int64() + cumsum_sorted = T.match_buffer(A, (batch_get_index_from_sorted, vocab_size_get_index_from_sorted)) + indices = T.match_buffer(B, (batch_get_index_from_sorted, vocab_size_get_index_from_sorted), "int64") + renorm_prob = T.match_buffer(C, (batch_get_index_from_sorted, 1)) usample = T.match_buffer(D, (out_batch, 1)) sample_indices = T.match_buffer(E, (out_batch, 1), "int64") output_index = T.match_buffer(F, (out_batch, 1), "int64") # with Ts.sblock("root"): - for ax0, ax1 in T.grid(out_batch, vocab_size): + for ax0, ax1 in T.grid(out_batch, vocab_size_get_index_from_sorted): with Ts.sblock("T_get_index_from_sorted"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(usample[v_ax0, T.int64(0)], cumsum_sorted[sample_indices[v_ax0, T.int64(0)], v_ax1 - T.int64(1):v_ax1 - T.int64(1) + T.int64(2)], sample_indices[v_ax0, T.int64(0)], renorm_prob[sample_indices[v_ax0, T.int64(0)], 0], indices[sample_indices[v_ax0, T.int64(0)], T.min(T.int64(0), v_ax1):T.min(T.int64(0), v_ax1) + (T.max(T.int64(0), v_ax1) + T.int64(1) - T.min(T.int64(0), v_ax1))]) Ts.writes(output_index[v_ax0, 0]) - if usample[v_ax0, T.int64(0)] < cumsum_sorted[sample_indices[v_ax0, T.int64(0)], v_ax1] / renorm_prob[sample_indices[v_ax0, T.int64(0)], 0] or v_ax1 + T.int64(1) == vocab_size: + if usample[v_ax0, T.int64(0)] < cumsum_sorted[sample_indices[v_ax0, T.int64(0)], v_ax1] / renorm_prob[sample_indices[v_ax0, T.int64(0)], 0] or v_ax1 + T.int64(1) == vocab_size_get_index_from_sorted: if v_ax1 == T.int64(0): output_index[v_ax0, 0] = indices[sample_indices[v_ax0, T.int64(0)], 0] else: @@ -1043,13 +1056,12 @@ def get_index_from_sorted(A: T.handle, B: T.handle, C: T.handle, D: T.handle, E: @Ts.prim_func(private=True) def get_renorm_prob(A: T.handle, B: T.handle, C: T.handle, D: T.handle): - batch, vocab_size = T.int64(), T.int64() - cumsum_sorted = T.match_buffer(A, (batch, vocab_size)) - top_p = T.match_buffer(B, (batch, 1)) - top_k = T.match_buffer(C, (batch, 1), "int64") - renorm_prob = T.match_buffer(D, (batch, 1)) + cumsum_sorted = T.match_buffer(A, (batch_get_renorm_prob, vocab_size_get_renorm_prob)) + top_p = T.match_buffer(B, (batch_get_renorm_prob, 1)) + top_k = T.match_buffer(C, (batch_get_renorm_prob, 1), "int64") + renorm_prob = T.match_buffer(D, (batch_get_renorm_prob, 1)) # with Ts.sblock("root"): - for ax0, ax1 in T.grid(batch, vocab_size): + for ax0, ax1 in T.grid(batch_get_renorm_prob, vocab_size_get_renorm_prob): with Ts.sblock("T_get_renorm_prob"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(cumsum_sorted[v_ax0, T.min(T.min(T.int64(0), v_ax1), v_ax1 + T.int64(1)):T.min(T.min(T.int64(0), v_ax1), v_ax1 + T.int64(1)) + (T.max(T.max(T.int64(0), v_ax1), v_ax1 + T.int64(1)) + T.int64(1) - T.min(T.min(T.int64(0), v_ax1), v_ax1 + T.int64(1)))], top_p[v_ax0, 0], top_k[v_ax0, 0]) @@ -1058,7 +1070,7 @@ def get_renorm_prob(A: T.handle, B: T.handle, C: T.handle, D: T.handle): renorm_prob[v_ax0, 0] = cumsum_sorted[v_ax0, 0] else: if cumsum_sorted[v_ax0, v_ax1] < top_p[v_ax0, 0] and v_ax1 + T.int64(1) < top_k[v_ax0, 0]: - if v_ax1 + T.int64(1) == vocab_size: + if v_ax1 + T.int64(1) == vocab_size_get_renorm_prob: renorm_prob[v_ax0, 0] = cumsum_sorted[v_ax0, v_ax1] else: if not (cumsum_sorted[v_ax0, v_ax1 + T.int64(1)] < top_p[v_ax0, 0] and v_ax1 + T.int64(1) + T.int64(1) < top_k[v_ax0, 0]): @@ -1146,6 +1158,9 @@ def foo( return z0 # fmt: off + batch = T.dynamic("batch") + vocab_size = T.dynamic("vocab_size") + @I.ir_module class Expected: @Ts.prim_func(private=True) @@ -1161,7 +1176,6 @@ def filter_with_top_p_top_k(A: T.Buffer((T.int64(2), T.int64(3)), "float32"), B: @Ts.prim_func(private=True) def get_renorm_cutoff(A: T.handle, B: T.handle, C: T.handle, D: T.handle, E: T.handle): - batch, vocab_size = T.int64(), T.int64() sorted_prob = T.match_buffer(A, (batch, vocab_size)) cumsum_sorted = T.match_buffer(B, (batch, vocab_size)) top_p = T.match_buffer(C, (batch, 1)) @@ -1259,10 +1273,12 @@ def foo(self, x: Tensor): z2 = op.topk(x, k=2, axis=-1) return z0, z1, z2 + seq_len = T.dynamic("seq_len") + @I.ir_module class Expected: @R.function - def foo(x: R.Tensor(("seq_len", 64), dtype="float16")): + def foo(x: R.Tensor((seq_len, 64), dtype="float16")): R.func_attr({"num_input": 1}) with R.dataflow(): sort = R.sort(x, axis=-1, descending=True) diff --git a/tests/python/relax/test_frontend_nn_subroutines.py b/tests/python/relax/test_frontend_nn_subroutines.py index a6fdf1958d83..9774140c05b7 100644 --- a/tests/python/relax/test_frontend_nn_subroutines.py +++ b/tests/python/relax/test_frontend_nn_subroutines.py @@ -42,14 +42,18 @@ def forward(self, input: relax.Expr) -> relax.Var: state = nn.op.matmul(input, self.weights) return self.activation(state) + batch_size_forward = I.dynamic("batch_size") + batch_size_layer = I.dynamic("batch_size") + batch_size_activation = I.dynamic("batch_size") + @I.ir_module class Expected: @R.function def forward( - state: R.Tensor(("batch_size", 64), dtype="float32"), + state: R.Tensor((batch_size_forward, 64), dtype="float32"), _io: R.Any, weights: R.Tensor((64, 32), dtype="float32"), - ) -> R.Tuple(R.Tensor(("batch_size", 32), dtype="float32"), R.Tuple(R.Any)): + ) -> R.Tuple(R.Tensor((batch_size_forward, 32), dtype="float32"), R.Tuple(R.Any)): R.func_attr({"num_input": 2}) with R.dataflow(): state = Expected.layer(state, weights) @@ -69,9 +73,9 @@ def _initialize_effect() -> R.Tuple(R.Any): @R.function(private=True) def layer( - state: R.Tensor(("batch_size", 64), dtype="float32"), + state: R.Tensor((batch_size_layer, 64), dtype="float32"), weights: R.Tensor((64, 32), dtype="float32"), - ) -> R.Tensor(("batch_size", 32), dtype="float32"): + ) -> R.Tensor((batch_size_layer, 32), dtype="float32"): with R.dataflow(): state = R.matmul(state, weights) state = Expected.activation(state) @@ -81,8 +85,8 @@ def layer( @R.function(private=True) def activation( - state: R.Tensor(("batch_size", 32), dtype="float32"), - ) -> R.Tensor(("batch_size", 32), dtype="float32"): + state: R.Tensor((batch_size_activation, 32), dtype="float32"), + ) -> R.Tensor((batch_size_activation, 32), dtype="float32"): with R.dataflow(): state = R.nn.silu(state) dataflow_output = state diff --git a/tests/python/relax/test_frontend_onnx.py b/tests/python/relax/test_frontend_onnx.py index abef4b7cf55b..8640c3d96389 100644 --- a/tests/python/relax/test_frontend_onnx.py +++ b/tests/python/relax/test_frontend_onnx.py @@ -1055,14 +1055,15 @@ def _make_expected_broadcast_ir_min( Expected IR module for the Min operation. """ + n = T.dynamic("n") + @I.ir_module class ExpectedMin: @R.function def main( - x: R.Tensor(("n", x_shape[1]), dtype="float32"), - y: R.Tensor(("n", y_shape[1]), dtype="float32"), - ) -> R.Tensor(("n", 4), dtype="float32"): - n = T.int64() + x: R.Tensor((n, x_shape[1]), dtype="float32"), + y: R.Tensor((n, y_shape[1]), dtype="float32"), + ) -> R.Tensor((n, 4), dtype="float32"): R.func_attr({"num_input": 2}) with R.dataflow(): lv = R.broadcast_to(x, R.shape((n, 4))) @@ -1089,14 +1090,15 @@ def _make_expected_broadcast_ir_max( Expected IR module for the Max operation. """ + n = T.dynamic("n") + @I.ir_module class ExpectedMax: @R.function def main( - x: R.Tensor(("n", x_shape[1]), dtype="float32"), - y: R.Tensor(("n", y_shape[1]), dtype="float32"), - ) -> R.Tensor(("n", 4), dtype="float32"): - n = T.int64() + x: R.Tensor((n, x_shape[1]), dtype="float32"), + y: R.Tensor((n, y_shape[1]), dtype="float32"), + ) -> R.Tensor((n, 4), dtype="float32"): R.func_attr({"num_input": 2}) with R.dataflow(): lv = R.broadcast_to(x, R.shape((n, 4))) @@ -2407,7 +2409,7 @@ def test_scatter_dynamic_shape(): (shared symbolic dims) it still lowers via scatter_elements and is correct; a dynamic indices that is *not* provably equal (e.g. a broadcast size-1 dim) must raise instead of silently emitting wrong values.""" - n = tvm.tirx.Var("N", "int64") + n = T.dynamic("N", "int64") rng = np.random.RandomState(1) batch = 5 @@ -2665,6 +2667,8 @@ def make_expected(tensor_shape: list[int], condition_shape: list[int] | None, ax if axis is None: flat_shape = (int(np.prod(tensor_shape)),) + num_nonzero = T.dynamic("num_nonzero") + @I.ir_module class ExpectedCompressFlat: @R.function @@ -2672,7 +2676,6 @@ def main( tensor: R.Tensor(tensor_shape, dtype="float32"), condition: R.Tensor(condition_shape, dtype="bool"), ): - num_nonzero = T.int64() R.func_attr({"num_input": 2}) with R.dataflow(): lv: R.Tensor((1, num_nonzero), dtype="int64") = R.match_cast( @@ -2688,6 +2691,8 @@ def main( return ExpectedCompressFlat + num_nonzero = T.dynamic("num_nonzero") + @I.ir_module class ExpectedCompressAxis: @R.function @@ -2695,7 +2700,6 @@ def main( tensor: R.Tensor(tensor_shape, dtype="float32"), condition: R.Tensor(condition_shape, dtype="bool"), ): - num_nonzero = T.int64() R.func_attr({"num_input": 2}) with R.dataflow(): lv: R.Tensor((1, num_nonzero), dtype="int64") = R.match_cast( @@ -3377,6 +3381,11 @@ def test_unsqueeze_dynamic_axes_ir(): model = helper.make_model(graph, producer_name="unsqueeze_dynamic_axes_ir_test") tvm_model = from_onnx(model, opset=13, keep_params_in_input=True) + unsqueeze_dim_0 = T.dynamic("unsqueeze_dim_0") + unsqueeze_dim_1 = T.dynamic("unsqueeze_dim_1") + unsqueeze_dim_2 = T.dynamic("unsqueeze_dim_2") + unsqueeze_dim_3 = T.dynamic("unsqueeze_dim_3") + @I.ir_module class Expected: @R.function @@ -3385,10 +3394,6 @@ def main( axes: R.Tensor((2,), dtype="int64"), ) -> R.Tensor(dtype="float32", ndim=4): R.func_attr({"num_input": 2}) - unsqueeze_dim_0 = T.int64() - unsqueeze_dim_1 = T.int64() - unsqueeze_dim_2 = T.int64() - unsqueeze_dim_3 = T.int64() with R.dataflow(): lv: R.Shape([32, 32]) = R.shape_of(a) lv1: R.Tensor((2,), dtype="bool") = R.less(axes, R.const(0, "int64")) @@ -4207,13 +4212,14 @@ def test_shape_start_end_symbolic(): keep_params_in_input=True, ) + B = T.dynamic("B") + @I.ir_module class Expected: @R.function def main( - data: R.Tensor((3, "B", 5, 6), dtype="float32"), + data: R.Tensor((3, B, 5, 6), dtype="float32"), ) -> R.Shape(ndim=2): - B = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Shape([B, 5]) = R.shape([B, 5]) @@ -4413,13 +4419,13 @@ def main( expected = ExpectedTriu else: + m = T.dynamic("m") @I.ir_module class ExpectedTril: @R.function - def main(x: R.Tensor(("m", 3), dtype="float32")) -> R.Tensor(("m", 3), dtype="float32"): + def main(x: R.Tensor((m, 3), dtype="float32")) -> R.Tensor((m, 3), dtype="float32"): R.func_attr({"num_input": 1}) - m = T.int64() with R.dataflow(): lv: R.Tensor((1,), dtype="int64") = R.shape_to_tensor(R.shape([m])) lv1: R.Tensor((3,), dtype="int64") = R.arange(0, 3, 1, dtype="int64") @@ -5496,15 +5502,16 @@ def test_dynamic_squeeze(): tvm_model = from_onnx(model, opset=13, keep_params_in_input=True) tvm_model["main"] = tvm_model["main"].without_attr("params") + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class Expected: @R.function def main( - x: R.Tensor((1, "A", "B"), dtype="float32"), + x: R.Tensor((1, A, B), dtype="float32"), axes: R.Tensor((1,), dtype="int64"), - ) -> R.Tensor(("A", "B"), dtype="float32"): - A = T.int64() - B = T.int64() + ) -> R.Tensor((A, B), dtype="float32"): R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((A, B), dtype="float32") = R.squeeze(x, axis=[0]) @@ -5531,6 +5538,10 @@ def test_squeeze_dynamic_axes_ir(): model = helper.make_model(graph, producer_name="squeeze_dynamic_axes_ir_test") tvm_model = from_onnx(model, opset=13, keep_params_in_input=True) + squeeze_num_keep_dims = T.dynamic("squeeze_num_keep_dims") + squeeze_dim_0 = T.dynamic("squeeze_dim_0") + squeeze_dim_1 = T.dynamic("squeeze_dim_1") + @I.ir_module class Expected: @R.function @@ -5539,9 +5550,6 @@ def main( axes: R.Tensor((2,), dtype="int64"), ) -> R.Tensor(dtype="float32", ndim=2): R.func_attr({"num_input": 2}) - squeeze_num_keep_dims = T.int64() - squeeze_dim_0 = T.int64() - squeeze_dim_1 = T.int64() with R.dataflow(): lv: R.Shape([1, 32, 1, 32]) = R.shape_of(x) lv1: R.Tensor((2,), dtype="bool") = R.less(axes, R.const(0, "int64")) @@ -5624,7 +5632,7 @@ def test_dynamic_shape_squeeze(axis): tvm_model["main"] = tvm_model["main"].without_attr("params") # Use an ordinary symbolic Var for the dynamic shape binding. - a = tvm.tirx.Var("A", "int64") + a = T.dynamic("A", "int64") x = relax.Var("x", relax.TensorType([a], "float32")) axes = relax.Var("axes", relax.TensorType([1], "int64")) gv = relax.Var("gv", tvm.ir.PrimType("int64")) @@ -7239,14 +7247,15 @@ def main(in_: R.Tensor((3, 1), dtype="float32")) -> R.Tensor((1, 1, 3, 1), dtype R.output(gv) return gv + batch = T.dynamic("batch") + @I.ir_module class ExpectedDynamicShape: @R.function def main( in_: R.Tensor((1, 32, 32), dtype="float32"), - in_2: R.Tensor(("batch", 32, 32), dtype="float32"), - ) -> R.Tensor(("batch", 32, 32), dtype="float32"): - batch = T.int64() + in_2: R.Tensor((batch, 32, 32), dtype="float32"), + ) -> R.Tensor((batch, 32, 32), dtype="float32"): R.func_attr({"num_input": 2}) with R.dataflow(): gv: R.Tensor((batch, 32, 32), dtype="float32") = R.broadcast_to( @@ -7832,67 +7841,71 @@ def main( R.output(gv) return gv + A = T.dynamic("A") + @I.ir_module class ExpectedShapeSlice1: @R.function def main( - x: R.Tensor(("A", 10, 5), dtype="float32"), + x: R.Tensor((A, 10, 5), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Shape(ndim=2): - A = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Shape([A, 10]) = R.shape([A, 10]) R.output(gv) return gv + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class ExpectedShapeSlice2: @R.function def main( - x: R.Tensor(("A", "B", 5), dtype="float32"), + x: R.Tensor((A, B, 5), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Shape(ndim=2): - A = T.int64() - B = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Shape([A, B]) = R.shape([A, B]) R.output(gv) return gv + C = T.dynamic("C") + @I.ir_module class ExpectedShapeSlice3: @R.function def main( - x: R.Tensor((20, 10, "C"), dtype="float32"), + x: R.Tensor((20, 10, C), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Tensor((2,), dtype="int64"): - C = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((2,), dtype="int64") = R.const([20, 10], "int64") R.output(gv) return gv + A = T.dynamic("A") + B = T.dynamic("B") + C = T.dynamic("C") + @I.ir_module class ExpectedShapeSlice4: @R.function def main( - x: R.Tensor(("A", "B", "C"), dtype="float32"), + x: R.Tensor((A, B, C), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Shape(ndim=2): - A = T.int64() - B = T.int64() - C = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Shape([A, B]) = R.shape([A, B]) @@ -7914,67 +7927,71 @@ def main( R.output(gv) return gv + A = T.dynamic("A") + @I.ir_module class ExpectedShapeSlice6: @R.function def main( - x: R.Tensor(("A", 10, 5), dtype="float32"), + x: R.Tensor((A, 10, 5), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Tensor((1,), dtype="int64"): - A = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((1,), dtype="int64") = R.const([10], "int64") R.output(gv) return gv + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class ExpectedShapeSlice7: @R.function def main( - x: R.Tensor(("A", "B", 5), dtype="float32"), + x: R.Tensor((A, B, 5), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Shape(ndim=1): - A = T.int64() - B = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Shape([B]) = R.shape([B]) R.output(gv) return gv + C = T.dynamic("C") + @I.ir_module class ExpectedShapeSlice8: @R.function def main( - x: R.Tensor((20, 10, "C"), dtype="float32"), + x: R.Tensor((20, 10, C), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Tensor((1,), dtype="int64"): - C = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((1,), dtype="int64") = R.const([10], "int64") R.output(gv) return gv + A = T.dynamic("A") + B = T.dynamic("B") + C = T.dynamic("C") + @I.ir_module class ExpectedShapeSlice9: @R.function def main( - x: R.Tensor(("A", "B", "C"), dtype="float32"), + x: R.Tensor((A, B, C), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Shape(ndim=1): - A = T.int64() - B = T.int64() - C = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Shape([B]) = R.shape([B]) @@ -7996,67 +8013,71 @@ def main( R.output(gv) return gv + A = T.dynamic("A") + @I.ir_module class ExpectedShapeSlice11: @R.function def main( - x: R.Tensor(("A", 10, 5), dtype="float32"), + x: R.Tensor((A, 10, 5), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Tensor((2,), dtype="int64"): - A = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((2,), dtype="int64") = R.const([10, 5], "int64") R.output(gv) return gv + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class ExpectedShapeSlice12: @R.function def main( - x: R.Tensor(("A", "B", 5), dtype="float32"), + x: R.Tensor((A, B, 5), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Shape(ndim=2): - A = T.int64() - B = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Shape([B, 5]) = R.shape([B, 5]) R.output(gv) return gv + C = T.dynamic("C") + @I.ir_module class ExpectedShapeSlice13: @R.function def main( - x: R.Tensor((20, 10, "C"), dtype="float32"), + x: R.Tensor((20, 10, C), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Shape(ndim=2): - C = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Shape([10, C]) = R.shape([10, C]) R.output(gv) return gv + A = T.dynamic("A") + B = T.dynamic("B") + C = T.dynamic("C") + @I.ir_module class ExpectedShapeSlice14: @R.function def main( - x: R.Tensor(("A", "B", "C"), dtype="float32"), + x: R.Tensor((A, B, C), dtype="float32"), starts: R.Tensor((1,), dtype="int64"), ends: R.Tensor((1,), dtype="int64"), axes: R.Tensor((1,), dtype="int64"), ) -> R.Shape(ndim=2): - A = T.int64() - B = T.int64() - C = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Shape([B, C]) = R.shape([B, C]) @@ -8957,6 +8978,11 @@ def verify_tile(dynamic, in_shape, repeats, out_shape, expected): expected.update_func(expected.get_global_var("tile"), tvm_model["tile"]) tvm.ir.assert_structural_equal(tvm_model, expected) + tile_input_dim_0 = T.dynamic("tile_input_dim_0") + tile_input_dim_1 = T.dynamic("tile_input_dim_1") + tile_input_dim_2 = T.dynamic("tile_input_dim_2") + tile_input_dim_3 = T.dynamic("tile_input_dim_3") + @I.ir_module class ExpectedTileDynamicInput: @Ts.prim_func(private=True) @@ -8967,27 +8993,23 @@ def tile(input: T.handle, T_tile: T.handle): def main( input: R.Tensor( ( - "tile_input_dim_0", - "tile_input_dim_1", - "tile_input_dim_2", - "tile_input_dim_3", + tile_input_dim_0, + tile_input_dim_1, + tile_input_dim_2, + tile_input_dim_3, ), dtype="float32", ), repeats: R.Tensor((4,), dtype="int64"), ) -> R.Tensor( ( - "tile_input_dim_0 * 2", - "tile_input_dim_1", - "tile_input_dim_2 * 3", - "tile_input_dim_3 * 2", + tile_input_dim_0 * 2, + tile_input_dim_1, + tile_input_dim_2 * 3, + tile_input_dim_3 * 2, ), dtype="float32", ): - tile_input_dim_0 = T.int64() - tile_input_dim_1 = T.int64() - tile_input_dim_2 = T.int64() - tile_input_dim_3 = T.int64() R.func_attr({"num_input": 1}) cls = ExpectedTileDynamicInput with R.dataflow(): @@ -9083,6 +9105,8 @@ def make_expected(dynamic_input, in_shape): ) if rank == 2: + tile_dim_0 = T.dynamic("tile_dim_0") + tile_dim_1 = T.dynamic("tile_dim_1") @I.ir_module class ExpectedTileRank2: @@ -9095,8 +9119,6 @@ def main( input: R.Tensor(input_shape, dtype="float32"), repeats: R.Tensor((2,), dtype="int64"), ) -> R.Tensor(dtype="float32", ndim=2): - tile_dim_0 = T.int64() - tile_dim_1 = T.int64() R.func_attr({"num_input": 2}) cls = ExpectedTileRank2 with R.dataflow(): @@ -9118,6 +9140,9 @@ def main( return ExpectedTileRank2 if rank == 3: + tile_dim_0 = T.dynamic("tile_dim_0") + tile_dim_1 = T.dynamic("tile_dim_1") + tile_dim_2 = T.dynamic("tile_dim_2") @I.ir_module class ExpectedTileRank3: @@ -9130,9 +9155,6 @@ def main( input: R.Tensor(input_shape, dtype="float32"), repeats: R.Tensor((3,), dtype="int64"), ) -> R.Tensor(dtype="float32", ndim=3): - tile_dim_0 = T.int64() - tile_dim_1 = T.int64() - tile_dim_2 = T.int64() R.func_attr({"num_input": 2}) cls = ExpectedTileRank3 with R.dataflow(): @@ -9155,6 +9177,10 @@ def main( return ExpectedTileRank3 if rank == 4: + tile_dim_0 = T.dynamic("tile_dim_0") + tile_dim_1 = T.dynamic("tile_dim_1") + tile_dim_2 = T.dynamic("tile_dim_2") + tile_dim_3 = T.dynamic("tile_dim_3") @I.ir_module class ExpectedTileRank4: @@ -9167,10 +9193,6 @@ def main( input: R.Tensor(input_shape, dtype="float32"), repeats: R.Tensor((4,), dtype="int64"), ) -> R.Tensor(dtype="float32", ndim=4): - tile_dim_0 = T.int64() - tile_dim_1 = T.int64() - tile_dim_2 = T.int64() - tile_dim_3 = T.int64() R.func_attr({"num_input": 2}) cls = ExpectedTileRank4 with R.dataflow(): @@ -10615,14 +10637,15 @@ def verify_flatten_dynamic_ir(axis, expected): tvm_model = from_onnx(model, keep_params_in_input=True) tvm.ir.assert_structural_equal(tvm_model, expected) + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class ExpectedDynamicAxis0: @R.function - def main(x: R.Tensor((1, "A", "B", 32), dtype="float32")) -> R.Tensor( - (1, "A * B * 32"), dtype="float32" + def main(x: R.Tensor((1, A, B, 32), dtype="float32")) -> R.Tensor( + (1, A * B * 32), dtype="float32" ): - A = T.int64() - B = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((1, A * B * 32), dtype="float32") = R.reshape( @@ -10631,28 +10654,30 @@ def main(x: R.Tensor((1, "A", "B", 32), dtype="float32")) -> R.Tensor( R.output(gv) return gv + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class ExpectedDynamicAxisNegative1: @R.function - def main(x: R.Tensor((1, "A", "B", 32), dtype="float32")) -> R.Tensor( - ("A * B", 32), dtype="float32" + def main(x: R.Tensor((1, A, B, 32), dtype="float32")) -> R.Tensor( + (A * B, 32), dtype="float32" ): - A = T.int64() - B = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((A * B, 32), dtype="float32") = R.reshape(x, R.shape([A * B, 32])) R.output(gv) return gv + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class ExpectedDynamicAxis2: @R.function - def main(x: R.Tensor((1, "A", "B", 32), dtype="float32")) -> R.Tensor( - ("A", "B * 32"), dtype="float32" + def main(x: R.Tensor((1, A, B, 32), dtype="float32")) -> R.Tensor( + (A, B * 32), dtype="float32" ): - A = T.int64() - B = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((A, B * 32), dtype="float32") = R.reshape(x, R.shape([A, B * 32])) @@ -10758,11 +10783,12 @@ def verify_nonzero(shape, expected): tvm_model = from_onnx(model, keep_params_in_input=True) tvm.ir.assert_structural_equal(tvm_model, expected) + nonzero_numbers = T.dynamic("nonzero_numbers") + @I.ir_module class ExpectedScalar: @R.function def main(x: R.Tensor((), dtype="bool")): - nonzero_numbers = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): lv: R.Tensor((1, nonzero_numbers), dtype="int64") = R.match_cast( @@ -10772,11 +10798,12 @@ def main(x: R.Tensor((), dtype="bool")): R.output(gv) return gv + nonzero_numbers = T.dynamic("nonzero_numbers") + @I.ir_module class ExpectedRank1: @R.function def main(x: R.Tensor((1,), dtype="bool")): - nonzero_numbers = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): lv: R.Tensor((1, nonzero_numbers), dtype="int64") = R.match_cast( @@ -10786,11 +10813,12 @@ def main(x: R.Tensor((1,), dtype="bool")): R.output(gv) return gv + nonzero_numbers = T.dynamic("nonzero_numbers") + @I.ir_module class ExpectedRank2: @R.function def main(x: R.Tensor((2, 3), dtype="bool")): - nonzero_numbers = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): lv: R.Tensor((2, nonzero_numbers), dtype="int64") = R.match_cast( @@ -10800,11 +10828,12 @@ def main(x: R.Tensor((2, 3), dtype="bool")): R.output(gv) return gv + nonzero_numbers = T.dynamic("nonzero_numbers") + @I.ir_module class ExpectedRank3: @R.function def main(x: R.Tensor((4, 5, 6), dtype="bool")): - nonzero_numbers = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): lv: R.Tensor((3, nonzero_numbers), dtype="int64") = R.match_cast( @@ -10814,11 +10843,12 @@ def main(x: R.Tensor((4, 5, 6), dtype="bool")): R.output(gv) return gv + nonzero_numbers = T.dynamic("nonzero_numbers") + @I.ir_module class ExpectedRank4: @R.function def main(x: R.Tensor((7, 8, 9, 10), dtype="bool")): - nonzero_numbers = T.int64() R.func_attr({"num_input": 1}) with R.dataflow(): lv: R.Tensor((4, nonzero_numbers), dtype="int64") = R.match_cast( @@ -11649,16 +11679,17 @@ def verify_symbolic_shape_deduction(with_reshape_flatten, expected): tvm_model = from_onnx(model, keep_params_in_input=True) tvm.ir.assert_structural_equal(tvm_model["main"].without_attr("params"), expected["main"]) + batch = T.dynamic("batch") + seq = T.dynamic("seq") + @I.ir_module class ExpectedWithReshapeFlatten: @R.function def main( - data: R.Tensor(("batch", "seq"), dtype="float32"), + data: R.Tensor((batch, seq), dtype="float32"), axes: R.Tensor((1,), dtype="int64"), target_shape: R.Tensor((1,), dtype="int64"), - ) -> R.Tensor(("batch",), dtype="float32"): - batch = T.int64() - seq = T.int64() + ) -> R.Tensor((batch,), dtype="float32"): R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((batch,), dtype="float32") = R.broadcast_to( @@ -11667,15 +11698,16 @@ def main( R.output(gv) return gv + batch = T.dynamic("batch") + seq = T.dynamic("seq") + @I.ir_module class ExpectedWithoutReshapeFlatten: @R.function def main( - data: R.Tensor(("batch", "seq"), dtype="float32"), + data: R.Tensor((batch, seq), dtype="float32"), axes: R.Tensor((1,), dtype="int64"), - ) -> R.Tensor(("batch",), dtype="float32"): - batch = T.int64() - seq = T.int64() + ) -> R.Tensor((batch,), dtype="float32"): R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((batch,), dtype="float32") = R.broadcast_to( @@ -11703,14 +11735,15 @@ def test_multi_inputs_with_same_symbolic_shape(): model = helper.make_model(graph, producer_name="test_multi_symbolic_shape_input") tvm_model = from_onnx(model, keep_params_in_input=True) + batch = T.dynamic("batch") + @I.ir_module class Expected: @R.function def main( - data1: R.Tensor(("batch", 1), dtype="float32"), - data2: R.Tensor(("batch", 1), dtype="float32"), - ) -> R.Tensor(("batch", 2), dtype="float32"): - batch = T.int64() + data1: R.Tensor((batch, 1), dtype="float32"), + data2: R.Tensor((batch, 1), dtype="float32"), + ) -> R.Tensor((batch, 2), dtype="float32"): R.func_attr({"num_input": 2}) with R.dataflow(): gv: R.Tensor((batch, 2), dtype="float32") = R.concat((data1, data2), axis=1) @@ -11846,12 +11879,13 @@ def test_shape_dim_string_expression_graph_add(): tvm_model = from_onnx(model, opset=14, keep_params_in_input=True) # fmt: off + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(("A", "B", "A + B"), dtype="float32")) -> R.Tensor(("A", "B", "A + B"), dtype="float32"): - A = T.int64() - B = T.int64() + def main(x: R.Tensor((A, B, A + B), dtype="float32")) -> R.Tensor((A, B, A + B), dtype="float32"): R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((A, B, A + B), dtype="float32") = x @@ -11880,12 +11914,13 @@ def test_shape_dim_string_expression_graph_subtract(): tvm_model = from_onnx(model, opset=14, keep_params_in_input=True) # fmt: off + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(("A", "B", "A - B"), dtype="float32")) -> R.Tensor(("A", "B", "A - B"), dtype="float32"): - A = T.int64() - B = T.int64() + def main(x: R.Tensor((A, B, A - B), dtype="float32")) -> R.Tensor((A, B, A - B), dtype="float32"): R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((A, B, A - B), dtype="float32") = x @@ -11914,12 +11949,13 @@ def test_shape_dim_string_expression_graph_mul(): tvm_model = from_onnx(model, opset=14, keep_params_in_input=True) # fmt: off + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(("A", "B", "A * B"), dtype="float32")) -> R.Tensor(("A", "B", "A * B"), dtype="float32"): - A = T.int64() - B = T.int64() + def main(x: R.Tensor((A, B, A * B), dtype="float32")) -> R.Tensor((A, B, A * B), dtype="float32"): R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((A, B, A * B), dtype="float32") = x @@ -11949,12 +11985,13 @@ def test_shape_dim_string_expression_graph_div_1(): tvm_model = from_onnx(model, opset=14, keep_params_in_input=True) # fmt: off + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(("A", "B", "A // B"), dtype="float32")) -> R.Tensor(("A", "B", "A // B"), dtype="float32"): - A = T.int64() - B = T.int64() + def main(x: R.Tensor((A, B, A // B), dtype="float32")) -> R.Tensor((A, B, A // B), dtype="float32"): R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((A, B, A // B), dtype="float32") = x @@ -11984,12 +12021,13 @@ def test_shape_dim_string_expression_graph_div_2(): tvm_model = from_onnx(model, opset=14, keep_params_in_input=True) # fmt: off + A = T.dynamic("A") + B = T.dynamic("B") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(("A", "B", "A // B"), dtype="float32")) -> R.Tensor(("A", "B", "A // B"), dtype="float32"): - A = T.int64() - B = T.int64() + def main(x: R.Tensor((A, B, A // B), dtype="float32")) -> R.Tensor((A, B, A // B), dtype="float32"): R.func_attr({"num_input": 1}) with R.dataflow(): gv: R.Tensor((A, B, A // B), dtype="float32") = x diff --git a/tests/python/relax/test_frontend_stablehlo.py b/tests/python/relax/test_frontend_stablehlo.py index f1b426d935c2..d6aa363b9969 100644 --- a/tests/python/relax/test_frontend_stablehlo.py +++ b/tests/python/relax/test_frontend_stablehlo.py @@ -19,7 +19,7 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -# ruff: noqa: E501, F841 +# ruff: noqa: E501 # pylint: disable=c-extension-no-member @@ -183,17 +183,18 @@ def test_add_dynamic(): mod = from_stablehlo(add_dyn) + n_0 = T.dynamic("n_0") + n_1 = T.dynamic("n_1") + n_2 = T.dynamic("n_2") + n_3 = T.dynamic("n_3") + @I.ir_module class Expected: @R.function def main( - arg0: R.Tensor(("n_0", "n_1"), dtype="float32"), - arg1: R.Tensor(("n_2", "n_3"), dtype="float32"), + arg0: R.Tensor((n_0, n_1), dtype="float32"), + arg1: R.Tensor((n_2, n_3), dtype="float32"), ) -> R.Tensor(dtype="float32", ndim=2): - n_0 = T.int64() - n_1 = T.int64() - n_2 = T.int64() - n_3 = T.int64() with R.dataflow(): lv: R.Tensor(dtype="float32", ndim=2) = R.add(arg0, arg1) gv: R.Tensor(dtype="float32", ndim=2) = lv diff --git a/tests/python/relax/test_frontend_tflite.py b/tests/python/relax/test_frontend_tflite.py index ee36bbced395..8f7b44a78839 100644 --- a/tests/python/relax/test_frontend_tflite.py +++ b/tests/python/relax/test_frontend_tflite.py @@ -1088,6 +1088,9 @@ class TfFillDynamic(tf.Module): def func(self, dims, value): return tf.fill(dims, value) + fill_dim_0 = T.dynamic("fill_dim_0") + fill_dim_1 = T.dynamic("fill_dim_1") + @I.ir_module class Expected: @R.function @@ -1095,8 +1098,6 @@ def main( dims: R.Tensor((2,), dtype="int32"), value: R.Tensor((), dtype="float32") ) -> R.Tensor(dtype="float32", ndim=2): R.func_attr({"num_input": 2}) - fill_dim_0 = T.int64() - fill_dim_1 = T.int64() with R.dataflow(): lv: R.Tensor((2,), dtype="int32") = R.match_cast( dims, R.Tensor((2,), dtype="int32") @@ -1123,13 +1124,14 @@ class TfRandomUniform(tf.Module): def func(self, shape): return tf.raw_ops.RandomUniform(shape=shape, dtype=tf.float32, seed=7, seed2=11) + random_uniform_dim_0 = T.dynamic("random_uniform_dim_0") + random_uniform_dim_1 = T.dynamic("random_uniform_dim_1") + @I.ir_module class Expected: @R.function def main(shape: R.Tensor((2,), dtype="int32")) -> R.Tensor(dtype="float32", ndim=2): R.func_attr({"num_input": 1}) - random_uniform_dim_0 = T.int64() - random_uniform_dim_1 = T.int64() with R.dataflow(): lv: R.Tensor((2,), dtype="int32") = R.match_cast( shape, R.Tensor((2,), dtype="int32") @@ -1167,13 +1169,14 @@ class TfRandomStandardNormal(tf.Module): def func(self, shape): return tf.raw_ops.RandomStandardNormal(shape=shape, dtype=tf.float32, seed=3, seed2=5) + random_standard_normal_dim_0 = T.dynamic("random_standard_normal_dim_0") + random_standard_normal_dim_1 = T.dynamic("random_standard_normal_dim_1") + @I.ir_module class Expected: @R.function def main(shape: R.Tensor((2,), dtype="int32")) -> R.Tensor(dtype="float32", ndim=2): R.func_attr({"num_input": 1}) - random_standard_normal_dim_0 = T.int64() - random_standard_normal_dim_1 = T.int64() with R.dataflow(): lv: R.Tensor((2,), dtype="int32") = R.match_cast( shape, R.Tensor((2,), dtype="int32") @@ -1227,6 +1230,8 @@ def func(self, logits, num_samples): seed2=17, ) + multinomial_num_samples = T.dynamic("multinomial_num_samples") + @I.ir_module class Expected: @R.function @@ -1235,7 +1240,6 @@ def main( num_samples: R.Tensor((), dtype="int32"), ) -> R.Tensor(dtype="int64", ndim=2): R.func_attr({"num_input": 2}) - multinomial_num_samples = T.int64() with R.dataflow(): lv: R.Tensor((), dtype="int32") = R.match_cast( num_samples, R.Tensor((), dtype="int32") @@ -13306,6 +13310,9 @@ def test_dilate_dynamic_dilations(): mod = from_tflite(tflite_model) mod["main"] = mod["main"].without_attr("params") + dilate_stride_0 = T.dynamic("dilate_stride_0") + dilate_stride_1 = T.dynamic("dilate_stride_1") + @I.ir_module class Expected: @R.function @@ -13314,8 +13321,6 @@ def main( tvmgen_tensor_1: R.Tensor((2,), dtype="int32"), ) -> R.Tensor(dtype="float32", ndim=2): R.func_attr({"num_input": 2}) - dilate_stride_0 = T.int64() - dilate_stride_1 = T.int64() with R.dataflow(): lv: R.Tensor((2,), dtype="int32") = R.match_cast( tvmgen_tensor_1, R.Tensor((2,), dtype="int32") diff --git a/tests/python/relax/test_inline_functions.py b/tests/python/relax/test_inline_functions.py index b50bfca60994..02354e54fef2 100644 --- a/tests/python/relax/test_inline_functions.py +++ b/tests/python/relax/test_inline_functions.py @@ -14,7 +14,6 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -# ruff: noqa: F401 import pytest @@ -175,22 +174,29 @@ def test_subroutine_with_symbolic_vars(): caller's `tirx::Var` symbolic variables should remain. """ + n_main = T.dynamic("n") + n_subroutine = T.dynamic("n") + @I.ir_module class Before: @R.function(private=True) - def main(A: R.Tensor(["n", 16], "int32")) -> R.Tensor(["n", 32], "int32"): + def main(A: R.Tensor([n_main, 16], "int32")) -> R.Tensor([n_main, 32], "int32"): B = A * A C = Before.subroutine(B) D = C + C return D @R.function(private=True) - def subroutine(B: R.Tensor(["n", 16], "int32")) -> R.Tensor(["n", 32], "int32"): + def subroutine(B: R.Tensor([n_subroutine, 16], "int32")) -> R.Tensor( + [n_subroutine, 32], "int32" + ): C = R.concat([B, B], axis=1) return C + n = T.dynamic("n") + @R.function(private=True) - def expected(A: R.Tensor(["n", 16], "int32")) -> R.Tensor(["n", 32], "int32"): + def expected(A: R.Tensor([n, 16], "int32")) -> R.Tensor([n, 32], "int32"): B = A * A C = R.concat([B, B], axis=1) D = C + C @@ -208,6 +214,8 @@ def test_subroutine_with_symbolic_vars_and_static_argument(): should remain. """ + n = T.dynamic("n") + @I.ir_module class Before: @R.function(private=True) @@ -218,7 +226,7 @@ def main(A: R.Tensor([16, 16], "int32")) -> R.Tensor([16, 32], "int32"): return D @R.function(private=True) - def subroutine(B: R.Tensor(["n", 16], "int32")) -> R.Tensor(["n", 32], "int32"): + def subroutine(B: R.Tensor([n, 16], "int32")) -> R.Tensor([n, 32], "int32"): C = R.concat([B, B], axis=1) return C @@ -274,6 +282,9 @@ def test_inline_multiple_instances_with_distinct_static_shapes(): different value for the symbolic variables it uses. """ + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Before: @R.function(private=True) @@ -283,7 +294,7 @@ def main(A: R.Tensor([16, 16]), B: R.Tensor([32, 32])): return (A_out, B_out) @R.function(private=True) - def subroutine(Input: R.Tensor(["n", "m"])) -> R.Tensor(["n", "m"]): + def subroutine(Input: R.Tensor([n, m])) -> R.Tensor([n, m]): Output = Input + Input return Output diff --git a/tests/python/relax/test_op_image.py b/tests/python/relax/test_op_image.py index 2644e0c3de9c..8361b9c788ad 100644 --- a/tests/python/relax/test_op_image.py +++ b/tests/python/relax/test_op_image.py @@ -465,7 +465,7 @@ def test_affine_grid_e2e(batch, target_h, target_w): @tvm.script.ir_module class AffineGridModule: @R.function - def main(theta: R.Tensor(("batch", 2, 3), "float32")) -> R.Tensor("float32", ndim=4): + def main(theta: R.Tensor((batch, 2, 3), "float32")) -> R.Tensor("float32", ndim=4): gv = R.image.affine_grid(theta, size=(target_h, target_w)) return gv diff --git a/tests/python/relax/test_op_index.py b/tests/python/relax/test_op_index.py index 3eeffa9cb6eb..83af134398a2 100644 --- a/tests/python/relax/test_op_index.py +++ b/tests/python/relax/test_op_index.py @@ -188,8 +188,8 @@ def test_take_infer_ty_shape_symbolic(): bb = relax.BlockBuilder() m = tirx.Var("m", "int64") n = tirx.Var("n", "int64") - i = tirx.Var("i", "int64") - j = tirx.Var("j", "int64") + i = T.dynamic("i", "int64") + j = T.dynamic("j", "int64") k = tirx.Var("k", "int64") x0 = relax.Var("x", R.Tensor((m, n), "float32")) x1 = relax.Var("x", R.Tensor((m, n))) @@ -781,8 +781,8 @@ def test_dynamic_strided_slice_infer_ty(): def test_dynamic_strided_slice_infer_ty_symbolic(): bb = relax.BlockBuilder() - i = tirx.Var("i", "int64") - j = tirx.Var("j", "int64") + i = T.dynamic("i", "int64") + j = T.dynamic("j", "int64") k = tirx.Var("k", "int64") l = tirx.Var("l", "int64") x0 = relax.Var("x", R.Tensor((i, j, k, l), "float32")) @@ -899,18 +899,20 @@ def test_dynamic_strided_slice_infer_ty_arg_wrong_shape_info(): def test_legalize_dynamic_begin_end(): """relax.op.strided_slice FLegalize must support dynamic begin/end""" + index = T.dynamic("index") + @I.ir_module class before: @R.function - def main(A: R.Tensor((16, 16), "float32"), B: R.Shape(["index"])) -> R.Tensor((1, 16)): - index = T.int64() + def main(A: R.Tensor((16, 16), "float32"), B: R.Shape([index])) -> R.Tensor((1, 16)): return R.strided_slice(A, [0], [index], [index + 1], assume_inbound=True) + index = T.dynamic("index") + @I.ir_module class expected: @R.function - def main(A: R.Tensor((16, 16), "float32"), B: R.Shape(["index"])) -> R.Tensor((1, 16)): - index = T.int64() + def main(A: R.Tensor((16, 16), "float32"), B: R.Shape([index])) -> R.Tensor((1, 16)): return R.call_tir( expected.strided_slice, (A, index), @@ -936,16 +938,19 @@ def strided_slice( def test_legalize_dynamic_begin_inf_end(): """relax.op.strided_slice FLegalize must support dynamic begin/end""" + index = T.dynamic("index") + @I.ir_module class before: @R.function - def main(A: R.Tensor((16, 16), "float32"), B: R.Shape(["index"])) -> R.Tensor((1, 16)): - index = T.int64() + def main(A: R.Tensor((16, 16), "float32"), B: R.Shape([index])) -> R.Tensor((1, 16)): return R.strided_slice( A, [0], [index], [T.int64(np.iinfo(np.int64).max)], assume_inbound=False ) # fmt: off + index = T.dynamic("index") + @I.ir_module class expected: @Ts.prim_func(private=True) @@ -961,8 +966,7 @@ def strided_slice(A: T.Buffer((T.int64(16), T.int64(16)), "float32"), index: T.i T_dynamic_strided_slice_with_axes[v_ax0, v_ax1] = A[T.min(T.max(T.if_then_else(index < T.int64(0), index + T.int64(16), index), T.int64(0)), T.int64(16)) + v_ax0, v_ax1] @R.function - def main(A: R.Tensor((16, 16), dtype="float32"), B: R.Shape(["index"])) -> R.Tensor(("T.max(16 - T.max(T.if_then_else(index < 0, index + 16, index), 0), 0)", 16), dtype="float32"): - index = T.int64() + def main(A: R.Tensor((16, 16), dtype="float32"), B: R.Shape([index])) -> R.Tensor((T.max(16 - T.max(T.if_then_else(index < 0, index + 16, index), 0), 0), 16), dtype="float32"): cls = expected gv = R.call_tir(cls.strided_slice, (A, index), out_ty=R.Tensor((T.max(16 - T.max(T.if_then_else(index < 0, index + 16, index), 0), 0), 16), dtype="float32")) return gv diff --git a/tests/python/relax/test_op_size.py b/tests/python/relax/test_op_size.py index 77c5ebef5af1..aef849abe3b9 100644 --- a/tests/python/relax/test_op_size.py +++ b/tests/python/relax/test_op_size.py @@ -20,6 +20,7 @@ import tvm import tvm.testing from tvm import relax +from tvm.script import ir as I from tvm.script import relax as R @@ -42,10 +43,13 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((), "int64"): def test_op_size_dynamic(): + m = I.dynamic("m") + n = I.dynamic("n") + @tvm.script.ir_module class Module: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor((), "int64"): + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((), "int64"): return R.size(x) x_np = np.random.rand(4, 5).astype("float32") diff --git a/tests/python/relax/test_op_take.py b/tests/python/relax/test_op_take.py index eab0d14836fe..0fd7efd57bec 100644 --- a/tests/python/relax/test_op_take.py +++ b/tests/python/relax/test_op_take.py @@ -148,11 +148,12 @@ def test_take_dynamic_prim_value_as_index(axis): target = "llvm" dev = tvm.cpu() + n = T.dynamic("n") + @I.ir_module class Module: @R.function - def main(A: R.Tensor(["n", "n"], "float16")): - n = T.int64() + def main(A: R.Tensor([n, n], "float16")): output = R.take(A, R.prim_value(n - 1), axis=axis) return output diff --git a/tests/python/relax/test_op_view.py b/tests/python/relax/test_op_view.py index 8824d3f89106..4e6d34d1e2a4 100644 --- a/tests/python/relax/test_op_view.py +++ b/tests/python/relax/test_op_view.py @@ -123,15 +123,17 @@ def func(A: R.Tensor([16])): def test_infer_shape_of_1d_dynamic_view(): + N = T.dynamic("N") + @R.function(private=True) - def explicit_ty(A: R.Tensor(["N"])) -> R.Tensor(["N // 2"]): - N = T.int64() + def explicit_ty(A: R.Tensor([N])) -> R.Tensor([N // 2]): B: R.Tensor([N // 2]) = R.memory.view(A, R.shape([N // 2])) return B + N = T.dynamic("N") + @R.function(private=True) - def inferred_ty(A: R.Tensor(["N"])): - N = T.int64() + def inferred_ty(A: R.Tensor([N])): B = R.memory.view(A, R.shape([N // 2])) return B @@ -139,15 +141,17 @@ def inferred_ty(A: R.Tensor(["N"])): def test_infer_shape_of_2d_dynamic_view_of_1d_source(): + N = T.dynamic("N") + @R.function(private=True) - def explicit_ty(A: R.Tensor(["N"])) -> R.Tensor(["N // 8", 8]): - N = T.int64() + def explicit_ty(A: R.Tensor([N])) -> R.Tensor([N // 8, 8]): B: R.Tensor([N // 8, 8]) = R.memory.view(A, R.shape([N // 8, 8])) return B + N = T.dynamic("N") + @R.function(private=True) - def inferred_ty(A: R.Tensor(["N"])): - N = T.int64() + def inferred_ty(A: R.Tensor([N])): B = R.memory.view(A, R.shape([N // 8, 8])) return B @@ -155,15 +159,17 @@ def inferred_ty(A: R.Tensor(["N"])): def test_infer_shape_of_2d_dynamic_view(): + N = T.dynamic("N") + @R.function(private=True) - def explicit_ty(A: R.Tensor(["N"])) -> R.Tensor(["N // 2"]): - N = T.int64() + def explicit_ty(A: R.Tensor([N])) -> R.Tensor([N // 2]): B: R.Tensor([N // 2]) = R.memory.view(A, R.shape([N // 2])) return B + N = T.dynamic("N") + @R.function(private=True) - def inferred_ty(A: R.Tensor(["N"])): - N = T.int64() + def inferred_ty(A: R.Tensor([N])): B = R.memory.view(A, R.shape([N // 2])) return B @@ -172,30 +178,30 @@ def inferred_ty(A: R.Tensor(["N"])): def test_error_if_1d_dynamic_view_larger_than_1d_source(): with pytest.raises(ValueError): + N = T.dynamic("N") @R.function - def func(A: R.Tensor(["N"])): - N = T.int64() + def func(A: R.Tensor([N])): B = R.memory.view(A, R.shape([N + 1])) return B def test_error_if_1d_dynamic_view_provably_larger_than_1d_source(): with pytest.raises(ValueError): + N = T.dynamic("N") @R.function - def func(A: R.Tensor(["N"])): - N = T.int64() + def func(A: R.Tensor([N])): B = R.memory.view(A, R.shape([N + T.if_then_else(N < 0, -1, 1)])) return B def test_error_if_2d_dynamic_view_provably_larger_than_1d_source(): with pytest.raises(ValueError): + N = T.dynamic("N") @R.function - def func(A: R.Tensor(["N"])): - N = T.int64() + def func(A: R.Tensor([N])): B = R.memory.view(A, R.shape([N // 4 + 1, 4])) return B @@ -215,9 +221,10 @@ def test_validity_of_dynamic_view_may_depend_on_runtime_value(): """ + N = T.dynamic("N") + @R.function - def func(A: R.Tensor(["N"])): - N = T.int64() + def func(A: R.Tensor([N])): B = R.memory.view(A, R.shape([(N + 3) // 4, 4])) return B diff --git a/tests/python/relax/test_optimize_layout_transform.py b/tests/python/relax/test_optimize_layout_transform.py index 8c4bfa705677..0fa2d80ee74f 100644 --- a/tests/python/relax/test_optimize_layout_transform.py +++ b/tests/python/relax/test_optimize_layout_transform.py @@ -270,6 +270,9 @@ def main( def test_tranform_layout_tir_remove_pad_transform_layout(): + p0 = T.dynamic("p0") + i0 = T.dynamic("i0") + @I.ir_module class Before: @Ts.prim_func(private=True) @@ -288,9 +291,7 @@ def relax_relu_replacement( @Ts.prim_func(private=True) def remove_pad(var_input: T.handle, var_output: T.handle): T.func_attr({"operator_name": "remove_pad", "tirx.noalias": True}) - p0 = T.int64() input = T.match_buffer(var_input, (p0,)) - i0 = T.int64() output = T.match_buffer(var_output, (i0,)) # with Ts.sblock("root"): for ax0 in range(i0): @@ -343,6 +344,9 @@ def main(x: R.Tensor((14,), dtype="float32")) -> R.Tensor((14,), dtype="float32" R.output(gv) return gv + p0 = T.dynamic("p0") + i0 = T.dynamic("i0") + @I.ir_module class Expected: @Ts.prim_func(private=True) @@ -361,9 +365,7 @@ def relax_relu_replacement( @Ts.prim_func(private=True) def remove_pad(var_input: T.handle, var_output: T.handle): T.func_attr({"operator_name": "remove_pad", "tirx.noalias": True}) - p0 = T.int64() input = T.match_buffer(var_input, (p0,)) - i0 = T.int64() output = T.match_buffer(var_output, (i0,)) # with Ts.sblock("root"): for ax0 in range(i0): diff --git a/tests/python/relax/test_pipeline.py b/tests/python/relax/test_pipeline.py index dd8a15f8194f..f35e53be7230 100644 --- a/tests/python/relax/test_pipeline.py +++ b/tests/python/relax/test_pipeline.py @@ -56,12 +56,14 @@ def test_pipeline_with_kv_cache(): target = tvm.target.Target("llvm", host="llvm") pipeline = relax.pipeline.get_default_pipeline(target) + m = T.dynamic("m") + L = T.dynamic("L") + @tvm.script.ir_module class Mod: @R.function - def create_kv_cache(reserve_slots: R.Shape(["m"])): + def create_kv_cache(reserve_slots: R.Shape([m])): # just allocate minimum slot since it is only used to signal dtype - m = T.int64() init_data = R.ones((1, 4), "float32") kv_cache = R.call_pure_packed( "vm.builtin.attention_kv_cache_create", @@ -76,10 +78,9 @@ def create_kv_cache(reserve_slots: R.Shape(["m"])): def main( x: R.Tensor((1, 4), "float32"), y: R.Tensor((1, 4), "float32"), - shape: R.Shape(["L", 4]), + shape: R.Shape([L, 4]), kv_cache: R.Any, ): - L = T.int64() # computation of the current value curr_value = R.add(x, y) # update cache diff --git a/tests/python/relax/test_pytorch_integration.py b/tests/python/relax/test_pytorch_integration.py index fe95a40cf0ac..48434f7f8c81 100644 --- a/tests/python/relax/test_pytorch_integration.py +++ b/tests/python/relax/test_pytorch_integration.py @@ -36,6 +36,8 @@ from tvm.script import s_tir as Ts from tvm.script import tirx as T +n_matmul = T.dynamic("n", "int32") + @R.py_module class PyTorchIntegrationModule(BasePyModule): @@ -67,12 +69,11 @@ def matmul( var_C: T.handle, ): """TIR function for matrix multiplication.""" - n = T.int32() - A = T.match_buffer(var_A, (n, 16), "float32") + A = T.match_buffer(var_A, (n_matmul, 16), "float32") B = T.match_buffer(var_B, (16, 20), "float32") - C = T.match_buffer(var_C, (n, 20), "float32") + C = T.match_buffer(var_C, (n_matmul, 20), "float32") - for i, j, k in T.grid(n, 20, 16): + for i, j, k in T.grid(n_matmul, 20, 16): with Ts.sblock("block"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) with Ts.init(): diff --git a/tests/python/relax/test_relax_operators.py b/tests/python/relax/test_relax_operators.py index 865acd454602..9541e7b06e75 100644 --- a/tests/python/relax/test_relax_operators.py +++ b/tests/python/relax/test_relax_operators.py @@ -36,10 +36,14 @@ exec_mode = tvm.testing.parameter("bytecode", "compiled") +m = T.dynamic("m") +n = T.dynamic("n") + + @tvm.script.ir_module class InputModule: @R.function - def foo(x: R.Tensor(("m", "n"), "int64")): + def foo(x: R.Tensor((m, n), "int64")): y = R.unique(x, sorted=False) y_sorted = R.unique(x) return y, y_sorted @@ -190,9 +194,10 @@ def func(condition: R.Tensor((), "bool"), x: R.Tensor((), "int32")): def test_assert_on_symbolic_var_passes(exec_mode): + N = T.dynamic("N") + @R.function(pure=False) - def func(x: R.Tensor(["N"], "int32")): - N = T.int64() + def func(x: R.Tensor([N], "int32")): _ = R.assert_op(R.prim_value(N % 8 == 0)) return x @@ -201,9 +206,10 @@ def func(x: R.Tensor(["N"], "int32")): def test_assert_on_symbolic_var_fails(exec_mode): + N = T.dynamic("N") + @R.function(pure=False) - def func(x: R.Tensor(["N"], "int32")): - N = T.int64() + def func(x: R.Tensor([N], "int32")): _ = R.assert_op(R.prim_value(N % 8 == 0)) return x @@ -268,6 +274,10 @@ def test_op_shape_of(exec_mode): assert constrained_shape == tvm_ffi.Shape([1]) +m = T.dynamic("m") +n = T.dynamic("n") + + @tvm.script.ir_module class ShapeToTensorTest: @R.function @@ -275,9 +285,7 @@ def const_shape(shape: R.Shape(ndim=-1)) -> R.Tensor(ndim=-1): return R.shape_to_tensor(shape) @R.function - def symbolic_shape(shape: R.Shape(("m", "n"))) -> R.Tensor(ndim=-1): - m = T.int64() - n = T.int64() + def symbolic_shape(shape: R.Shape((m, n))) -> R.Tensor(ndim=-1): return R.shape_to_tensor(shape) @@ -586,9 +594,10 @@ def func(condition: T.bool): def test_computed_prim_value_as_branch_condition(exec_mode): """The primitive scalar condition may be computed within the function""" + N = T.dynamic("N") + @R.function - def func(x: R.Tensor(["N"], "int64")): - N = T.int64() + def func(x: R.Tensor([N], "int64")): if R.prim_value(N % 16 == 0): out = R.prim_value(5) else: diff --git a/tests/python/relax/test_relax_to_pyfunc_converter.py b/tests/python/relax/test_relax_to_pyfunc_converter.py index 3eb189cb658a..21b381ae0782 100644 --- a/tests/python/relax/test_relax_to_pyfunc_converter.py +++ b/tests/python/relax/test_relax_to_pyfunc_converter.py @@ -33,6 +33,14 @@ from tvm.script import s_tir as Ts from tvm.script import tirx as T +n_symbolic_add = T.dynamic("n") +batch_symbolic_matmul = T.dynamic("batch") +m = T.dynamic("m") +k = T.dynamic("k") +n_symbolic_matmul = T.dynamic("n") +batch_symbolic_expand_dims = T.dynamic("batch") +seq_len = T.dynamic("seq_len") + @I.ir_module class ComprehensiveTestModule: @@ -91,21 +99,22 @@ def complex_function(x: R.Tensor((5,), "float32"), y: R.Tensor((5,), "float32")) return R.nn.relu(tir_result) @R.function - def symbolic_add(x: R.Tensor(("n",), "float32"), y: R.Tensor(("n",), "float32")) -> R.Tensor( - ("n",), "float32" - ): + def symbolic_add( + x: R.Tensor((n_symbolic_add,), "float32"), y: R.Tensor((n_symbolic_add,), "float32") + ) -> R.Tensor((n_symbolic_add,), "float32"): return R.add(x, y) @R.function def symbolic_matmul( - x: R.Tensor(("batch", "m", "k"), "float32"), y: R.Tensor(("batch", "k", "n"), "float32") - ) -> R.Tensor(("batch", "m", "n"), "float32"): + x: R.Tensor((batch_symbolic_matmul, m, k), "float32"), + y: R.Tensor((batch_symbolic_matmul, k, n_symbolic_matmul), "float32"), + ) -> R.Tensor((batch_symbolic_matmul, m, n_symbolic_matmul), "float32"): return R.matmul(x, y) @R.function - def symbolic_expand_dims(x: R.Tensor(("batch", "seq_len"), "float32")) -> R.Tensor( - ("batch", "seq_len", 1), "float32" - ): + def symbolic_expand_dims( + x: R.Tensor((batch_symbolic_expand_dims, seq_len), "float32"), + ) -> R.Tensor((batch_symbolic_expand_dims, seq_len, 1), "float32"): return R.expand_dims(x, axis=2) @R.function diff --git a/tests/python/relax/test_runtime_builtin_rnn_state.py b/tests/python/relax/test_runtime_builtin_rnn_state.py index 410481ae704a..85a07e3e03fe 100644 --- a/tests/python/relax/test_runtime_builtin_rnn_state.py +++ b/tests/python/relax/test_runtime_builtin_rnn_state.py @@ -212,6 +212,8 @@ def rnn_state_get( dtype: str, ): # fmt: off + batch_size = T.dynamic("batch_size", "int32") + @Ts.prim_func def _rnn_state_get( var_storage: T.handle, @@ -219,7 +221,6 @@ def _rnn_state_get( var_history_slot_ids: T.handle, var_output: T.handle, ): - batch_size = T.int32() storage = T.match_buffer(var_storage, (reserved_nseq, max_history, *shape), dtype) seq_slot_ids = T.match_buffer(var_seq_slot_ids, (batch_size,), "int32") @@ -247,6 +248,8 @@ def rnn_state_set( dtype: str, ): # fmt: off + batch_size = T.dynamic("batch_size", "int32") + @Ts.prim_func def _rnn_state_set( var_storage: T.handle, @@ -254,7 +257,6 @@ def _rnn_state_set( var_history_slot_ids: T.handle, var_data: T.handle, ): - batch_size = T.int32() storage = T.match_buffer(var_storage, (reserved_nseq, max_history, *shape), dtype) seq_slot_ids = T.match_buffer(var_seq_slot_ids, (batch_size,), "int32") diff --git a/tests/python/relax/test_testing_nn.py b/tests/python/relax/test_testing_nn.py index f9a508b8863c..3784325c84f3 100644 --- a/tests/python/relax/test_testing_nn.py +++ b/tests/python/relax/test_testing_nn.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: F401, RUF005 +# ruff: noqa: RUF005 import tvm import tvm.testing from tvm import relax @@ -81,28 +81,32 @@ def forward(self, input: relax.Expr) -> relax.Var: state = relax.op.matmul(input, self.weights) return self.activation(state) + batch_size_main = T.dynamic("batch_size") + batch_size_layer = T.dynamic("batch_size") + batch_size_activation = T.dynamic("batch_size") + @I.ir_module class Expected: @R.function def main( - state: R.Tensor(("batch_size", 64), dtype="float32"), + state: R.Tensor((batch_size_main, 64), dtype="float32"), weights: R.Tensor((64, 32), dtype="float32"), - ) -> R.Tensor(("batch_size", 32), dtype="float32"): + ) -> R.Tensor((batch_size_main, 32), dtype="float32"): state = Expected.layer(state, weights) return state @R.function(private=True) def layer( - state: R.Tensor(("batch_size", 64), dtype="float32"), + state: R.Tensor((batch_size_layer, 64), dtype="float32"), weights: R.Tensor((64, 32), dtype="float32"), - ) -> R.Tensor(("batch_size", 32), dtype="float32"): + ) -> R.Tensor((batch_size_layer, 32), dtype="float32"): state = R.matmul(state, weights) state = Expected.activation(state) return state @R.function(private=True) - def activation(state: R.Tensor(("batch_size", 32), dtype="float32")) -> R.Tensor( - ("batch_size", 32), dtype="float32" + def activation(state: R.Tensor((batch_size_activation, 32), dtype="float32")) -> R.Tensor( + (batch_size_activation, 32), dtype="float32" ): state = R.nn.relu(state) return state diff --git a/tests/python/relax/test_tir_call_source_kernel.py b/tests/python/relax/test_tir_call_source_kernel.py index 397d584facce..1d269ccfe42e 100644 --- a/tests/python/relax/test_tir_call_source_kernel.py +++ b/tests/python/relax/test_tir_call_source_kernel.py @@ -42,41 +42,43 @@ def test_tir_call_source_kernel(): BLOCK_SIZE = 64 + m_add = T.dynamic("m") + m_main = T.dynamic("m") + @I.ir_module class Module: @Ts.prim_func def add(x_handle: T.handle, y_handle: T.handle, output_handle: T.handle) -> None: T.func_attr({"global_symbol": "add"}) - m = T.int64() - x = T.match_buffer(x_handle, (m,), "float32") - y = T.match_buffer(y_handle, (m,), "float32") - output = T.match_buffer(output_handle, (m,), "float32") + x = T.match_buffer(x_handle, (m_add,), "float32") + y = T.match_buffer(y_handle, (m_add,), "float32") + output = T.match_buffer(output_handle, (m_add,), "float32") with Ts.sblock("root"): - Ts.reads(x[0:m], y[0:m]) - Ts.writes(output[0:m]) + Ts.reads(x[0:m_add], y[0:m_add]) + Ts.writes(output[0:m_add]) T.call_kernel( add_cuda_source, - ((T.ceildiv(m, BLOCK_SIZE),), (BLOCK_SIZE,)), + ((T.ceildiv(m_add, BLOCK_SIZE),), (BLOCK_SIZE,)), x.data, y.data, output.data, - m, + m_add, kernel_name="add_kernel", ) @R.function - def main(x: R.Tensor(("m",), "float32"), y: R.Tensor(("m",), "float32")): - m = T.int64() + def main(x: R.Tensor((m_main,), "float32"), y: R.Tensor((m_main,), "float32")): with R.dataflow(): - output = R.call_tir(Module.add, [x, y], relax.TensorType((m,), "float32")) + output = R.call_tir(Module.add, [x, y], relax.TensorType((m_main,), "float32")) R.output(output) return output + m = T.dynamic("m") + @I.ir_module class Parsed: @Ts.prim_func def add(x_handle: T.handle, y_handle: T.handle, output_handle: T.handle): - m = T.int64() x = T.match_buffer(x_handle, (m,)) y = T.match_buffer(y_handle, (m,)) output = T.match_buffer(output_handle, (m,)) diff --git a/tests/python/relax/test_transform.py b/tests/python/relax/test_transform.py index bb5816342295..daf50d62ba93 100644 --- a/tests/python/relax/test_transform.py +++ b/tests/python/relax/test_transform.py @@ -77,11 +77,13 @@ def rewrite_match_cast_symbol(block, _mod, _ctx): def test_to_non_dataflow(): + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class TestToNonDataflow: @R.function - def foo(x: R.Tensor(("m", "n"), "float32")): - m, n = T.int64(), T.int64() + def foo(x: R.Tensor((m, n), "float32")): with R.dataflow(): lv0 = R.call_dps_packed( "test.op.identity", @@ -136,22 +138,26 @@ def fvisit(e): def test_call_tir_rewrite(): + m_exp = T.dynamic("m") + n_exp = T.dynamic("n") + m_foo = T.dynamic("m") + n_foo = T.dynamic("n") + @tvm.script.ir_module class TestCallTIRRewrite: @Ts.prim_func def exp(A_handle: T.handle, B_handle: T.handle): - m = T.int64() - n = T.int64() - A = T.match_buffer(A_handle, (m, n), "float32") - B = T.match_buffer(B_handle, (m, n), "float32") + A = T.match_buffer(A_handle, (m_exp, n_exp), "float32") + B = T.match_buffer(B_handle, (m_exp, n_exp), "float32") T.evaluate(0) @R.function - def foo(x: R.Tensor(("m", "n"), "float32")): + def foo(x: R.Tensor((m_foo, n_foo), "float32")): # we expect RemovePurityChecking to have been used before this point R.func_attr({"relax.force_pure": True}) - m, n = T.int64(), T.int64() - gv0 = R.call_tir(TestCallTIRRewrite.exp, (x,), R.Tensor((m, n), dtype="float32")) + gv0 = R.call_tir( + TestCallTIRRewrite.exp, (x,), R.Tensor((m_foo, n_foo), dtype="float32") + ) return gv0 mod = TestCallTIRRewrite @@ -323,13 +329,15 @@ def nested() -> R.Any: def test_call_dps_packed_rewrite(): + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class TestCallDPSPackedRewrite: @R.function - def foo(x: R.Tensor(("m", "n"), "float32")): + def foo(x: R.Tensor((m, n), "float32")): # we expect RemovePurityChecking to have been used before this point R.func_attr({"relax.force_pure": True}) - m, n = T.int64(), T.int64() gv0 = R.call_dps_packed("test.op.identity", (x,), R.Tensor((m, n), dtype="float32")) return gv0 diff --git a/tests/python/relax/test_transform_adjust_matmul_order.py b/tests/python/relax/test_transform_adjust_matmul_order.py index 29603489f34a..170702a67629 100644 --- a/tests/python/relax/test_transform_adjust_matmul_order.py +++ b/tests/python/relax/test_transform_adjust_matmul_order.py @@ -173,6 +173,10 @@ def main( Expected = Before +lora_r_before = T.dynamic("lora_r") +lora_r_expected = T.dynamic("lora_r") + + class TestLHSDynamic(Base): """Prefer (x*A)*B instead of x*(A*B) @@ -192,8 +196,8 @@ class Before: @R.function def main( x: R.Tensor([16]), - A: R.Tensor([16, "lora_r"]), - B: R.Tensor(["lora_r", 32]), + A: R.Tensor([16, lora_r_before]), + B: R.Tensor([lora_r_before, 32]), ) -> R.Tensor([32]): weight: R.Tensor([16, 32]) = R.matmul(A, B) out: R.Tensor([32]) = R.matmul(x, weight) @@ -204,15 +208,18 @@ class Expected: @R.function def main( x: R.Tensor([16]), - A: R.Tensor([16, "lora_r"]), - B: R.Tensor(["lora_r", 32]), + A: R.Tensor([16, lora_r_expected]), + B: R.Tensor([lora_r_expected, 32]), ) -> R.Tensor([32]): - lora_r = T.int64() - x: R.Tensor([lora_r]) = R.matmul(x, A) + x: R.Tensor([lora_r_expected]) = R.matmul(x, A) x: R.Tensor([32]) = R.matmul(x, B) return x +lora_r_before = T.dynamic("lora_r") +lora_r_expected = T.dynamic("lora_r") + + class TestRHSDynamic(Base): """Prefer A*(B*x) instead of (A*B)*x @@ -229,8 +236,8 @@ class Before: @R.function def main( x: R.Tensor([16]), - A: R.Tensor([32, "lora_r"]), - B: R.Tensor(["lora_r", 16]), + A: R.Tensor([32, lora_r_before]), + B: R.Tensor([lora_r_before, 16]), ) -> R.Tensor([32]): weight: R.Tensor([32, 16]) = R.matmul(A, B) out: R.Tensor([32]) = R.matmul(weight, x) @@ -241,11 +248,10 @@ class Expected: @R.function def main( x: R.Tensor([16]), - A: R.Tensor([32, "lora_r"]), - B: R.Tensor(["lora_r", 16]), + A: R.Tensor([32, lora_r_expected]), + B: R.Tensor([lora_r_expected, 16]), ) -> R.Tensor([32]): - lora_r = T.int64() - x: R.Tensor([lora_r]) = R.matmul(B, x) + x: R.Tensor([lora_r_expected]) = R.matmul(B, x) x: R.Tensor([32]) = R.matmul(A, x) return x @@ -264,6 +270,10 @@ class TestIdempotentRHSDynamic(Base): Expected = TestRHSDynamic.Expected +batch_size_before = T.dynamic("batch_size") +lora_r_before = T.dynamic("lora_r") + + class TestDynamicWithBatchSymbolic1(Base): """When both batch_size and lora_r are symbolic and it cannot be proven which is cheaper, LHS or RHS, maintain the existing order. @@ -290,13 +300,12 @@ class TestDynamicWithBatchSymbolic1(Base): class Before: @R.function def main( - x: R.Tensor(["batch_size", 1, 16]), - A: R.Tensor([16, "lora_r"]), - B: R.Tensor(["lora_r", 32]), - ) -> R.Tensor(["batch_size", 1, 32]): - batch_size = T.int64() + x: R.Tensor([batch_size_before, 1, 16]), + A: R.Tensor([16, lora_r_before]), + B: R.Tensor([lora_r_before, 32]), + ) -> R.Tensor([batch_size_before, 1, 32]): weight: R.Tensor([16, 32]) = R.matmul(A, B) - out: R.Tensor([batch_size, 1, 32]) = R.matmul(x, weight) + out: R.Tensor([batch_size_before, 1, 32]) = R.matmul(x, weight) return out Expected = Before @@ -368,6 +377,10 @@ def main( return out +batch_size_before = T.dynamic("batch_size") +lora_r_before = T.dynamic("lora_r") + + class TestDynamicWithBatchSymbolic2(Base): """When both batch_size and lora_r are symbolic and it cannot be proven which is cheaper, LHS or RHS, maintain the existing order. @@ -394,13 +407,12 @@ class TestDynamicWithBatchSymbolic2(Base): class Before: @R.function def main( - x: R.Tensor(["batch_size", 16, 1]), - A: R.Tensor([32, "lora_r"]), - B: R.Tensor(["lora_r", 16]), - ) -> R.Tensor(["batch_size", 32, 1]): - batch_size = T.int64() + x: R.Tensor([batch_size_before, 16, 1]), + A: R.Tensor([32, lora_r_before]), + B: R.Tensor([lora_r_before, 16]), + ) -> R.Tensor([batch_size_before, 32, 1]): weight: R.Tensor([32, 16]) = R.matmul(A, B) - out: R.Tensor([batch_size, 32, 1]) = R.matmul(weight, x) + out: R.Tensor([batch_size_before, 32, 1]) = R.matmul(weight, x) return out Expected = Before @@ -472,6 +484,12 @@ def main( return out +M_before = T.dynamic("M") +N_before = T.dynamic("N") +P_before = T.dynamic("P") +Q_before = T.dynamic("Q") + + class TestNoOpForFullyDynamicOnLHS(Base): """Keep existing order if no benefit can be proven @@ -494,9 +512,9 @@ class TestNoOpForFullyDynamicOnLHS(Base): class Before: @R.function def main( - A: R.Tensor(["M", "N"]), - B: R.Tensor(["N", "P"]), - C: R.Tensor(["P", "Q"]), + A: R.Tensor([M_before, N_before]), + B: R.Tensor([N_before, P_before]), + C: R.Tensor([P_before, Q_before]), ): out = R.matmul(R.matmul(A, B), C) return out @@ -504,6 +522,12 @@ def main( Expected = Before +M_before = T.dynamic("M") +N_before = T.dynamic("N") +P_before = T.dynamic("P") +Q_before = T.dynamic("Q") + + class TestNoOpForFullyDynamicOnRHS(Base): """Keep existing order if no benefit can be proven @@ -515,9 +539,9 @@ class TestNoOpForFullyDynamicOnRHS(Base): class Before: @R.function def main( - A: R.Tensor(["M", "N"]), - B: R.Tensor(["N", "P"]), - C: R.Tensor(["P", "Q"]), + A: R.Tensor([M_before, N_before]), + B: R.Tensor([N_before, P_before]), + C: R.Tensor([P_before, Q_before]), ): out = R.matmul(A, R.matmul(B, C)) return out @@ -626,6 +650,10 @@ def main( Expected = Before +lora_r_before = T.dynamic("lora_r") +lora_r_expected = T.dynamic("lora_r") + + class TestRHSPermuteDimsDynamic(Base): """Prefer (x*A)*B instead of x*(A*B) @@ -645,8 +673,8 @@ class Before: @R.function def main( x: R.Tensor([16]), - A: R.Tensor([32, "lora_r"]), - B: R.Tensor(["lora_r", 16]), + A: R.Tensor([32, lora_r_before]), + B: R.Tensor([lora_r_before, 16]), ) -> R.Tensor([32]): linear_weight: R.Tensor([32, 16]) = R.matmul(A, B) matmul_weight: R.Tensor([16, 32]) = R.permute_dims(linear_weight) @@ -658,17 +686,22 @@ class Expected: @R.function def main( x: R.Tensor([16]), - A: R.Tensor([32, "lora_r"]), - B: R.Tensor(["lora_r", 16]), + A: R.Tensor([32, lora_r_expected]), + B: R.Tensor([lora_r_expected, 16]), ) -> R.Tensor([32]): - lora_r = T.int64() B_transpose = R.permute_dims(B) - x: R.Tensor([lora_r]) = R.matmul(x, B_transpose) + x: R.Tensor([lora_r_expected]) = R.matmul(x, B_transpose) A_transpose = R.permute_dims(A) x: R.Tensor([32]) = R.matmul(x, A_transpose) return x +lora_r_before = T.dynamic("lora_r") +batch_size_before = T.dynamic("batch_size") +lora_r_expected = T.dynamic("lora_r") +batch_size_expected = T.dynamic("batch_size") + + class TestRHSPermuteDimsWithDynamicBatch(Base): """Prefer (x*A)*B instead of x*(A*B) @@ -698,44 +731,44 @@ class TestRHSPermuteDimsWithDynamicBatch(Base): class Before: @R.function def main( - x: R.Tensor(["batch_size", 4096]), - A: R.Tensor([4096, "lora_r"]), - B: R.Tensor(["lora_r", 4096]), - ) -> R.Tensor(["batch_size", 4096]): + x: R.Tensor([batch_size_before, 4096]), + A: R.Tensor([4096, lora_r_before]), + B: R.Tensor([lora_r_before, 4096]), + ) -> R.Tensor([batch_size_before, 4096]): R.func_attr( { "tir_var_upper_bound": {"lora_r": 2048, "batch_size": 2048}, } ) - lora_r = T.int64() # noqa: F841 - batch_size = T.int64() linear_weight: R.Tensor([4096, 4096]) = R.matmul(A, B) matmul_weight: R.Tensor([4096, 4096]) = R.permute_dims(linear_weight) - out: R.Tensor([batch_size, 4096]) = R.matmul(x, matmul_weight) + out: R.Tensor([batch_size_before, 4096]) = R.matmul(x, matmul_weight) return out @I.ir_module class Expected: @R.function def main( - x: R.Tensor(["batch_size", 4096]), - A: R.Tensor([4096, "lora_r"]), - B: R.Tensor(["lora_r", 4096]), - ) -> R.Tensor(["batch_size", 4096]): + x: R.Tensor([batch_size_expected, 4096]), + A: R.Tensor([4096, lora_r_expected]), + B: R.Tensor([lora_r_expected, 4096]), + ) -> R.Tensor([batch_size_expected, 4096]): R.func_attr( { "tir_var_upper_bound": {"lora_r": 2048, "batch_size": 2048}, } ) - lora_r = T.int64() - batch_size = T.int64() B_transpose = R.permute_dims(B) - x: R.Tensor([batch_size, lora_r]) = R.matmul(x, B_transpose) + x: R.Tensor([batch_size_expected, lora_r_expected]) = R.matmul(x, B_transpose) A_transpose = R.permute_dims(A) - x: R.Tensor([batch_size, 4096]) = R.matmul(x, A_transpose) + x: R.Tensor([batch_size_expected, 4096]) = R.matmul(x, A_transpose) return x +lora_r_before = T.dynamic("lora_r") +lora_r_expected = T.dynamic("lora_r") + + class TestRHSPermuteDimsDynamicWithSquareMatrix(Base): """Prefer (x*A)*B instead of x*(A*B) @@ -753,8 +786,8 @@ class Before: @R.function def main( x: R.Tensor([32]), - A: R.Tensor([32, "lora_r"]), - B: R.Tensor(["lora_r", 32]), + A: R.Tensor([32, lora_r_before]), + B: R.Tensor([lora_r_before, 32]), ) -> R.Tensor([32]): linear_weight: R.Tensor([32, 32]) = R.matmul(A, B) matmul_weight: R.Tensor([32, 32]) = R.permute_dims(linear_weight) @@ -766,12 +799,11 @@ class Expected: @R.function def main( x: R.Tensor([32]), - A: R.Tensor([32, "lora_r"]), - B: R.Tensor(["lora_r", 32]), + A: R.Tensor([32, lora_r_expected]), + B: R.Tensor([lora_r_expected, 32]), ) -> R.Tensor([32]): - lora_r = T.int64() B_transpose = R.permute_dims(B) - x: R.Tensor([lora_r]) = R.matmul(x, B_transpose) + x: R.Tensor([lora_r_expected]) = R.matmul(x, B_transpose) A_transpose = R.permute_dims(A) x: R.Tensor([32]) = R.matmul(x, A_transpose) return x diff --git a/tests/python/relax/test_transform_alter_op_impl.py b/tests/python/relax/test_transform_alter_op_impl.py index 10eff2b2dd68..222922122304 100644 --- a/tests/python/relax/test_transform_alter_op_impl.py +++ b/tests/python/relax/test_transform_alter_op_impl.py @@ -256,6 +256,9 @@ def relu(arg0: T.Buffer((14,), "float32"), output: T.Buffer((14,), "float32")): Ts.writes(output[v_ax0]) output[v_ax0] = T.max(arg0[v_ax0], T.float32(0)) + p0 = T.dynamic("p0") + i0 = T.dynamic("i0") + @I.ir_module class Expected: @R.function @@ -299,9 +302,7 @@ def relax_relu_replacement( @Ts.prim_func(private=True) def remove_pad(var_input: T.handle, var_output: T.handle): T.func_attr({"operator_name": "remove_pad", "tirx.noalias": True}) - p0 = T.int64() input = T.match_buffer(var_input, (p0,)) - i0 = T.int64() output = T.match_buffer(var_output, (i0,)) # with Ts.sblock("root"): for ax0 in range(i0): diff --git a/tests/python/relax/test_transform_annotate_tir_op_pattern.py b/tests/python/relax/test_transform_annotate_tir_op_pattern.py index 88e13bc8cabd..8f413cae45ab 100644 --- a/tests/python/relax/test_transform_annotate_tir_op_pattern.py +++ b/tests/python/relax/test_transform_annotate_tir_op_pattern.py @@ -37,21 +37,22 @@ class OpPatternKind(enum.IntEnum): def test_annotate_opkind_outewisefusable(): + m = T.dynamic("m", "int32") + n = T.dynamic("n", "int32") + k = T.dynamic("k", "int32") + @tvm.script.ir_module class InputModule: @Ts.prim_func def tir_matmul(x: T.handle, y: T.handle, z: T.handle) -> None: T.func_attr({"global_symbol": "tir_matmul"}) - m = T.int32() - n = T.int32() - k = T.int32() A = T.match_buffer(x, (m, n)) B = T.match_buffer(y, (n, k)) C = T.match_buffer(z, (m, k)) - for i, j, k in T.grid(m, k, n): + for i, j, k_index in T.grid(m, k, n): with Ts.sblock("matmul"): - vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) + vi, vj, vk = Ts.axis.remap("SSR", [i, j, k_index]) with Ts.init(): C[vi, vj] = T.float32(0) C[vi, vj] = C[vi, vj] + A[vi, vk] * B[vk, vj] @@ -70,21 +71,22 @@ def tir_matmul(x: T.handle, y: T.handle, z: T.handle) -> None: ], ) def test_annotate_opkind_outewisefusable_with_cast(cast_pattern): + m = T.dynamic("m", "int32") + n = T.dynamic("n", "int32") + k = T.dynamic("k", "int32") + @tvm.script.ir_module class InputModule: @Ts.prim_func def tir_matmul(x: T.handle, y: T.handle, z: T.handle) -> None: T.func_attr({"global_symbol": "tir_matmul"}) - m = T.int32() - n = T.int32() - k = T.int32() A = T.match_buffer(x, (m, n), "float16") B = T.match_buffer(y, (n, k), "float16") C = T.match_buffer(z, (m, k), "float32") - for i, j, k in T.grid(m, k, n): + for i, j, k_index in T.grid(m, k, n): with Ts.sblock("matmul"): - vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) + vi, vj, vk = Ts.axis.remap("SSR", [i, j, k_index]) with Ts.init(): C[vi, vj] = T.float32(0) C[vi, vj] = C[vi, vj] + cast_pattern(A[vi, vk], B[vk, vj]) @@ -104,9 +106,9 @@ def tir_matmul(x: T.handle, y: T.handle, z: T.handle, m: T.int64, n: T.int64, k: B = T.match_buffer(y, (n, k)) C = T.match_buffer(z, (m, k)) - for i, j, k in T.grid(m, k, n): + for i, j, k_index in T.grid(m, k, n): with Ts.sblock("matmul"): - vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) + vi, vj, vk = Ts.axis.remap("SSR", [i, j, k_index]) with Ts.init(): C[vi, vj] = T.float32(0) C[vi, vj] = C[vi, vj] + A[vi, vk] * B[vk, vj] diff --git a/tests/python/relax/test_transform_attach_global_symbol.py b/tests/python/relax/test_transform_attach_global_symbol.py index ac05bed09fd8..39b76f11a45b 100644 --- a/tests/python/relax/test_transform_attach_global_symbol.py +++ b/tests/python/relax/test_transform_attach_global_symbol.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: F401, F841 +# ruff: noqa: F401 import pytest @@ -29,57 +29,65 @@ def test_basic(): + m_tir_matmul = T.dynamic("m") + n_tir_matmul = T.dynamic("n") + k_tir_matmul = T.dynamic("k") + m_main = T.dynamic("m") + n_main = T.dynamic("n") + k_main = T.dynamic("k") + @tvm.script.ir_module class Before: @Ts.prim_func def tir_matmul(x: T.handle, y: T.handle, z: T.handle) -> None: - m = T.int64() - n = T.int64() - k = T.int64() - A = T.match_buffer(x, (m, n)) - B = T.match_buffer(y, (n, k)) - C = T.match_buffer(z, (m, k)) - - for i, j, k in T.grid(m, k, n): + A = T.match_buffer(x, (m_tir_matmul, n_tir_matmul)) + B = T.match_buffer(y, (n_tir_matmul, k_tir_matmul)) + C = T.match_buffer(z, (m_tir_matmul, k_tir_matmul)) + + for i, j, k_tir_matmul_index in T.grid(m_tir_matmul, k_tir_matmul, n_tir_matmul): with Ts.sblock("matmul"): - vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) + vi, vj, vk = Ts.axis.remap("SSR", [i, j, k_tir_matmul_index]) with Ts.init(): C[vi, vj] = T.float32(0) C[vi, vj] = C[vi, vj] + A[vi, vk] * B[vk, vj] @R.function(private=True) def main( - x: R.Tensor(("m", "n"), "float32"), w: R.Tensor(("n", "k"), "float32") + x: R.Tensor((m_main, n_main), "float32"), w: R.Tensor((n_main, k_main), "float32") ) -> R.Tensor: - m, n, k = T.int64(), T.int64(), T.int64() - gv0 = R.call_tir(Before.tir_matmul, (x, w), R.Tensor((m, k), dtype="float32")) + gv0 = R.call_tir(Before.tir_matmul, (x, w), R.Tensor((m_main, k_main), dtype="float32")) return gv0 + m_tir_matmul = T.dynamic("m") + n_tir_matmul = T.dynamic("n") + k_tir_matmul = T.dynamic("k") + m_main = T.dynamic("m") + n_main = T.dynamic("n") + k_main = T.dynamic("k") + @tvm.script.ir_module class Expected: @Ts.prim_func def tir_matmul(x: T.handle, y: T.handle, z: T.handle) -> None: T.func_attr({"global_symbol": "tir_matmul"}) - m = T.int64() - n = T.int64() - k = T.int64() - A = T.match_buffer(x, (m, n)) - B = T.match_buffer(y, (n, k)) - C = T.match_buffer(z, (m, k)) - - for i, j, k in T.grid(m, k, n): + A = T.match_buffer(x, (m_tir_matmul, n_tir_matmul)) + B = T.match_buffer(y, (n_tir_matmul, k_tir_matmul)) + C = T.match_buffer(z, (m_tir_matmul, k_tir_matmul)) + + for i, j, k_tir_matmul_index in T.grid(m_tir_matmul, k_tir_matmul, n_tir_matmul): with Ts.sblock("matmul"): - vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) + vi, vj, vk = Ts.axis.remap("SSR", [i, j, k_tir_matmul_index]) with Ts.init(): C[vi, vj] = T.float32(0) C[vi, vj] = C[vi, vj] + A[vi, vk] * B[vk, vj] @R.function def main( - x: R.Tensor(("m", "n"), "float32"), w: R.Tensor(("n", "k"), "float32") + x: R.Tensor((m_main, n_main), "float32"), w: R.Tensor((n_main, k_main), "float32") ) -> R.Tensor: - m, n, k = T.int64(), T.int64(), T.int64() - gv0 = R.call_tir(Expected.tir_matmul, (x, w), R.Tensor((m, k), dtype="float32")) + gv0 = R.call_tir( + Expected.tir_matmul, (x, w), R.Tensor((m_main, k_main), dtype="float32") + ) return gv0 before = Before diff --git a/tests/python/relax/test_transform_bind_params.py b/tests/python/relax/test_transform_bind_params.py index fe107a7a946d..e890f88b32b2 100644 --- a/tests/python/relax/test_transform_bind_params.py +++ b/tests/python/relax/test_transform_bind_params.py @@ -14,7 +14,6 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -# ruff: noqa: F841 import numpy as np import pytest @@ -76,20 +75,21 @@ def main(x: R.Tensor((16, 16), "float32"), w: R.Tensor((16, 16), "float32")) -> def test_bind_params_symbolic_vars(): + batch = T.dynamic("batch") + k = T.dynamic("k") + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Before: @R.function def main( - x: R.Tensor(("batch", "m"), dtype="float32"), - w0: R.Tensor(("n", "m"), dtype="float32"), - b0: R.Tensor(("n",), dtype="float32"), - w1: R.Tensor(("k", "n"), dtype="float32"), - b1: R.Tensor(("k",), dtype="float32"), - ) -> R.Tensor(("batch", "k"), dtype="float32"): - batch = T.int64() - k = T.int64() - m = T.int64() - n = T.int64() + x: R.Tensor((batch, m), dtype="float32"), + w0: R.Tensor((n, m), dtype="float32"), + b0: R.Tensor((n,), dtype="float32"), + w1: R.Tensor((k, n), dtype="float32"), + b1: R.Tensor((k,), dtype="float32"), + ) -> R.Tensor((batch, k), dtype="float32"): with R.dataflow(): lv0 = R.call_dps_packed( "linear0", (x, w0, b0), out_ty=R.Tensor((batch, n), dtype="float32") diff --git a/tests/python/relax/test_transform_bind_symbolic_vars.py b/tests/python/relax/test_transform_bind_symbolic_vars.py index a517c553f6ad..48bbc0c9335f 100644 --- a/tests/python/relax/test_transform_bind_symbolic_vars.py +++ b/tests/python/relax/test_transform_bind_symbolic_vars.py @@ -29,17 +29,19 @@ def test_bind_tensors(): """Symbolic variables may occur in Tensor shapes""" + batch = T.dynamic("batch") + n = T.dynamic("n") + k = T.dynamic("k") + m = T.dynamic("m") + @tvm.script.ir_module class Before: @R.function def main( - x: R.Tensor(("batch", "m"), dtype="float32"), - w0: R.Tensor(("m", "n"), dtype="float32"), - w1: R.Tensor(("k", 10), dtype="float32"), - ) -> R.Tensor(("batch", "k"), dtype="float32"): - batch = T.int64() - n = T.int64() - k = T.int64() + x: R.Tensor((batch, m), dtype="float32"), + w0: R.Tensor((m, n), dtype="float32"), + w1: R.Tensor((k, 10), dtype="float32"), + ) -> R.Tensor((batch, k), dtype="float32"): with R.dataflow(): lv0 = R.call_dps_packed( "test0", (x, w0), out_ty=R.Tensor((batch, n), dtype="float32") @@ -54,15 +56,17 @@ def main( target_func_name = "main" After = relax.transform.BindSymbolicVars(symvar_map, target_func_name)(Before) + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Expected: @R.function def main( - x: R.Tensor((1, "m"), dtype="float32"), - w0: R.Tensor(("m", "n"), dtype="float32"), + x: R.Tensor((1, m), dtype="float32"), + w0: R.Tensor((m, n), dtype="float32"), w1: R.Tensor((3, 10), dtype="float32"), ) -> R.Tensor((1, 3), dtype="float32"): - n = T.int64() with R.dataflow(): lv0 = R.call_dps_packed("test0", (x, w0), out_ty=R.Tensor((1, n), dtype="float32")) out = R.call_dps_packed( @@ -77,17 +81,19 @@ def main( def test_bind_shape(): """Symbolic variables may occur in ShapeExpr""" + batch = T.dynamic("batch") + n = T.dynamic("n") + k = T.dynamic("k") + m = T.dynamic("m") + @tvm.script.ir_module class Before: @R.function def main( - x: R.Shape(("batch", "m")), - w0: R.Shape(("m", "n")), - w1: R.Shape(("k", 10)), - ) -> R.Shape(("batch", "k")): - batch = T.int64() - n = T.int64() - k = T.int64() + x: R.Shape((batch, m)), + w0: R.Shape((m, n)), + w1: R.Shape((k, 10)), + ) -> R.Shape((batch, k)): with R.dataflow(): lv0 = R.call_dps_packed("test0", (x, w0), out_ty=R.Tensor((batch, n))) out = R.call_dps_packed("test1", (lv0, w1), out_ty=R.Tensor((batch, k))) @@ -98,13 +104,13 @@ def main( target_func_name = "main" After = relax.transform.BindSymbolicVars(symvar_map, target_func_name)(Before) + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Expected: @R.function - def main(x: R.Shape([1, "m"]), w0: R.Shape(["m", "n"]), w1: R.Shape([3, 10])) -> R.Shape( - [1, 3] - ): - n = T.int64() + def main(x: R.Shape([1, m]), w0: R.Shape([m, n]), w1: R.Shape([3, 10])) -> R.Shape([1, 3]): with R.dataflow(): lv0 = R.call_dps_packed("test0", (x, w0), out_ty=R.Tensor((1, n))) out = R.call_dps_packed("test1", (lv0, w1), out_ty=R.Tensor((1, 3))) @@ -117,18 +123,19 @@ def main(x: R.Shape([1, "m"]), w0: R.Shape(["m", "n"]), w1: R.Shape([3, 10])) -> def test_arith(): """Symbolic shapes may use TIR arithmetic expressions""" + batch = T.dynamic("batch") + m = T.dynamic("m") + n = T.dynamic("n") + k = T.dynamic("k") + @tvm.script.ir_module class Before: @R.function def main( - x: R.Tensor(("batch", "m-1"), dtype="float32"), - w0: R.Tensor(("m", "n"), dtype="float32"), - w1: R.Tensor(("k", 10), dtype="float32"), - ) -> R.Tensor(("batch", "k*m"), dtype="float32"): - batch = T.int64() - m = T.int64() - n = T.int64() - k = T.int64() + x: R.Tensor((batch, m - 1), dtype="float32"), + w0: R.Tensor((m, n), dtype="float32"), + w1: R.Tensor((k, 10), dtype="float32"), + ) -> R.Tensor((batch, k * m), dtype="float32"): with R.dataflow(): lv0 = R.call_dps_packed( "test0", @@ -147,15 +154,16 @@ def main( target_func_name = "main" After = relax.transform.BindSymbolicVars(symvar_map, target_func_name)(Before) + n = T.dynamic("n") + @I.ir_module class Expected: @R.function def main( x: R.Tensor((1, 2), dtype="float32"), - w0: R.Tensor((3, "n"), dtype="float32"), + w0: R.Tensor((3, n), dtype="float32"), w1: R.Tensor((2, 10), dtype="float32"), ) -> R.Tensor((1, 6), dtype="float32"): - n = T.int64() with R.dataflow(): lv0 = R.call_dps_packed( "test0", (x, w0), out_ty=R.Tensor((1, n + 3), dtype="float32") @@ -172,24 +180,32 @@ def main( def test_bind_multiple_variables_by_name(): """String names may be used to replace across multiple functions""" + m_main_1 = T.dynamic("m") + n_main_1 = T.dynamic("n") + m_main_2 = T.dynamic("m") + n_main_2 = T.dynamic("n") + @tvm.script.ir_module class Before: @R.function - def main_1(x: R.Tensor(("m", "n"), dtype="float32")): + def main_1(x: R.Tensor((m_main_1, n_main_1), dtype="float32")): return x @R.function - def main_2(x: R.Tensor(("m", "n"), dtype="float32")): + def main_2(x: R.Tensor((m_main_2, n_main_2), dtype="float32")): return x + m_main_1 = T.dynamic("m") + m_main_2 = T.dynamic("m") + @tvm.script.ir_module class Expected: @R.function - def main_1(x: R.Tensor(("m", 16), dtype="float32")): + def main_1(x: R.Tensor((m_main_1, 16), dtype="float32")): return x @R.function - def main_2(x: R.Tensor(("m", 16), dtype="float32")): + def main_2(x: R.Tensor((m_main_2, 16), dtype="float32")): return x After = relax.transform.BindSymbolicVars({"n": 16})(Before) @@ -199,24 +215,33 @@ def main_2(x: R.Tensor(("m", 16), dtype="float32")): def test_bind_single_variable_by_identity(): """TIR variables may be used to replace a specific var""" + m_main_1 = T.dynamic("m") + n_main_1 = T.dynamic("n") + m_main_2 = T.dynamic("m") + n_main_2 = T.dynamic("n") + @tvm.script.ir_module class Before: @R.function - def main_1(x: R.Tensor(("m", "n"), dtype="float32")): + def main_1(x: R.Tensor((m_main_1, n_main_1), dtype="float32")): return x @R.function - def main_2(x: R.Tensor(("m", "n"), dtype="float32")): + def main_2(x: R.Tensor((m_main_2, n_main_2), dtype="float32")): return x + m_main_1 = T.dynamic("m") + m_main_2 = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main_1(x: R.Tensor(("m", 16), dtype="float32")): + def main_1(x: R.Tensor((m_main_1, 16), dtype="float32")): return x @R.function - def main_2(x: R.Tensor(("m", "n"), dtype="float32")): + def main_2(x: R.Tensor((m_main_2, n), dtype="float32")): return x main_1_n = Before["main_1"].params[0].ty.shape[1] @@ -227,24 +252,33 @@ def main_2(x: R.Tensor(("m", "n"), dtype="float32")): def test_bind_single_variable_by_function_name(): """Variable name and function name may be used to replace a specific var""" + m_main_1 = T.dynamic("m") + n_main_1 = T.dynamic("n") + m_main_2 = T.dynamic("m") + n_main_2 = T.dynamic("n") + @tvm.script.ir_module class Before: @R.function - def main_1(x: R.Tensor(("m", "n"), dtype="float32")): + def main_1(x: R.Tensor((m_main_1, n_main_1), dtype="float32")): return x @R.function - def main_2(x: R.Tensor(("m", "n"), dtype="float32")): + def main_2(x: R.Tensor((m_main_2, n_main_2), dtype="float32")): return x + m_main_1 = T.dynamic("m") + m_main_2 = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main_1(x: R.Tensor(("m", 16), dtype="float32")): + def main_1(x: R.Tensor((m_main_1, 16), dtype="float32")): return x @R.function - def main_2(x: R.Tensor(("m", "n"), dtype="float32")): + def main_2(x: R.Tensor((m_main_2, n), dtype="float32")): return x After = relax.transform.BindSymbolicVars({"n": 16}, "main_1")(Before) @@ -254,10 +288,13 @@ def main_2(x: R.Tensor(("m", "n"), dtype="float32")): def test_error_for_unused_replacement(): """Each replacement must be used""" + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Before: @R.function - def main(x: R.Tensor(("m", "n"), dtype="float32")): + def main(x: R.Tensor((m, n), dtype="float32")): return x with pytest.raises(RuntimeError): diff --git a/tests/python/relax/test_transform_bundle_model_params.py b/tests/python/relax/test_transform_bundle_model_params.py index 60c56be30f95..54078e7f9d9c 100644 --- a/tests/python/relax/test_transform_bundle_model_params.py +++ b/tests/python/relax/test_transform_bundle_model_params.py @@ -299,8 +299,8 @@ class Before: def main( cond: R.Tensor((), "bool"), extent: T.int64, - weight: R.Tensor(["extent"], "float32"), - ) -> R.Tensor(["extent"], "float32"): + weight: 'R.Tensor([extent], "float32")', # noqa: F821 + ) -> 'R.Tensor([extent], "float32")': # noqa: F821 R.func_attr({"num_input": 1}) if cond: out = R.add(weight, weight) @@ -338,7 +338,7 @@ class Before: def main( x: R.Tensor(dtype="float32", ndim=1), extent: T.int64, - weight: R.Tensor(["extent"], "float32"), + weight: 'R.Tensor([extent], "float32")', ): R.func_attr({"num_input": 1}) out = R.add(x, weight) diff --git a/tests/python/relax/test_transform_canonicalize_bindings.py b/tests/python/relax/test_transform_canonicalize_bindings.py index 4265f244f84a..1b102293371c 100644 --- a/tests/python/relax/test_transform_canonicalize_bindings.py +++ b/tests/python/relax/test_transform_canonicalize_bindings.py @@ -151,22 +151,26 @@ def main(x: R.Tensor) -> R.Any: def test_match_cast(): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class TestMatchCast: @R.function def main(x: R.Tensor): q = x - m, n = T.int64(), T.int64() z = R.match_cast(q, R.Tensor((m, n))) w = z return w + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Expected: @R.function def main(x: R.Tensor): # can't get rid of z because its ty is different from x's - m, n = T.int64(), T.int64() z = R.match_cast(x, R.Tensor((m, n))) return z @@ -174,11 +178,13 @@ def main(x: R.Tensor): def test_same_shape(): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class TestSameShape: @R.function - def main(x: R.Tensor(("m", "n"), "float32")): - m, n = T.int64(), T.int64() + def main(x: R.Tensor((m, n), "float32")): y = x # trivial check z = R.match_cast(x, R.Tensor((m, n), "float32")) @@ -186,10 +192,13 @@ def main(x: R.Tensor(("m", "n"), "float32")): q = R.add(w, y) return R.add(q, w) + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "float32")): + def main(x: R.Tensor((m, n), "float32")): # the trivial check is canonicalized into a var binding # and then eliminated q = R.add(x, x) @@ -199,6 +208,9 @@ def main(x: R.Tensor(("m", "n"), "float32")): def test_change_shape(): + o = T.dynamic("o") + p = T.dynamic("p") + @I.ir_module class TestChangeShape: @R.function @@ -209,17 +221,18 @@ def main(x: R.Tensor(ndim=2)): # rather than a symbolic shape, these new shape vars # cannot be expressed in terms of previous variables. # Therefore, the match cast must be retained. - o, p = T.int64(), T.int64() z = R.match_cast(x, R.Tensor((o, p))) w = z q = R.add(w, y) return R.add(q, w) + o = T.dynamic("o") + p = T.dynamic("p") + @I.ir_module class Expected: @R.function def main(x: R.Tensor(ndim=2)): - o, p = T.int64(), T.int64() z = R.match_cast(x, R.Tensor((o, p))) # the ty field on q will need to be updated q = R.add(z, x) @@ -229,28 +242,33 @@ def main(x: R.Tensor(ndim=2)): def test_replace_symbolic_variable_and_remove_match_cast(): + o = T.dynamic("o") + p = T.dynamic("p") + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class TestChangeShape: @R.function - def main(x: R.Tensor(("m", "n"))): + def main(x: R.Tensor((m, n))): y = x # The MatchCast is non-trivial, as it introduces new shape # vars. However, the new shape vars are redundant, and # are replaced by canonicalization. After replacing the # new shape vars, the MatchCast is trivial and may be # removed. - o, p = T.int64(), T.int64() z = R.match_cast(x, R.Tensor((o, p))) w = z q = R.add(w, y) return R.add(q, w) + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"))): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n))): q: R.Tensor([m, n]) = R.add(x, x) return R.add(q, x) @@ -270,21 +288,28 @@ def test_replace_symbolic_variable_and_remove_match_cast_of_tuple(): """ + o = T.dynamic("o") + p = T.dynamic("p") + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Before: @R.function - def main(x: R.Tuple(R.Tensor(("m", "n")))): + def main(x: R.Tuple(R.Tensor((m, n)))): y = x - o, p = T.int64(), T.int64() z = R.match_cast(x, R.Tuple(R.Tensor((o, p)))) w = z q = R.add(w[0], y[0]) return R.add(q, w[0]) + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Expected: @R.function - def main(x: R.Tuple(R.Tensor(("m", "n")))): + def main(x: R.Tuple(R.Tensor((m, n)))): q = R.add(x[0], x[0]) return R.add(q, x[0]) @@ -368,6 +393,10 @@ def test_fold_variables_from_match_cast(): """ + N1 = T.dynamic("N1") + M = T.dynamic("M") + N2 = T.dynamic("N2") + @I.ir_module class Before: @R.function @@ -376,10 +405,6 @@ def main( A: R.Tensor([16, 16], dtype="float32"), B: R.Tensor([16, 16], dtype="float32"), ): - N1 = T.int64() - M = T.int64() - N2 = T.int64() - # The symbolic variables `N1`, `N2` and `M` are defined by # these `R.match_cast` statements. Since the inputs have # a known shape, the values of these symbolic variables @@ -407,6 +432,10 @@ def main( ) return (proj_A, proj_B) + N1 = T.dynamic("N1") + M = T.dynamic("M") + N2 = T.dynamic("N2") + @I.ir_module class Expected: @R.function @@ -418,9 +447,6 @@ def main( # Shape annotations use the inferred static values, but runtime # primitive arguments remain symbolic. Keep the match-casts that # define the symbols used by those runtime arguments. - N1 = T.int64() - M = T.int64() - N2 = T.int64() lhs_A = R.match_cast(A, R.Tensor([N1, M], dtype="float32")) lhs_B = R.match_cast(B, R.Tensor([N2, M], dtype="float32")) @@ -455,6 +481,10 @@ def test_inconsistent_match_cast_raises_error(): """ + N1 = T.dynamic("N1") + M = T.dynamic("M") + N2 = T.dynamic("N2") + @I.ir_module class Before: @R.function @@ -463,10 +493,6 @@ def main( A: R.Tensor([16, 16], dtype="float32"), B: R.Tensor([32, 32], dtype="float32"), ): - N1 = T.int64() - M = T.int64() - N2 = T.int64() - # These R.match_cast statements define inconsistent values # for the symbolic shape parameters. lhs_A = R.match_cast(A, R.Tensor([N1, M], dtype="float32")) @@ -503,18 +529,18 @@ def test_match_cast_may_have_distinct_values_in_branches(): """ + N = T.dynamic("N") + M = T.dynamic("M") + @I.ir_module class Before: @R.function def main( - state: R.Tensor(["N"], dtype="float32"), - A: R.Tensor(["M", 16], dtype="float32"), - B: R.Tensor(["M", 32], dtype="float32"), + state: R.Tensor([N], dtype="float32"), + A: R.Tensor([M, 16], dtype="float32"), + B: R.Tensor([M, 32], dtype="float32"), scale: T.float32, ): - N = T.int64() - M = T.int64() - if N == 16: weights: R.Tensor([M, 16], "float32") = A * scale weights: R.Tensor([M, N], "float32") = R.match_cast( @@ -534,18 +560,18 @@ def main( return out + N = T.dynamic("N") + M = T.dynamic("M") + @I.ir_module class Expected: @R.function def main( - state: R.Tensor(["N"], dtype="float32"), - A: R.Tensor(["M", 16], dtype="float32"), - B: R.Tensor(["M", 32], dtype="float32"), + state: R.Tensor([N], dtype="float32"), + A: R.Tensor([M, 16], dtype="float32"), + B: R.Tensor([M, 32], dtype="float32"), scale: T.float32, ): - N = T.int64() - M = T.int64() - if N == 16: # Prior to the R.match_cast, the weights: R.Tensor([M, 16], "float32") = A * scale @@ -870,10 +896,12 @@ def test_canonicalize_with_updated_ty(): in order to provide better type. """ + n = T.dynamic("n") + @I.ir_module class Before: @R.function(private=True) - def main(A: R.Tensor(("n", 16), dtype="int32")) -> R.Tensor(("n", 16), dtype="int32"): + def main(A: R.Tensor((n, 16), dtype="int32")) -> R.Tensor((n, 16), dtype="int32"): # CanonicalizeBindings recognizes this trivial binding, and # replaces `B` with `A`. B = A @@ -889,11 +917,12 @@ def main(A: R.Tensor(("n", 16), dtype="int32")) -> R.Tensor(("n", 16), dtype="in # version of `C` with `ndim=2`. return C + n = T.dynamic("n") + @I.ir_module class Expected: @R.function(private=True) - def main(A: R.Tensor(("n", 16), dtype="int32")) -> R.Tensor(("n", 16), dtype="int32"): - n = T.int64() + def main(A: R.Tensor((n, 16), dtype="int32")) -> R.Tensor((n, 16), dtype="int32"): C: R.Tensor([n, 16], "int32") = R.add(A, A) return C @@ -1195,11 +1224,13 @@ def test_canonicalization_causes_ty_update(): class. """ + vocab_size = T.dynamic("vocab_size") + @I.ir_module class Before: @R.function def transform_params( - A: R.Tensor(("vocab_size", 4096), dtype="float16"), + A: R.Tensor((vocab_size, 4096), dtype="float16"), B: R.Tensor((6144, 4096), dtype="float16"), ): with R.dataflow(): @@ -1230,14 +1261,15 @@ def transform_params( # with `shape=[vocab_size,4096]`. return E + vocab_size = T.dynamic("vocab_size") + @I.ir_module class Expected: @R.function def transform_params( - A: R.Tensor(("vocab_size", 4096), dtype="float16"), + A: R.Tensor((vocab_size, 4096), dtype="float16"), B: R.Tensor((6144, 4096), dtype="float16"), ): - vocab_size = T.int64() with R.dataflow(): E: R.Tuple( R.Tensor((vocab_size, 4096), dtype="float16"), diff --git a/tests/python/relax/test_transform_codegen_pass.py b/tests/python/relax/test_transform_codegen_pass.py index 7f4293b136a8..166331f61ab1 100644 --- a/tests/python/relax/test_transform_codegen_pass.py +++ b/tests/python/relax/test_transform_codegen_pass.py @@ -297,57 +297,62 @@ def rename_main(mod): def test_dynamic_shape(): import tvm.relax.backend.cuda.cublas + r1_main = T.dynamic("r1") + r2 = T.dynamic("r2") + r1_fused_relax_matmul_cublas = T.dynamic("r1") + @I.ir_module class Before: @R.function def main( x: R.Tensor((1, 4096), dtype="float16"), - w1: R.Tensor((4096, "r1"), dtype="float16"), - w2: R.Tensor((4096, "r2"), dtype="float16"), - ) -> R.Tuple(R.Tensor((1, "r1"), dtype="float16"), R.Tensor((1, "r2"), dtype="float16")): - r1 = T.int64() - r2 = T.int64() + w1: R.Tensor((4096, r1_main), dtype="float16"), + w2: R.Tensor((4096, r2), dtype="float16"), + ) -> R.Tuple(R.Tensor((1, r1_main), dtype="float16"), R.Tensor((1, r2), dtype="float16")): cls = Before with R.dataflow(): - lv: R.Tensor((1, r1), dtype="float16") = cls.fused_relax_matmul_cublas(x, w1) + lv: R.Tensor((1, r1_main), dtype="float16") = cls.fused_relax_matmul_cublas(x, w1) lv1: R.Tensor((1, r2), dtype="float16") = cls.fused_relax_matmul_cublas(x, w2) gv: R.Tuple( - R.Tensor((1, r1), dtype="float16"), R.Tensor((1, r2), dtype="float16") + R.Tensor((1, r1_main), dtype="float16"), R.Tensor((1, r2), dtype="float16") ) = (lv, lv1) R.output(gv) return gv @R.function def fused_relax_matmul_cublas( - x: R.Tensor((1, 4096), dtype="float16"), w1: R.Tensor((4096, "r1"), dtype="float16") - ) -> R.Tensor((1, "r1"), dtype="float16"): - r1 = T.int64() + x: R.Tensor((1, 4096), dtype="float16"), + w1: R.Tensor((4096, r1_fused_relax_matmul_cublas), dtype="float16"), + ) -> R.Tensor((1, r1_fused_relax_matmul_cublas), dtype="float16"): R.func_attr({"Codegen": "cublas"}) @R.function def gv( x_1: R.Tensor((1, 4096), dtype="float16"), - w1_1: R.Tensor((4096, r1), dtype="float16"), - ) -> R.Tensor((1, r1), dtype="float16"): + w1_1: R.Tensor((4096, r1_fused_relax_matmul_cublas), dtype="float16"), + ) -> R.Tensor((1, r1_fused_relax_matmul_cublas), dtype="float16"): R.func_attr({"Composite": "cublas.matmul"}) with R.dataflow(): - gv_1: R.Tensor((1, r1), dtype="float16") = R.matmul(x_1, w1_1, out_dtype=None) + gv_1: R.Tensor((1, r1_fused_relax_matmul_cublas), dtype="float16") = R.matmul( + x_1, w1_1, out_dtype=None + ) R.output(gv_1) return gv_1 - gv1: R.Tensor((1, r1), dtype="float16") = gv(x, w1) + gv1: R.Tensor((1, r1_fused_relax_matmul_cublas), dtype="float16") = gv(x, w1) return gv1 + r1 = T.dynamic("r1") + r2 = T.dynamic("r2") + @I.ir_module class Expected: @R.function def main( x: R.Tensor((1, 4096), dtype="float16"), - w1: R.Tensor((4096, "r1"), dtype="float16"), - w2: R.Tensor((4096, "r2"), dtype="float16"), - ) -> R.Tuple(R.Tensor((1, "r1"), dtype="float16"), R.Tensor((1, "r2"), dtype="float16")): - r1 = T.int64() - r2 = T.int64() + w1: R.Tensor((4096, r1), dtype="float16"), + w2: R.Tensor((4096, r2), dtype="float16"), + ) -> R.Tuple(R.Tensor((1, r1), dtype="float16"), R.Tensor((1, r2), dtype="float16")): with R.dataflow(): lv = R.call_dps_packed( "fused_relax_matmul_cublas", diff --git a/tests/python/relax/test_transform_combine_parallel_matmul.py b/tests/python/relax/test_transform_combine_parallel_matmul.py index 0269d7b070a2..9f85e1a84d7c 100644 --- a/tests/python/relax/test_transform_combine_parallel_matmul.py +++ b/tests/python/relax/test_transform_combine_parallel_matmul.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: E731, F401, F841 +# ruff: noqa: E731, F401 import pytest import tvm.testing @@ -538,13 +538,14 @@ def test_combine_matmul_of_static_and_dynamic_shapes(): """ + M = T.dynamic("M") + @R.function(private=True) def before( x: R.Tensor((2, 1024, 640), "float32"), w0: R.Tensor((640, 640), "float32"), - w1: R.Tensor((640, "M"), "float32"), + w1: R.Tensor((640, M), "float32"), ): - M = T.int64() with R.dataflow(): lv0 = R.matmul(x, w0) lv1 = R.matmul(x, w1) @@ -552,15 +553,16 @@ def before( R.output(out) return out + M = T.dynamic("M") + @R.function(private=True) def expected( x: R.Tensor((2, 1024, 640), dtype="float32"), w0: R.Tensor((640, 640), dtype="float32"), - w1: R.Tensor((640, "M"), dtype="float32"), + w1: R.Tensor((640, M), dtype="float32"), ) -> R.Tuple( - R.Tensor((2, 1024, 640), dtype="float32"), R.Tensor((2, 1024, "M"), dtype="float32") + R.Tensor((2, 1024, 640), dtype="float32"), R.Tensor((2, 1024, M), dtype="float32") ): - M = T.int64() with R.dataflow(): lv: R.Tensor((640, 640 + M), dtype="float32") = R.concat((w0, w1), axis=1) lv1: R.Tensor((2, 1024, 640 + M), dtype="float32") = R.matmul( @@ -594,13 +596,14 @@ def test_combine_matmul_of_dynamic_and_static_shapes(): concatenated weights. """ + M = T.dynamic("M") + @R.function(private=True) def before( x: R.Tensor((2, 1024, 640), "float32"), - w0: R.Tensor((640, "M"), "float32"), + w0: R.Tensor((640, M), "float32"), w1: R.Tensor((640, 640), "float32"), ): - M = T.int64() with R.dataflow(): lv0 = R.matmul(x, w0) lv1 = R.matmul(x, w1) @@ -608,15 +611,16 @@ def before( R.output(out) return out + M = T.dynamic("M") + @R.function(private=True) def expected( x: R.Tensor((2, 1024, 640), dtype="float32"), - w0: R.Tensor((640, "M"), dtype="float32"), + w0: R.Tensor((640, M), dtype="float32"), w1: R.Tensor((640, 640), dtype="float32"), ) -> R.Tuple( - R.Tensor((2, 1024, "M"), dtype="float32"), R.Tensor((2, 1024, 640), dtype="float32") + R.Tensor((2, 1024, M), dtype="float32"), R.Tensor((2, 1024, 640), dtype="float32") ): - M = T.int64() with R.dataflow(): lv: R.Tensor((640, 640 + M), dtype="float32") = R.concat((w1, w0), axis=1) lv1: R.Tensor((2, 1024, 640 + M), dtype="float32") = R.matmul( @@ -650,14 +654,16 @@ def test_limit_one_dynamic_shape_in_combined_matmul(): matmul. """ + M = T.dynamic("M") + N = T.dynamic("N") + @R.function(private=True) def before( x: R.Tensor((2, 1024, 640), "float32"), - w0: R.Tensor((640, "M"), "float32"), + w0: R.Tensor((640, M), "float32"), w1: R.Tensor((640, 640), "float32"), - w2: R.Tensor((640, "N"), "float32"), + w2: R.Tensor((640, N), "float32"), ): - M = T.int64() with R.dataflow(): lv0 = R.matmul(x, w0) lv1 = R.matmul(x, w1) @@ -666,18 +672,20 @@ def before( R.output(out) return out + M = T.dynamic("M") + N = T.dynamic("N") + @R.function(private=True) def expected( x: R.Tensor((2, 1024, 640), dtype="float32"), - w0: R.Tensor((640, "M"), dtype="float32"), + w0: R.Tensor((640, M), dtype="float32"), w1: R.Tensor((640, 640), dtype="float32"), - w2: R.Tensor((640, "N"), "float32"), + w2: R.Tensor((640, N), "float32"), ) -> R.Tuple( - R.Tensor((2, 1024, "M"), dtype="float32"), + R.Tensor((2, 1024, M), dtype="float32"), R.Tensor((2, 1024, 640), dtype="float32"), - R.Tensor((2, 1024, "N"), dtype="float32"), + R.Tensor((2, 1024, N), dtype="float32"), ): - M = T.int64() with R.dataflow(): concat_weights = R.concat((w1, w0), axis=1) concat_output = R.matmul(x, concat_weights, out_dtype="float32") diff --git a/tests/python/relax/test_transform_compute_prim_value.py b/tests/python/relax/test_transform_compute_prim_value.py index 84038a5ee702..ccaff7467074 100644 --- a/tests/python/relax/test_transform_compute_prim_value.py +++ b/tests/python/relax/test_transform_compute_prim_value.py @@ -24,19 +24,21 @@ def test_prim_value_in_assert_condition(): + N = T.dynamic("N") + @I.ir_module class Before: @R.function(pure=False) - def main(A: R.Tensor(["N"])): - N = T.int64() + def main(A: R.Tensor([N])): _ = R.assert_op(N % 16 == 0) return A + N = T.dynamic("N") + @I.ir_module class Expected: @R.function(pure=False) - def main(A: R.Tensor(["N"])): - N = T.int64() + def main(A: R.Tensor([N])): condition: T.bool = Expected.compute_symbolic_expr(R.prim_value(N)) _ = R.assert_op(condition) return A @@ -51,22 +53,24 @@ def compute_symbolic_expr(N: T.int64) -> T.bool: def test_prim_value_in_branch_condition(): + N = T.dynamic("N") + @I.ir_module class Before: @R.function(pure=False) - def main(A: R.Tensor(["N"])): - N = T.int64() + def main(A: R.Tensor([N])): if R.prim_value(N % 16 == 0): out = R.call_packed("fast_vectorized_impl", A, ty_args=[A.ty]) else: out = R.call_packed("slow_non_vectorized_impl", A, ty_args=[A.ty]) return out + N = T.dynamic("N") + @I.ir_module class Expected: @R.function(pure=False) - def main(A: R.Tensor(["N"])): - N = T.int64() + def main(A: R.Tensor([N])): condition: T.bool = Expected.compute_symbolic_expr(R.prim_value(N)) if condition: out = R.call_packed("fast_vectorized_impl", A, ty_args=[A.ty]) diff --git a/tests/python/relax/test_transform_convert_layout.py b/tests/python/relax/test_transform_convert_layout.py index 6bb55ea71a0d..bf7ccd104e4a 100644 --- a/tests/python/relax/test_transform_convert_layout.py +++ b/tests/python/relax/test_transform_convert_layout.py @@ -118,6 +118,11 @@ def main( def test_conv2d_symbolic(): + N = T.dynamic("N") + C = T.dynamic("C") + H = T.dynamic("H") + W = T.dynamic("W") + @I.ir_module class Input: @R.function @@ -125,22 +130,22 @@ def main(x: R.Tensor("float32", ndim=4), w: R.Tensor("float32", ndim=4)) -> R.Te None, "float32", ndim=4 ): with R.dataflow(): - N, C, H, W = T.int64(), T.int64(), T.int64(), T.int64() lv0 = R.match_cast(x, R.Tensor((N, C, H, W), "float32")) gv: R.Tensor("float32", ndim=4) = R.nn.conv2d(lv0, w, out_dtype="float32") R.output(gv) return gv + N = T.dynamic("N") + C = T.dynamic("C") + H = T.dynamic("H") + W = T.dynamic("W") + @I.ir_module class Expected: @R.function def main( x: R.Tensor(dtype="float32", ndim=4), w: R.Tensor(dtype="float32", ndim=4) ) -> R.Tensor(dtype="float32", ndim=4): - N = T.int64() - C = T.int64() - H = T.int64() - W = T.int64() with R.dataflow(): lv0: R.Tensor((N, C, H, W), dtype="float32") = R.match_cast( x, R.Tensor((N, C, H, W), dtype="float32") @@ -167,6 +172,11 @@ def main( def test_conv2d_matchcast_bias(): + N = T.dynamic("N") + C = T.dynamic("C") + H = T.dynamic("H") + W = T.dynamic("W") + @I.ir_module class Input: @R.function @@ -175,22 +185,22 @@ def main(x: R.Tensor("float32", ndim=4), w: R.Tensor("float32", ndim=4)) -> R.Te ): with R.dataflow(): lv0: R.Tensor("float32", ndim=4) = R.nn.conv2d(x, w, out_dtype="float32") - N, C, H, W = T.int64(), T.int64(), T.int64(), T.int64() lv1 = R.match_cast(lv0, R.Tensor((N, C, H, W), "float32")) gv = R.add(lv1, w) R.output(gv) return gv + N = T.dynamic("N") + H = T.dynamic("H") + W = T.dynamic("W") + C = T.dynamic("C") + @I.ir_module class Expected: @R.function def main( x: R.Tensor(dtype="float32", ndim=4), w: R.Tensor(dtype="float32", ndim=4) ) -> R.Tensor(dtype="float32", ndim=4): - N = T.int64() - H = T.int64() - W = T.int64() - C = T.int64() with R.dataflow(): lv: R.Tensor(dtype="float32", ndim=4) = R.permute_dims(x, axes=[0, 2, 3, 1]) lv1: R.Tensor(dtype="float32", ndim=4) = R.permute_dims(w, axes=[0, 2, 3, 1]) @@ -1909,6 +1919,12 @@ def main( def test_conv2d_symbolic_sub_indexed(): + N = T.dynamic("N") + H = T.dynamic("H") + W = T.dynamic("W") + Hw = T.dynamic("Hw") + Ww = T.dynamic("Ww") + @I.ir_module class Input: @R.function @@ -1916,8 +1932,6 @@ def main(x: R.Tensor("float32", ndim=4), w: R.Tensor("float32", ndim=4)) -> R.Te "float32", ndim=4 ): with R.dataflow(): - N, H, W = T.int64(), T.int64(), T.int64() - Hw, Ww = T.int64(), T.int64() lv0 = R.match_cast(x, R.Tensor((N, 16, H, W), "float32")) lv1 = R.match_cast(w, R.Tensor((4, 16, Hw, Ww), "float32")) gv: R.Tensor( @@ -1926,17 +1940,18 @@ def main(x: R.Tensor("float32", ndim=4), w: R.Tensor("float32", ndim=4)) -> R.Te R.output(gv) return gv + N = T.dynamic("N") + H = T.dynamic("H") + W = T.dynamic("W") + Hw = T.dynamic("Hw") + Ww = T.dynamic("Ww") + @I.ir_module class Expected: @R.function def main( x: R.Tensor(dtype="float32", ndim=4), w: R.Tensor(dtype="float32", ndim=4) ) -> R.Tensor(dtype="float32", ndim=4): - N = T.int64() - H = T.int64() - W = T.int64() - Hw = T.int64() - Ww = T.int64() with R.dataflow(): lv0: R.Tensor((N, 16, H, W), dtype="float32") = R.match_cast( x, R.Tensor((N, 16, H, W), dtype="float32") @@ -1984,6 +1999,16 @@ def main( def test_conv2d_matchcast_bias_sub_indexed(): + N = T.dynamic("N") + H = T.dynamic("H") + W = T.dynamic("W") + Hw = T.dynamic("Hw") + Ww = T.dynamic("Ww") + Nb = T.dynamic("Nb") + Cb = T.dynamic("Cb") + Hb = T.dynamic("Hb") + Wb = T.dynamic("Wb") + @I.ir_module class Input: @R.function @@ -1993,17 +2018,24 @@ def main( bias: R.Tensor("float32", ndim=4), ) -> R.Tensor(None, "float32", ndim=4): with R.dataflow(): - N, H, W = T.int64(), T.int64(), T.int64() - Hw, Ww = T.int64(), T.int64() lv0 = R.match_cast(x, R.Tensor((N, 16, H, W), "float32")) lv1 = R.match_cast(w, R.Tensor((4, 16, Hw, Ww), "float32")) lv2: R.Tensor("float32", ndim=4) = R.nn.conv2d(lv0, lv1, out_dtype="float32") - Nb, Cb, Hb, Wb = T.int64(), T.int64(), T.int64(), T.int64() lv_bias = R.match_cast(bias, R.Tensor((Nb, Cb, Hb, Wb), "float32")) gv = R.add(lv2, lv_bias) R.output(gv) return gv + N = T.dynamic("N") + H = T.dynamic("H") + W = T.dynamic("W") + Hw = T.dynamic("Hw") + Ww = T.dynamic("Ww") + Nb = T.dynamic("Nb") + Cb = T.dynamic("Cb") + Hb = T.dynamic("Hb") + Wb = T.dynamic("Wb") + @I.ir_module class Expected_NHWC4c: @R.function @@ -2012,9 +2044,6 @@ def main( w: R.Tensor(dtype="float32", ndim=4), bias: R.Tensor(dtype="float32", ndim=4), ) -> R.Tensor(dtype="float32", ndim=4): - N, H, W = T.int64(), T.int64(), T.int64() - Hw, Ww = T.int64(), T.int64() - Nb, Cb, Hb, Wb = T.int64(), T.int64(), T.int64(), T.int64() with R.dataflow(): lv0: R.Tensor((N, 16, H, W), dtype="float32") = R.match_cast( x, R.Tensor((N, 16, H, W), dtype="float32") @@ -2072,6 +2101,16 @@ def main( R.output(gv) return gv + N = T.dynamic("N") + H = T.dynamic("H") + W = T.dynamic("W") + Hw = T.dynamic("Hw") + Ww = T.dynamic("Ww") + Nb = T.dynamic("Nb") + Cb = T.dynamic("Cb") + Hb = T.dynamic("Hb") + Wb = T.dynamic("Wb") + @I.ir_module class Expected_NCHW4c: @R.function @@ -2080,9 +2119,6 @@ def main( w: R.Tensor(dtype="float32", ndim=4), bias: R.Tensor(dtype="float32", ndim=4), ) -> R.Tensor(dtype="float32", ndim=4): - N, H, W = T.int64(), T.int64(), T.int64() - Hw, Ww = T.int64(), T.int64() - Nb, Cb, Hb, Wb = T.int64(), T.int64(), T.int64(), T.int64() with R.dataflow(): lv0: R.Tensor((N, 16, H, W), dtype="float32") = R.match_cast( x, R.Tensor((N, 16, H, W), dtype="float32") @@ -2141,6 +2177,16 @@ def main( def test_conv2d_layout_incompatible_fallback(): + N = T.dynamic("N") + H = T.dynamic("H") + W = T.dynamic("W") + Hw = T.dynamic("Hw") + Ww = T.dynamic("Ww") + Nb = T.dynamic("Nb") + Cb = T.dynamic("Cb") + Hb = T.dynamic("Hb") + Wb = T.dynamic("Wb") + @I.ir_module class Input: @R.function @@ -2150,17 +2196,24 @@ def main( bias: R.Tensor("float32", ndim=4), ) -> R.Tensor(None, "float32", ndim=4): with R.dataflow(): - N, H, W = T.int64(), T.int64(), T.int64() - Hw, Ww = T.int64(), T.int64() lv0 = R.match_cast(x, R.Tensor((N, 15, H, W), "float32")) lv1 = R.match_cast(w, R.Tensor((4, 15, Hw, Ww), "float32")) lv2: R.Tensor("float32", ndim=4) = R.nn.conv2d(lv0, lv1, out_dtype="float32") - Nb, Cb, Hb, Wb = T.int64(), T.int64(), T.int64(), T.int64() lv_bias = R.match_cast(bias, R.Tensor((Nb, Cb, Hb, Wb), "float32")) gv = R.add(lv2, lv_bias) R.output(gv) return gv + N = T.dynamic("N") + H = T.dynamic("H") + W = T.dynamic("W") + Hw = T.dynamic("Hw") + Ww = T.dynamic("Ww") + Nb = T.dynamic("Nb") + Cb = T.dynamic("Cb") + Hb = T.dynamic("Hb") + Wb = T.dynamic("Wb") + @I.ir_module class Expected: @R.function @@ -2169,9 +2222,6 @@ def main( w: R.Tensor(dtype="float32", ndim=4), bias: R.Tensor(dtype="float32", ndim=4), ) -> R.Tensor(dtype="float32", ndim=4): - N, H, W = T.int64(), T.int64(), T.int64() - Hw, Ww = T.int64(), T.int64() - Nb, Cb, Hb, Wb = T.int64(), T.int64(), T.int64(), T.int64() with R.dataflow(): lv0: R.Tensor((N, 15, H, W), dtype="float32") = R.match_cast( x, R.Tensor((N, 15, H, W), dtype="float32") diff --git a/tests/python/relax/test_transform_cse.py b/tests/python/relax/test_transform_cse.py index 593b81a39751..643c5784a10b 100644 --- a/tests/python/relax/test_transform_cse.py +++ b/tests/python/relax/test_transform_cse.py @@ -493,6 +493,11 @@ def foo(x: R.Tensor((2, 3), dtype="float32"), y: R.Tensor((2, 3), dtype="float32 def test_match_cast_with_symbolic_vars(): + n = T.dynamic("n") + m = T.dynamic("m") + p = T.dynamic("p") + q = T.dynamic("q") + @I.ir_module class Before: @R.function @@ -500,32 +505,29 @@ def foo(x: R.Tensor(dtype="float32"), y: R.Tensor(dtype="float32")): with R.dataflow(): A1 = R.add(x, y) - n = T.int64() - m = T.int64() B1 = R.match_cast(A1, R.Tensor([n, m], "float32")) A2 = R.add(x, y) - p = T.int64() - q = T.int64() B2 = R.match_cast(A2, R.Tensor([p, q], "float32")) gv = R.multiply(B1, B2) R.output(gv) return gv + n = T.dynamic("n") + m = T.dynamic("m") + p = T.dynamic("p") + q = T.dynamic("q") + @I.ir_module class Expected: @R.function def foo(x: R.Tensor(dtype="float32"), y: R.Tensor(dtype="float32")): with R.dataflow(): A1 = R.add(x, y) - n = T.int64() - m = T.int64() B1 = R.match_cast(A1, R.Tensor([n, m], "float32")) A2 = A1 - p = T.int64() - q = T.int64() B2 = R.match_cast(A1, R.Tensor([p, q], "float32")) gv = R.multiply(B1, B2) diff --git a/tests/python/relax/test_transform_dead_code_elimination.py b/tests/python/relax/test_transform_dead_code_elimination.py index 603e0f099631..4dd94d15f6df 100644 --- a/tests/python/relax/test_transform_dead_code_elimination.py +++ b/tests/python/relax/test_transform_dead_code_elimination.py @@ -281,6 +281,16 @@ def main(x: R.Tensor((16, 16), "float32")) -> R.Tensor((16, 16), "float32"): def test_unused_relax_func_symbolic_shape(): # Test with relax function w/ symbolic shape. + m_tir_matmul = T.dynamic("m") + n_tir_matmul = T.dynamic("n") + k_tir_matmul = T.dynamic("k") + m_unused_func = T.dynamic("m") + n_unused_func = T.dynamic("n") + k_unused_func = T.dynamic("k") + m_main = T.dynamic("m") + k_main = T.dynamic("k") + n_main = T.dynamic("n") + @tvm.script.ir_module(check_well_formed=False) class InputModule: @Ts.prim_func @@ -289,28 +299,31 @@ def tir_matmul( y_handle: T.handle, z_handle: T.handle, ) -> None: - m = T.int64() - n = T.int64() - k = T.int64() - x = T.match_buffer(x_handle, (m, n), "float32") - y = T.match_buffer(y_handle, (n, k), "float32") - z = T.match_buffer(z_handle, (m, k), "float32") - for i, j, k in T.grid(m, k, n): + x = T.match_buffer(x_handle, (m_tir_matmul, n_tir_matmul), "float32") + y = T.match_buffer(y_handle, (n_tir_matmul, k_tir_matmul), "float32") + z = T.match_buffer(z_handle, (m_tir_matmul, k_tir_matmul), "float32") + for i, j, k_tir_matmul_index in T.grid(m_tir_matmul, k_tir_matmul, n_tir_matmul): with Ts.sblock("matmul"): - vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) + vi, vj, vk = Ts.axis.remap("SSR", [i, j, k_tir_matmul_index]) with Ts.init(): z[vi, vj] = 0.0 z[vi, vj] = z[vi, vj] + x[vi, vk] * y[vk, vj] @R.function(private=True) - def unused_func(x: R.Tensor(("m", "n"), "float32"), w: R.Tensor(("n", "k"), "float32")): + def unused_func( + x: R.Tensor((m_unused_func, n_unused_func), "float32"), + w: R.Tensor((n_unused_func, k_unused_func), "float32"), + ): gv0 = R.add(x, w) return gv0 @R.function - def main(x: R.Tensor(("m", "n"), "float32"), w: R.Tensor(("n", "k"), "float32")): - m, k = T.int64(), T.int64() - gv0 = R.call_tir(InputModule.tir_matmul, (x, w), R.Tensor((m, k), dtype="float32")) + def main( + x: R.Tensor((m_main, n_main), "float32"), w: R.Tensor((n_main, k_main), "float32") + ): + gv0 = R.call_tir( + InputModule.tir_matmul, (x, w), R.Tensor((m_main, k_main), dtype="float32") + ) return gv0 mod = InputModule diff --git a/tests/python/relax/test_transform_decompose_ops.py b/tests/python/relax/test_transform_decompose_ops.py index bf5f84fc71f8..2da0d09ff2a7 100644 --- a/tests/python/relax/test_transform_decompose_ops.py +++ b/tests/python/relax/test_transform_decompose_ops.py @@ -367,13 +367,14 @@ def main(t: R.Tensor([3], dtype="int64")): gv: R.Shape(ndim=3) = R.tensor_to_shape(t) return gv + x = T.dynamic("x") + x_1 = T.dynamic("x_1") + x_2 = T.dynamic("x_2") + @I.ir_module class Expected: @R.function def main(t: R.Tensor([3], dtype="int64")) -> R.Shape(ndim=3): - x = T.int64() - x_1 = T.int64() - x_2 = T.int64() gv: R.Shape(ndim=3) = R.call_pure_packed( "vm.builtin.tensor_to_shape", t, ty_args=(R.Shape(ndim=3),) ) diff --git a/tests/python/relax/test_transform_fold_constant.py b/tests/python/relax/test_transform_fold_constant.py index 0f47559b745f..a4c5aca644d8 100644 --- a/tests/python/relax/test_transform_fold_constant.py +++ b/tests/python/relax/test_transform_fold_constant.py @@ -180,16 +180,21 @@ def expected(c1: R.Tensor((16, 16), "float32")): def test_fold_mixed_case(): + n_addone = T.dynamic("n", "int32") + m_addone = T.dynamic("m", "int32") + n_before = T.dynamic("n") + m_before = T.dynamic("m") + n_expected = T.dynamic("n") + m_expected = T.dynamic("m") + @tvm.script.ir_module class Module: # TIR function can handle different cases. @Ts.prim_func def addone(a: T.handle, b: T.handle) -> None: - n = T.int32() - m = T.int32() - A = T.match_buffer(a, (n, m)) - B = T.match_buffer(b, (n, m)) - for i, j in T.grid(n, m): + A = T.match_buffer(a, (n_addone, m_addone)) + B = T.match_buffer(b, (n_addone, m_addone)) + for i, j in T.grid(n_addone, m_addone): with Ts.sblock("addone"): vi, vj = Ts.axis.remap("SS", [i, j]) B[vi, vj] = A[vi, vj] + T.float32(1) @@ -207,11 +212,10 @@ def sub( @R.function def before(c0: R.Tensor((16, 16), "float32"), x: R.Tensor("float32", ndim=2)): - n, m = T.int64(), T.int64() cls = Module - x0 = R.match_cast(x, R.Tensor((n, m), "float32")) + x0 = R.match_cast(x, R.Tensor((n_before, m_before), "float32")) # this line cannot be folded because n is unknown - lv0 = relax.call_tir(cls.addone, (c0,), R.Tensor((n, 16), dtype="float32")) + lv0 = relax.call_tir(cls.addone, (c0,), R.Tensor((n_before, 16), dtype="float32")) # this line can be folded lv1 = relax.call_tir(cls.addone, (c0,), R.Tensor((16, 16), dtype="float32")) # this line can be folded because all inputs are const @@ -227,11 +231,10 @@ def expected( c2: R.Tensor((16, 16), "float32"), x: R.Tensor("float32", ndim=2), ): - n, m = T.int64(), T.int64() cls = Module - x0 = R.match_cast(x, R.Tensor((n, m), "float32")) + x0 = R.match_cast(x, R.Tensor((n_expected, m_expected), "float32")) # this line cannot be folded because n is unknown - lv0 = relax.call_tir(cls.addone, (c0,), R.Tensor((n, 16), dtype="float32")) + lv0 = relax.call_tir(cls.addone, (c0,), R.Tensor((n_expected, 16), dtype="float32")) # this line can not be folded because x's shape is unknown lv3 = relax.call_tir(cls.sub, (c2, x), R.Tensor((16, 16), dtype="float32")) return (lv0, lv3) @@ -589,6 +592,8 @@ def expected(c1: R.Tensor((2048,), "float32")): def test_call_tir_with_primitive_args_not_folded(): """call_tir with symbolic primitive arguments cannot be const-evaluated.""" + m = T.dynamic("m") + @tvm.script.ir_module class Module: @Ts.prim_func(private=True) @@ -599,8 +604,7 @@ def shape_to_tensor(m: T.int64, out: T.Buffer((T.int64(1),), "int64")): out[vi] = m @R.function - def main(x: R.Tensor(("m",), "float32")): - m = T.int64() + def main(x: R.Tensor((m,), "float32")): cls = Module gv = relax.call_tir(cls.shape_to_tensor, (m,), R.Tensor((1,), "int64")) return gv diff --git a/tests/python/relax/test_transform_fuse_ops.py b/tests/python/relax/test_transform_fuse_ops.py index 7421293d572d..eda3b65c8d6f 100644 --- a/tests/python/relax/test_transform_fuse_ops.py +++ b/tests/python/relax/test_transform_fuse_ops.py @@ -1275,10 +1275,13 @@ def main(inp_0: R.Tensor((1, 784), dtype="float32"), inp_1: R.Tensor((1, 128), d def test_symbolic_shape_aware_fuse(): + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Before: @R.function - def main(x: R.Tensor(["n", "m"], "float32")): + def main(x: R.Tensor([n, m], "float32")): with R.dataflow(): lv0 = R.emit_te(topi.add, x, R.const(1, "float32")) lv1 = R.emit_te(topi.exp, lv0) @@ -1286,12 +1289,18 @@ def main(x: R.Tensor(["n", "m"], "float32")): R.output(gv) return gv + n_fused_add_exp_squeeze = T.dynamic("n") + m_fused_add_exp_squeeze = T.dynamic("m") + n_main = T.dynamic("n") + m_main = T.dynamic("m") + @I.ir_module class Expected: @R.function(private=True) def fused_add_exp_squeeze( - x: R.Tensor(["n", "m"], "float32"), p0: R.Tensor([], "float32") - ) -> R.Tensor(["n", "m"], dtype="float32"): + x: R.Tensor([n_fused_add_exp_squeeze, m_fused_add_exp_squeeze], "float32"), + p0: R.Tensor([], "float32"), + ) -> R.Tensor([n_fused_add_exp_squeeze, m_fused_add_exp_squeeze], dtype="float32"): R.func_attr({"Primitive": True}) with R.dataflow(): lv0 = R.emit_te(topi.add, x, p0) @@ -1301,7 +1310,9 @@ def fused_add_exp_squeeze( return gv @R.function - def main(x: R.Tensor(["n", "m"], "float32")) -> R.Tensor(["n", "m"], dtype="float32"): + def main(x: R.Tensor([n_main, m_main], "float32")) -> R.Tensor( + [n_main, m_main], dtype="float32" + ): cls = Expected with R.dataflow(): gv = cls.fused_add_exp_squeeze(x, R.const(1, "float32")) @@ -1312,11 +1323,12 @@ def main(x: R.Tensor(["n", "m"], "float32")) -> R.Tensor(["n", "m"], dtype="floa def test_symbolic_shape_aware_fuse_2(): + n = T.dynamic("n") + @I.ir_module class Before: @R.function - def main(s: R.Shape(["n"])): - n = T.int64() + def main(s: R.Shape([n])): with R.dataflow(): lv0 = R.emit_te(topi.full, [n, n], "float32", 0) lv1 = R.emit_te(topi.trilu, lv0, tvm.tirx.const(1, "int32"), upper=True) @@ -1324,28 +1336,40 @@ def main(s: R.Shape(["n"])): R.output(gv) return gv + n_fused_full_trilu_broadcast_to = T.dynamic("n") + n_main = T.dynamic("n") + @I.ir_module class Expected: @R.function(private=True) def fused_full_trilu_broadcast_to( - s: R.Shape(["n"]), - ) -> R.Tensor([1, 1, "n", "n"], "float32"): + s: R.Shape([n_fused_full_trilu_broadcast_to]), + ) -> R.Tensor( + [1, 1, n_fused_full_trilu_broadcast_to, n_fused_full_trilu_broadcast_to], "float32" + ): R.func_attr({"Primitive": True}) - n = T.int64() with R.dataflow(): - lv0 = R.emit_te(topi.full, [n, n], "float32", 0) + lv0 = R.emit_te( + topi.full, + [n_fused_full_trilu_broadcast_to, n_fused_full_trilu_broadcast_to], + "float32", + 0, + ) lv1 = R.emit_te(topi.trilu, lv0, tvm.tirx.const(1, "int32"), upper=True) - gv = R.emit_te(topi.broadcast_to, lv1, [1, 1, n, n]) + gv = R.emit_te( + topi.broadcast_to, + lv1, + [1, 1, n_fused_full_trilu_broadcast_to, n_fused_full_trilu_broadcast_to], + ) R.output(gv) return gv @R.function - def main(s: R.Shape(["n"])) -> R.Tensor((1, 1, "n", "n"), dtype="float32"): + def main(s: R.Shape([n_main])) -> R.Tensor((1, 1, n_main, n_main), dtype="float32"): cls = Expected - n = T.int64() with R.dataflow(): - gv: R.Tensor([1, 1, n, n], "float32") = cls.fused_full_trilu_broadcast_to( - R.shape([n]) + gv: R.Tensor([1, 1, n_main, n_main], "float32") = cls.fused_full_trilu_broadcast_to( + R.shape([n_main]) ) R.output(gv) return gv @@ -1354,6 +1378,8 @@ def main(s: R.Shape(["n"])) -> R.Tensor((1, 1, "n", "n"), dtype="float32"): def test_symbolic_prim_arg_after_tensor_arg(): + n = T.dynamic("n") + @I.ir_module class Before: @Ts.prim_func(private=True) @@ -1378,9 +1404,8 @@ def exp(x_handle: T.handle, n: T.int64, out_handle: T.handle): @R.function def main( - x: R.Tensor((1, "n"), dtype="float32"), - ) -> R.Tensor((1, "n"), dtype="float32"): - n = T.int64() + x: R.Tensor((1, n), dtype="float32"), + ) -> R.Tensor((1, n), dtype="float32"): cls = Before with R.dataflow(): lv = R.call_tir( @@ -1420,6 +1445,8 @@ def main( def test_symbolic_prim_arg_before_tensor_arg(): + n = T.dynamic("n") + @I.ir_module class Before: @Ts.prim_func(private=True) @@ -1444,9 +1471,8 @@ def exp(n: T.int64, x_handle: T.handle, out_handle: T.handle): @R.function def main( - x: R.Tensor((1, "n"), dtype="float32"), - ) -> R.Tensor((1, "n"), dtype="float32"): - n = T.int64() + x: R.Tensor((1, n), dtype="float32"), + ) -> R.Tensor((1, n), dtype="float32"): cls = Before with R.dataflow(): lv = R.call_tir( @@ -1486,6 +1512,8 @@ def main( def test_symbolic_prim_arg_reused_from_derived_tensor_shape(): + n = T.dynamic("n") + @I.ir_module class Before: @Ts.prim_func(private=True) @@ -1526,10 +1554,9 @@ def exp(x_handle: T.handle, n: T.int64, out_handle: T.handle): @R.function def main( - source: R.Tensor(("n",), dtype="float32"), - x: R.Tensor((1, "(n - 1) // 4 + 1"), dtype="float32"), - ) -> R.Tensor((1, "(n - 1) // 4 + 1"), dtype="float32"): - n = T.int64() + source: R.Tensor((n,), dtype="float32"), + x: R.Tensor((1, (n - 1) // 4 + 1), dtype="float32"), + ) -> R.Tensor((1, (n - 1) // 4 + 1), dtype="float32"): cls = Before with R.dataflow(): lv = R.call_tir( @@ -1567,6 +1594,9 @@ def main( def test_symbolic_prim_arg_not_bound_by_derived_tensor_shape(): + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Before: @Ts.prim_func(private=True) @@ -1591,11 +1621,9 @@ def exp(x_handle: T.handle, n: T.int64, m: T.int64, out_handle: T.handle): @R.function def main( - shape: R.Shape(["n", "m"]), - x: R.Tensor((1, "n + 1"), dtype="float32"), - ) -> R.Tensor((1, "n + 1"), dtype="float32"): - n = T.int64() - m = T.int64() + shape: R.Shape([n, m]), + x: R.Tensor((1, n + 1), dtype="float32"), + ) -> R.Tensor((1, n + 1), dtype="float32"): cls = Before with R.dataflow(): lv = R.call_tir( @@ -1761,6 +1789,8 @@ def collect_packed_calls(expr): def test_symbolic_prim_arg_used_only_by_output_shape(): + n = T.dynamic("n") + @I.ir_module class Before: @Ts.prim_func(private=True) @@ -1784,9 +1814,8 @@ def double(x_handle: T.handle, n: T.int64, out_handle: T.handle): @R.function def main( - source: R.Tensor(("n",), dtype="float32"), - ) -> R.Tensor(("n",), dtype="float32"): - n = T.int64() + source: R.Tensor((n,), dtype="float32"), + ) -> R.Tensor((n,), dtype="float32"): cls = Before with R.dataflow(): lv = R.call_tir( @@ -1826,11 +1855,12 @@ def main( def test_shape_expr_arg(): + n = T.dynamic("n") + @I.ir_module class Before: @R.function - def main(s: R.Shape(["n"]), kv_cache: R.Any): - n = T.int64() + def main(s: R.Shape([n]), kv_cache: R.Any): with R.dataflow(): lv0 = R.emit_te(topi.full, [n, n], "float32", 0) lv1 = R.emit_te(topi.trilu, lv0, tvm.tirx.const(1, "int32"), upper=True) @@ -1844,34 +1874,46 @@ def main(s: R.Shape(["n"]), kv_cache: R.Any): R.output(gv, lv2) return gv, lv2 + n_fused_full_trilu_broadcast_to = T.dynamic("n") + n_main = T.dynamic("n") + @I.ir_module class Expected: @R.function(private=True) def fused_full_trilu_broadcast_to( - s: R.Shape(["n"]), - ) -> R.Tensor([1, 1, "n", "n"], "float32"): + s: R.Shape([n_fused_full_trilu_broadcast_to]), + ) -> R.Tensor( + [1, 1, n_fused_full_trilu_broadcast_to, n_fused_full_trilu_broadcast_to], "float32" + ): R.func_attr({"Primitive": True}) - n = T.int64() with R.dataflow(): - lv0 = R.emit_te(topi.full, [n, n], "float32", 0) + lv0 = R.emit_te( + topi.full, + [n_fused_full_trilu_broadcast_to, n_fused_full_trilu_broadcast_to], + "float32", + 0, + ) lv1 = R.emit_te(topi.trilu, lv0, tvm.tirx.const(1, "int32"), upper=True) - gv = R.emit_te(topi.broadcast_to, lv1, [1, 1, n, n]) + gv = R.emit_te( + topi.broadcast_to, + lv1, + [1, 1, n_fused_full_trilu_broadcast_to, n_fused_full_trilu_broadcast_to], + ) R.output(gv) return gv @R.function - def main(s: R.Shape(["n"]), kv_cache: R.Any): + def main(s: R.Shape([n_main]), kv_cache: R.Any): cls = Expected - n = T.int64() with R.dataflow(): - lv: R.Tensor([1, 1, n, n], "float32") = cls.fused_full_trilu_broadcast_to( - R.shape([n]) + lv: R.Tensor([1, 1, n_main, n_main], "float32") = cls.fused_full_trilu_broadcast_to( + R.shape([n_main]) ) gv = R.call_pure_packed( "vm.builtin.attention_kv_cache_view", kv_cache, - R.shape([1 + n, 32, 128]), - ty_args=(R.Tensor((1 + n, 32, 128), dtype="float32"),), + R.shape([1 + n_main, 32, 128]), + ty_args=(R.Tensor((1 + n_main, 32, 128), dtype="float32"),), ) R.output(gv, lv) return gv, lv @@ -1880,12 +1922,13 @@ def main(s: R.Shape(["n"]), kv_cache: R.Any): def test_skipping_match_cast(): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Module: @R.function def main(A: R.Tensor((10, 20), dtype="float32")) -> R.Tensor(dtype="float32", ndim=2): - m = T.int64() - n = T.int64() with R.dataflow(): lv: R.Tensor((m, n), dtype="float32") = R.match_cast( A, R.Tensor((m, n), dtype="float32") diff --git a/tests/python/relax/test_transform_fuse_ops_by_pattern.py b/tests/python/relax/test_transform_fuse_ops_by_pattern.py index 1a06f6245c1d..a1ded6cfbc9e 100644 --- a/tests/python/relax/test_transform_fuse_ops_by_pattern.py +++ b/tests/python/relax/test_transform_fuse_ops_by_pattern.py @@ -1138,13 +1138,16 @@ def test_error_on_repeated_variable_definitions(): def test_matmul_symbolic_var(): + batch_size = T.dynamic("batch_size") + M = T.dynamic("M") + @I.ir_module class Before: @R.function def main( - x: R.Tensor(["batch_size", 1024], "float16"), + x: R.Tensor([batch_size, 1024], "float16"), w1: R.Tensor([1024, 1024], "float16"), - w2: R.Tensor([1024, "M"], "float16"), + w2: R.Tensor([1024, M], "float16"), ): with R.dataflow(): matmul1 = R.matmul(x, w1) @@ -1153,16 +1156,22 @@ def main( R.output(out) return out + batch_size_main = T.dynamic("batch_size") + M_main = T.dynamic("M") + batch_size_fused_relax_matmul_cublas = T.dynamic("batch_size") + batch_size_fused_relax_matmul1_cublas = T.dynamic("batch_size") + M_fused_relax_matmul1_cublas = T.dynamic("M") + @I.ir_module class Expected: @R.function def main( - x: R.Tensor(["batch_size", 1024], "float16"), + x: R.Tensor([batch_size_main, 1024], "float16"), w1: R.Tensor([1024, 1024], "float16"), - w2: R.Tensor([1024, "M"], "float16"), + w2: R.Tensor([1024, M_main], "float16"), ) -> R.Tuple( - R.Tensor(["batch_size", 1024], "float16"), - R.Tensor(["batch_size", "M"], "float16"), + R.Tensor([batch_size_main, 1024], "float16"), + R.Tensor([batch_size_main, M_main], "float16"), ): cls = Expected with R.dataflow(): @@ -1174,17 +1183,16 @@ def main( @R.function def fused_relax_matmul_cublas( - x: R.Tensor(["batch_size", 1024], "float16"), + x: R.Tensor([batch_size_fused_relax_matmul_cublas, 1024], "float16"), w1: R.Tensor([1024, 1024], "float16"), - ) -> R.Tensor(["batch_size", 1024], "float16"): - batch_size = T.int64() + ) -> R.Tensor([batch_size_fused_relax_matmul_cublas, 1024], "float16"): R.func_attr({"Codegen": "cublas"}) @R.function def inner_func( - x: R.Tensor([batch_size, 1024], "float16"), + x: R.Tensor([batch_size_fused_relax_matmul_cublas, 1024], "float16"), w1: R.Tensor([1024, 1024], "float16"), - ) -> R.Tensor([batch_size, 1024], "float16"): + ) -> R.Tensor([batch_size_fused_relax_matmul_cublas, 1024], "float16"): R.func_attr({"Composite": "cublas.matmul"}) with R.dataflow(): out = R.matmul(x, w1) @@ -1196,18 +1204,20 @@ def inner_func( @R.function def fused_relax_matmul1_cublas( - x: R.Tensor(["batch_size", 1024], "float16"), - w2: R.Tensor([1024, "M"], "float16"), - ) -> R.Tensor(["batch_size", "M"], "float16"): - batch_size = T.int64() - M = T.int64() + x: R.Tensor([batch_size_fused_relax_matmul1_cublas, 1024], "float16"), + w2: R.Tensor([1024, M_fused_relax_matmul1_cublas], "float16"), + ) -> R.Tensor( + [batch_size_fused_relax_matmul1_cublas, M_fused_relax_matmul1_cublas], "float16" + ): R.func_attr({"Codegen": "cublas"}) @R.function def inner_func( - x: R.Tensor([batch_size, 1024], "float16"), - w2: R.Tensor((1024, M), "float16"), - ) -> R.Tensor([batch_size, M], "float16"): + x: R.Tensor([batch_size_fused_relax_matmul1_cublas, 1024], "float16"), + w2: R.Tensor((1024, M_fused_relax_matmul1_cublas), "float16"), + ) -> R.Tensor( + [batch_size_fused_relax_matmul1_cublas, M_fused_relax_matmul1_cublas], "float16" + ): R.func_attr({"Composite": "cublas.matmul"}) with R.dataflow(): out = R.matmul(x, w2) diff --git a/tests/python/relax/test_transform_fuse_tir.py b/tests/python/relax/test_transform_fuse_tir.py index d1a1a29f6780..967bff135998 100644 --- a/tests/python/relax/test_transform_fuse_tir.py +++ b/tests/python/relax/test_transform_fuse_tir.py @@ -705,12 +705,18 @@ def main(x: R.Tensor((2, 3), "float32")): def test_symbolic_shape_aware_fuse(): + n_fused_add_exp_squeeze = T.dynamic("n") + m_fused_add_exp_squeeze = T.dynamic("m") + n_main = T.dynamic("n") + m_main = T.dynamic("m") + @I.ir_module class Before: @R.function def fused_add_exp_squeeze( - x: R.Tensor(["n", "m"], "float32"), p0: R.Tensor([], "float32") - ) -> R.Tensor(["n", "m"], dtype="float32"): + x: R.Tensor([n_fused_add_exp_squeeze, m_fused_add_exp_squeeze], "float32"), + p0: R.Tensor([], "float32"), + ) -> R.Tensor([n_fused_add_exp_squeeze, m_fused_add_exp_squeeze], dtype="float32"): R.func_attr({"Primitive": True}) with R.dataflow(): lv0 = R.emit_te(topi.add, x, p0) @@ -720,7 +726,9 @@ def fused_add_exp_squeeze( return gv @R.function - def main(x: R.Tensor(["n", "m"], "float32")) -> R.Tensor(["n", "m"], dtype="float32"): + def main(x: R.Tensor([n_main, m_main], "float32")) -> R.Tensor( + [n_main, m_main], dtype="float32" + ): cls = Before with R.dataflow(): gv = cls.fused_add_exp_squeeze(x, R.const(1, "float32")) @@ -730,10 +738,13 @@ def main(x: R.Tensor(["n", "m"], "float32")) -> R.Tensor(["n", "m"], dtype="floa def fused_add_exp_squeeze(x, p0): return topi.squeeze(topi.exp(topi.add(x, p0))) + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(["n", "m"], "float32")) -> R.Tensor(["n", "m"], dtype="float32"): + def main(x: R.Tensor([n, m], "float32")) -> R.Tensor([n, m], dtype="float32"): with R.dataflow(): gv = R.emit_te(fused_add_exp_squeeze, x, R.const(1, "float32")) R.output(gv) @@ -743,12 +754,13 @@ def main(x: R.Tensor(["n", "m"], "float32")) -> R.Tensor(["n", "m"], dtype="floa def test_fuse_of_dynamic_kernel_with_var_params_and_static_args(): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Before: @Ts.prim_func(private=True) def dynamic_tir_kernel(a: T.handle, b: T.handle): - m = T.int64() - n = T.int64() A = T.match_buffer(a, [m, n], "float32") B = T.match_buffer(b, [m, n], "float32") @@ -811,12 +823,13 @@ def test_fuse_of_dynamic_kernel_with_expression_params_and_static_args(): Here, the kernel requires arguments (m*n), and is provided """ + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Before: @Ts.prim_func(private=True) def dynamic_tir_kernel(a: T.handle, b: T.handle, c: T.handle, d: T.handle): - m = T.int64() - n = T.int64() A = T.match_buffer(a, [m * n], "float32") B = T.match_buffer(b, [m], "float32") C = T.match_buffer(c, [n], "float32") @@ -899,14 +912,17 @@ def test_symbolic_shape_aware_fuse_with_allocation(): def te_mean(x, axis): return topi.divide(topi.sum(x, axis, keepdims=True), 4096) + n_fused_mean_add_tir_sqrt_divide_multiply = T.dynamic("n") + n_main = T.dynamic("n") + @I.ir_module class Before: @R.function def fused_mean_add_tir_sqrt_divide_multiply( - x: R.Tensor((1, "n", 4096), dtype="float32"), - y: R.Tensor((1, "n", 4096), dtype="float32"), + x: R.Tensor((1, n_fused_mean_add_tir_sqrt_divide_multiply, 4096), dtype="float32"), + y: R.Tensor((1, n_fused_mean_add_tir_sqrt_divide_multiply, 4096), dtype="float32"), rms_norm_weight: R.Tensor((4096,), dtype="float32"), - ) -> R.Tensor((1, "n", 4096), dtype="float32"): + ) -> R.Tensor((1, n_fused_mean_add_tir_sqrt_divide_multiply, 4096), dtype="float32"): R.func_attr({"Primitive": True}) with R.dataflow(): lv0 = R.emit_te(te_mean, x, axis=2) @@ -919,10 +935,10 @@ def fused_mean_add_tir_sqrt_divide_multiply( @R.function def main( - x: R.Tensor((1, "n", 4096), dtype="float32"), - y: R.Tensor((1, "n", 4096), dtype="float32"), + x: R.Tensor((1, n_main, 4096), dtype="float32"), + y: R.Tensor((1, n_main, 4096), dtype="float32"), rms_norm_weight: R.Tensor((4096,), dtype="float32"), - ) -> R.Tensor((1, "n", 4096), dtype="float32"): + ) -> R.Tensor((1, n_main, 4096), dtype="float32"): cls = Before with R.dataflow(): gv = cls.fused_mean_add_tir_sqrt_divide_multiply(x, y, rms_norm_weight) @@ -936,14 +952,16 @@ def fused_mean_add_tir_sqrt_divide_multiply(x, y, rms_norm_weight): lv3 = topi.divide(y, lv2) return topi.multiply(rms_norm_weight, lv3) + n = T.dynamic("n") + @I.ir_module class Expected: @R.function def main( - x: R.Tensor((1, "n", 4096), dtype="float32"), - y: R.Tensor((1, "n", 4096), dtype="float32"), + x: R.Tensor((1, n, 4096), dtype="float32"), + y: R.Tensor((1, n, 4096), dtype="float32"), rms_norm_weight: R.Tensor((4096,), dtype="float32"), - ) -> R.Tensor((1, "n", 4096), dtype="float32"): + ) -> R.Tensor((1, n, 4096), dtype="float32"): with R.dataflow(): gv = R.emit_te(fused_mean_add_tir_sqrt_divide_multiply, x, y, rms_norm_weight) R.output(gv) @@ -953,6 +971,9 @@ def main( def test_symbolic_var_in_call_tir_args(): + m_fused = T.dynamic("m") + m_main = T.dynamic("m") + @I.ir_module class Before: @Ts.prim_func(private=True) @@ -971,16 +992,15 @@ def foo( def fused( x: R.Tensor((1, 1, 32, 128), dtype="float32"), y: R.Tensor((2048, 128), dtype="float32"), - len: R.Shape(["m"]), + len: R.Shape([m_fused]), ) -> R.Tensor((1, 1, 32, 128), dtype="float32"): R.func_attr({"Primitive": True}) - m = T.int64() cls = Before with R.dataflow(): lv1 = R.emit_te(topi.add, x, x) gv = R.call_tir( cls.foo, - [lv1, y, m], + [lv1, y, m_fused], out_ty=R.Tensor((1, 1, 32, 128), dtype="float32"), ) R.output(gv) @@ -990,7 +1010,7 @@ def fused( def main( x: R.Tensor((1, 1, 32, 128), dtype="float32"), y: R.Tensor((2048, 128), dtype="float32"), - len: R.Shape(["m"]), + len: R.Shape([m_main]), ) -> R.Tensor((1, 1, 32, 128), dtype="float32"): cls = Before with R.dataflow(): @@ -998,6 +1018,8 @@ def main( R.output(gv) return gv + m = T.dynamic("m") + @I.ir_module class Expected: @Ts.prim_func(private=True) @@ -1024,9 +1046,8 @@ def fused( def main( x: R.Tensor((1, 1, 32, 128), dtype="float32"), y: R.Tensor((2048, 128), dtype="float32"), - len: R.Shape(["m"]), + len: R.Shape([m]), ) -> R.Tensor((1, 1, 32, 128), dtype="float32"): - m = T.int64() cls = Expected with R.dataflow(): gv = R.call_tir( @@ -1161,14 +1182,17 @@ def main(inp_0: R.Tensor((1, 4, 64, 64), dtype="float32")) -> R.Tensor( def test_tir_expression_in_shape(): + n_fused_transpose_matmul = T.dynamic("n") + n_main = T.dynamic("n") + @I.ir_module class Module: @R.function def fused_transpose_matmul( x: R.Tensor((3, 4), dtype="float32"), - y: R.Tensor(("n - 1", 4), dtype="float32"), - tir_vars: R.Shape(["n"]), - ) -> R.Tensor(("n - 1", 3), dtype="float32"): + y: R.Tensor((n_fused_transpose_matmul - 1, 4), dtype="float32"), + tir_vars: R.Shape([n_fused_transpose_matmul]), + ) -> R.Tensor((n_fused_transpose_matmul - 1, 3), dtype="float32"): R.func_attr({"Primitive": True}) with R.dataflow(): lv = R.emit_te(topi.transpose, x) @@ -1179,15 +1203,17 @@ def fused_transpose_matmul( @R.function def main( x: R.Tensor((3, 4), dtype="float32"), - y: R.Tensor(("n - 1", 4), dtype="float32"), - tir_vars: R.Shape(["n"]), - ) -> R.Tensor(("n - 1", 3), dtype="float32"): + y: R.Tensor((n_main - 1, 4), dtype="float32"), + tir_vars: R.Shape([n_main]), + ) -> R.Tensor((n_main - 1, 3), dtype="float32"): cls = Module with R.dataflow(): lv = cls.fused_transpose_matmul(x, y, tir_vars) R.output(lv) return lv + n = T.dynamic("n") + @I.ir_module class Expected: @Ts.prim_func(private=True) @@ -1218,10 +1244,9 @@ def fused_transpose_matmul( @R.function def main( x: R.Tensor((3, 4), dtype="float32"), - y: R.Tensor(("n - 1", 4), dtype="float32"), - tir_vars: R.Shape(["n"]), - ) -> R.Tensor(("n - 1", 3), dtype="float32"): - n = T.int64() + y: R.Tensor((n - 1, 4), dtype="float32"), + tir_vars: R.Shape([n]), + ) -> R.Tensor((n - 1, 3), dtype="float32"): cls = Expected with R.dataflow(): lv = R.call_tir( @@ -1457,6 +1482,12 @@ def test_symbolic_var_in_buffer_shape(): typically determined from the DLTensor's known shape.) """ + sequence_length_foo = T.dynamic("sequence_length") + sequence_length_fused = T.dynamic("sequence_length") + m_fused = T.dynamic("m") + sequence_length_main = T.dynamic("sequence_length") + m_main = T.dynamic("m") + @I.ir_module class Before: @Ts.prim_func(private=True) @@ -1466,52 +1497,56 @@ def foo( m: T.int64, rotary_handle: T.handle, ): - sequence_length = T.int64() - X = T.match_buffer( - X_handle, [T.int64(1), sequence_length, T.int64(32), T.int64(128)], "float32" + X_handle, [T.int64(1), sequence_length_foo, T.int64(32), T.int64(128)], "float32" ) rotary = T.match_buffer( - rotary_handle, [T.int64(1), sequence_length, T.int64(32), T.int64(128)], "float32" + rotary_handle, + [T.int64(1), sequence_length_foo, T.int64(32), T.int64(128)], + "float32", ) - for i0, i1, i2, i3 in T.grid(T.int64(1), sequence_length, T.int64(32), T.int64(128)): + for i0, i1, i2, i3 in T.grid( + T.int64(1), sequence_length_foo, T.int64(32), T.int64(128) + ): with Ts.sblock("rotary"): v0, v1, v2, v3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) rotary[v0, v1, v2, v3] = Y[m + v1 - 1, v3] * X[v0, v1, v2, v3] @R.function def fused( - x: R.Tensor((1, "sequence_length", 32, 128), dtype="float32"), + x: R.Tensor((1, sequence_length_fused, 32, 128), dtype="float32"), y: R.Tensor((2048, 128), dtype="float32"), - len: R.Shape(["m"]), - ) -> R.Tensor((1, "sequence_length", 32, 128), dtype="float32"): + len: R.Shape([m_fused]), + ) -> R.Tensor((1, sequence_length_fused, 32, 128), dtype="float32"): R.func_attr({"Primitive": True}) - sequence_length = T.int64() - m = T.int64() cls = Before with R.dataflow(): lv1 = R.emit_te(topi.add, x, x) gv = R.call_tir( cls.foo, - [lv1, y, m], - out_ty=R.Tensor((1, sequence_length, 32, 128), dtype="float32"), + [lv1, y, m_fused], + out_ty=R.Tensor((1, sequence_length_fused, 32, 128), dtype="float32"), ) R.output(gv) return gv @R.function def main( - x: R.Tensor((1, "sequence_length", 32, 128), dtype="float32"), + x: R.Tensor((1, sequence_length_main, 32, 128), dtype="float32"), y: R.Tensor((2048, 128), dtype="float32"), - len: R.Shape(["m"]), - ) -> R.Tensor((1, "sequence_length", 32, 128), dtype="float32"): + len: R.Shape([m_main]), + ) -> R.Tensor((1, sequence_length_main, 32, 128), dtype="float32"): cls = Before with R.dataflow(): gv = cls.fused(x, y, len) R.output(gv) return gv + sequence_length_fused = T.dynamic("sequence_length") + sequence_length_main = T.dynamic("sequence_length") + m = T.dynamic("m") + @I.ir_module class Expected: @Ts.prim_func(private=True) @@ -1523,43 +1558,45 @@ def fused( ): T.func_attr({"tirx.noalias": True}) - sequence_length = T.int64() - X = T.match_buffer( - X_handle, [T.int64(1), sequence_length, T.int64(32), T.int64(128)], "float32" + X_handle, [T.int64(1), sequence_length_fused, T.int64(32), T.int64(128)], "float32" ) rotary = T.match_buffer( - rotary_handle, [T.int64(1), sequence_length, T.int64(32), T.int64(128)], "float32" + rotary_handle, + [T.int64(1), sequence_length_fused, T.int64(32), T.int64(128)], + "float32", ) - T_add = Ts.sblock_alloc_buffer((T.int64(1), sequence_length, T.int64(32), T.int64(128))) + T_add = Ts.sblock_alloc_buffer( + (T.int64(1), sequence_length_fused, T.int64(32), T.int64(128)) + ) for ax0, ax1, ax2, ax3 in T.grid( - T.int64(1), sequence_length, T.int64(32), T.int64(128) + T.int64(1), sequence_length_fused, T.int64(32), T.int64(128) ): with Ts.sblock("T_add"): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) T_add[v_ax0, v_ax1, v_ax2, v_ax3] = ( X[v_ax0, v_ax1, v_ax2, v_ax3] + X[v_ax0, v_ax1, v_ax2, v_ax3] ) - for i0, i1, i2, i3 in T.grid(T.int64(1), sequence_length, T.int64(32), T.int64(128)): + for i0, i1, i2, i3 in T.grid( + T.int64(1), sequence_length_fused, T.int64(32), T.int64(128) + ): with Ts.sblock("rotary"): v0, v1, v2, v3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) rotary[v0, v1, v2, v3] = Y[m + v1 - T.int64(1), v3] * T_add[v0, v1, v2, v3] @R.function def main( - x: R.Tensor((1, "sequence_length", 32, 128), dtype="float32"), + x: R.Tensor((1, sequence_length_main, 32, 128), dtype="float32"), y: R.Tensor((2048, 128), dtype="float32"), - len: R.Shape(["m"]), - ) -> R.Tensor((1, "sequence_length", 32, 128), dtype="float32"): - sequence_length = T.int64() - m = T.int64() + len: R.Shape([m]), + ) -> R.Tensor((1, sequence_length_main, 32, 128), dtype="float32"): cls = Expected with R.dataflow(): gv = R.call_tir( cls.fused, (x, y, m), - out_ty=R.Tensor([1, sequence_length, 32, 128], "float32"), + out_ty=R.Tensor([1, sequence_length_main, 32, 128], "float32"), ) R.output(gv) return gv @@ -1570,6 +1607,8 @@ def main( def test_symbolic_var_called_with_static_shape(): """A dynamic PrimFunc may be called with a static shape""" + num_elements = T.dynamic("num_elements") + @I.ir_module class Before: @Ts.prim_func(private=True) @@ -1577,8 +1616,6 @@ def sum_1d( X_handle: T.handle, Y: T.Buffer([T.int64(1)], "float32"), ): - num_elements = T.int64() - X = T.match_buffer(X_handle, [num_elements], "float32") for i in range(num_elements): @@ -1645,6 +1682,8 @@ def main( def test_symbolic_var_called_with_multiple_static_shapes(): """A dynamic PrimFunc may be called with different shapes each time""" + num_elements = T.dynamic("num_elements") + @I.ir_module class Before: @Ts.prim_func(private=True) @@ -1652,8 +1691,6 @@ def sum_1d( X_handle: T.handle, Sum: T.Buffer([T.int64(1)], "float32"), ): - num_elements = T.int64() - X = T.match_buffer(X_handle, [num_elements], "float32") for i in range(num_elements): diff --git a/tests/python/relax/test_transform_gradient.py b/tests/python/relax/test_transform_gradient.py index 9c7452aa2696..c5f6f4a7def9 100644 --- a/tests/python/relax/test_transform_gradient.py +++ b/tests/python/relax/test_transform_gradient.py @@ -1077,14 +1077,16 @@ def main( def test_tir_copy(): + n = T.dynamic("n") + @I.ir_module class Before: @R.function def main( - x0: R.Tensor(("n", "n"), "float32"), - x1: R.Tensor(("n", "n"), "float32"), - x2: R.Tensor(("n", "n"), "float32"), - x3: R.Tensor(("n", "n"), "float32"), + x0: R.Tensor((n, n), "float32"), + x1: R.Tensor((n, n), "float32"), + x2: R.Tensor((n, n), "float32"), + x3: R.Tensor((n, n), "float32"), ): with R.dataflow(): lv0 = R.add(x0, x1) diff --git a/tests/python/relax/test_transform_gradient_te_register.py b/tests/python/relax/test_transform_gradient_te_register.py index 31f9c446b42e..b20474d7f42c 100644 --- a/tests/python/relax/test_transform_gradient_te_register.py +++ b/tests/python/relax/test_transform_gradient_te_register.py @@ -286,17 +286,21 @@ def main(a: R.Tensor((5, 5), dtype="float32")) -> R.Tensor((), dtype="float32"): def get_expected_3(): # fmt: off + n_f_mul = T.dynamic("n") + n_f_mul_grad = T.dynamic("n") + n_main_adjoint = T.dynamic("n") + n_main = T.dynamic("n") + @I.ir_module class Expected: @Ts.prim_func(private=True) def f_mul(var_A: T.handle, var_B: T.handle, var_f_mul: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() - A = T.match_buffer(var_A, (n, n)) - B = T.match_buffer(var_B, (n, n)) - f_mul_1 = T.match_buffer(var_f_mul, (n, n)) + A = T.match_buffer(var_A, (n_f_mul, n_f_mul)) + B = T.match_buffer(var_B, (n_f_mul, n_f_mul)) + f_mul_1 = T.match_buffer(var_f_mul, (n_f_mul, n_f_mul)) # with Ts.sblock("root"): - for i0, i1 in T.grid(n, n): + for i0, i1 in T.grid(n_f_mul, n_f_mul): with Ts.sblock("f_mul"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(A[v_i0, v_i1], B[v_i0, v_i1]) @@ -306,20 +310,19 @@ def f_mul(var_A: T.handle, var_B: T.handle, var_f_mul: T.handle): @Ts.prim_func(private=True) def f_mul_grad(var_A: T.handle, var_B: T.handle, var_C: T.handle, var_f_mul_grad_1: T.handle, var_f_mul_grad_2: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() - A = T.match_buffer(var_A, (n, n)) - B = T.match_buffer(var_B, (n, n)) - C = T.match_buffer(var_C, (n, n)) - f_mul_grad_1 = T.match_buffer(var_f_mul_grad_1, (n, n)) - f_mul_grad_2 = T.match_buffer(var_f_mul_grad_2, (n, n)) + A = T.match_buffer(var_A, (n_f_mul_grad, n_f_mul_grad)) + B = T.match_buffer(var_B, (n_f_mul_grad, n_f_mul_grad)) + C = T.match_buffer(var_C, (n_f_mul_grad, n_f_mul_grad)) + f_mul_grad_1 = T.match_buffer(var_f_mul_grad_1, (n_f_mul_grad, n_f_mul_grad)) + f_mul_grad_2 = T.match_buffer(var_f_mul_grad_2, (n_f_mul_grad, n_f_mul_grad)) # with Ts.sblock("root"): - for i0, i1 in T.grid(n, n): + for i0, i1 in T.grid(n_f_mul_grad, n_f_mul_grad): with Ts.sblock("f_mul_grad_1"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(C[v_i0, v_i1], A[v_i0, v_i1]) Ts.writes(f_mul_grad_1[v_i0, v_i1]) f_mul_grad_1[v_i0, v_i1] = C[v_i0, v_i1] * A[v_i0, v_i1] - for i0, i1 in T.grid(n, n): + for i0, i1 in T.grid(n_f_mul_grad, n_f_mul_grad): with Ts.sblock("f_mul_grad_2"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(B[v_i0, v_i1], A[v_i0, v_i1]) @@ -327,28 +330,26 @@ def f_mul_grad(var_A: T.handle, var_B: T.handle, var_C: T.handle, var_f_mul_grad f_mul_grad_2[v_i0, v_i1] = B[v_i0, v_i1] * A[v_i0, v_i1] @R.function - def main_adjoint(a: R.Tensor(("n", "n"), dtype="float32"), b: R.Tensor(("n", "n"), dtype="float32")) -> R.Tuple(R.Tensor((), dtype="float32"), R.Tuple(R.Tensor(("n", "n"), dtype="float32"), R.Tensor(("n", "n"), dtype="float32"))): - n = T.int64() + def main_adjoint(a: R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32"), b: R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32")) -> R.Tuple(R.Tensor((), dtype="float32"), R.Tuple(R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32"), R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32"))): cls = Expected with R.dataflow(): - lv = R.call_tir(cls.f_mul, (a, b), out_ty=R.Tensor((n, n), dtype="float32")) + lv = R.call_tir(cls.f_mul, (a, b), out_ty=R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32")) gv: R.Tensor((), dtype="float32") = R.sum(lv, axis=None, keepdims=False) gv_adjoint: R.Tensor((), dtype="float32") = R.ones(R.shape([]), dtype="float32") - lv_adjoint: R.Tensor((n, n), dtype="float32") = R.broadcast_to(gv_adjoint, R.shape([n, n])) - lv_1 = R.call_tir(cls.f_mul_grad, (lv_adjoint, a, b), out_ty=[R.Tensor((n, n), dtype="float32"), R.Tensor((n, n), dtype="float32")]) - a_adjoint: R.Tensor((n, n), dtype="float32") = lv_1[0] - b_adjoint: R.Tensor((n, n), dtype="float32") = lv_1[1] - a_adjoint_out: R.Tensor((n, n), dtype="float32") = a_adjoint - b_adjoint_out: R.Tensor((n, n), dtype="float32") = b_adjoint + lv_adjoint: R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32") = R.broadcast_to(gv_adjoint, R.shape([n_main_adjoint, n_main_adjoint])) + lv_1 = R.call_tir(cls.f_mul_grad, (lv_adjoint, a, b), out_ty=[R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32"), R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32")]) + a_adjoint: R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32") = lv_1[0] + b_adjoint: R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32") = lv_1[1] + a_adjoint_out: R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32") = a_adjoint + b_adjoint_out: R.Tensor((n_main_adjoint, n_main_adjoint), dtype="float32") = b_adjoint R.output(gv, a_adjoint_out, b_adjoint_out) return (gv, (a_adjoint_out, b_adjoint_out)) @R.function - def main(a: R.Tensor(("n", "n"), dtype="float32"), b: R.Tensor(("n", "n"), dtype="float32")) -> R.Tensor((), dtype="float32"): - n = T.int64() + def main(a: R.Tensor((n_main, n_main), dtype="float32"), b: R.Tensor((n_main, n_main), dtype="float32")) -> R.Tensor((), dtype="float32"): cls = Expected with R.dataflow(): - lv = R.call_tir_with_grad(cls.f_mul, (a, b), out_ty=R.Tensor((n, n), dtype="float32"), te_grad_name="f_mul_grad") + lv = R.call_tir_with_grad(cls.f_mul, (a, b), out_ty=R.Tensor((n_main, n_main), dtype="float32"), te_grad_name="f_mul_grad") gv: R.Tensor((), dtype="float32") = R.sum(lv, axis=None, keepdims=False) R.output(gv) return gv diff --git a/tests/python/relax/test_transform_ipc_allreduce_rewrite.py b/tests/python/relax/test_transform_ipc_allreduce_rewrite.py index 3d9df7029ed6..155479e8dbb0 100644 --- a/tests/python/relax/test_transform_ipc_allreduce_rewrite.py +++ b/tests/python/relax/test_transform_ipc_allreduce_rewrite.py @@ -24,12 +24,13 @@ def test_ipc_allreduce_rewrite(): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Module: @R.function(pure=False) - def main(shape: R.Shape(["m", "n"])): # type: ignore - m = T.int64() - n = T.int64() + def main(shape: R.Shape([m, n])): # type: ignore alloc: R.Tensor((m, n), dtype="float16") = R.builtin.alloc_tensor( # type: ignore R.shape([m, n]), R.dtype("float16"), R.prim_value(0), R.str("global") ) @@ -42,12 +43,13 @@ def main(shape: R.Shape(["m", "n"])): # type: ignore ) return alloc1 + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Expected: @R.function(pure=False) - def main(shape: R.Shape(["m", "n"])): # type: ignore - m = T.int64() - n = T.int64() + def main(shape: R.Shape([m, n])): # type: ignore alloc: R.Tensor((m, n), dtype="float16") = R.builtin.alloc_tensor( # type: ignore R.shape([m, n]), R.dtype("float16"), R.prim_value(0), R.str("ipc_memory") ) @@ -74,12 +76,13 @@ def main(shape: R.Shape(["m", "n"])): # type: ignore def test_ipc_allreduce_spread_along_reshape(): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Module: @R.function(pure=False) - def main(shape: R.Shape(["m", "n"])): # type: ignore - m = T.int64() - n = T.int64() + def main(shape: R.Shape([m, n])): # type: ignore alloc: R.Tensor((m, n), dtype="float16") = R.builtin.alloc_tensor( # type: ignore R.shape([m, n]), R.dtype("float16"), R.prim_value(0), R.str("global") ) @@ -92,14 +95,15 @@ def main(shape: R.Shape(["m", "n"])): # type: ignore ) return alloc1 + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Expected: @R.function(pure=False) def main( - shape: R.Shape(["m", "n"]), # type: ignore - ) -> R.Tensor(("m * n",), dtype="float16"): # type: ignore - m = T.int64() - n = T.int64() + shape: R.Shape([m, n]), # type: ignore + ) -> R.Tensor((m * n,), dtype="float16"): # type: ignore alloc: R.Tensor((m, n), dtype="float16") = R.builtin.alloc_tensor( # type: ignore R.shape([m, n]), R.dtype("float16"), R.prim_value(0), R.str("ipc_memory") ) @@ -128,12 +132,13 @@ def main( def test_ipc_allreduce_skip_reducer_other_than_sum(): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Module: @R.function(pure=False) - def main(shape: R.Shape(["m", "n"])): # type: ignore - m = T.int64() - n = T.int64() + def main(shape: R.Shape([m, n])): # type: ignore alloc: R.Tensor((m, n), dtype="float16") = R.builtin.alloc_tensor( # type: ignore R.shape([m, n]), R.dtype("float16"), R.prim_value(0), R.str("global") ) diff --git a/tests/python/relax/test_transform_lambda_lift.py b/tests/python/relax/test_transform_lambda_lift.py index bb6eab79ddcc..22ed281a2038 100644 --- a/tests/python/relax/test_transform_lambda_lift.py +++ b/tests/python/relax/test_transform_lambda_lift.py @@ -440,6 +440,9 @@ def main_inner(): def test_symbolic_variable_defined_by_inner_func(): + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Before: @R.function @@ -447,13 +450,16 @@ def main(x1: R.Tensor((10, 5), "float32"), y1: R.Tensor((10, 5), "float32")) -> (10, 5), "float32" ): @R.function - def inner(x2: R.Tensor(("n", "m"), "float32"), y2: R.Tensor(("n", "m"), "float32")): + def inner(x2: R.Tensor((n, m), "float32"), y2: R.Tensor((n, m), "float32")): sum_inner = R.add(x2, y2) return sum_inner sum_main = inner(x1, y1) return sum_main + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Expected: @R.function @@ -465,8 +471,8 @@ def main(x1: R.Tensor((10, 5), "float32"), y1: R.Tensor((10, 5), "float32")) -> @R.function(private=True) def main_inner( - x2: R.Tensor(("n", "m"), "float32"), y2: R.Tensor(("n", "m"), "float32") - ) -> R.Tensor(("n", "m"), "float32"): + x2: R.Tensor((n, m), "float32"), y2: R.Tensor((n, m), "float32") + ) -> R.Tensor((n, m), "float32"): sum_inner = R.add(x2, y2) return sum_inner @@ -475,21 +481,22 @@ def main_inner( def test_runtime_symbolic_variable_defined_by_inner_func(): + n_from_param = T.dynamic("n") + n_from_match_cast = T.dynamic("n") + @I.ir_module class Before: @R.function def main(x: R.Tensor((4,), "float32")): @R.function - def from_param(y: R.Tensor(("n",), "float32")): - n = T.int64() - z = R.ones((n,), "float32") + def from_param(y: R.Tensor((n_from_param,), "float32")): + z = R.ones((n_from_param,), "float32") return z @R.function def from_match_cast(y: R.Tensor(ndim=1, dtype="float32")): - n = T.int64() - y2 = R.match_cast(y, R.Tensor((n,), "float32")) - z = R.ones((n,), "float32") + y2 = R.match_cast(y, R.Tensor((n_from_match_cast,), "float32")) + z = R.ones((n_from_match_cast,), "float32") return z a = from_param(x) @@ -503,15 +510,15 @@ def from_match_cast(y: R.Tensor(ndim=1, dtype="float32")): def test_symbolic_variable_defined_by_outer_func(): + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Before: @R.function - def main( - x1: R.Tensor(("n", "m"), "float32"), y1: R.Tensor(("n", "m"), "float32") - ) -> R.Tensor(("n", "m"), "float32"): - n = T.int64() - m = T.int64() - + def main(x1: R.Tensor((n, m), "float32"), y1: R.Tensor((n, m), "float32")) -> R.Tensor( + (n, m), "float32" + ): @R.function def inner(x2: R.Tensor((n, m), "float32"), y2: R.Tensor((n, m), "float32")): sum_inner = R.add(x2, y2) @@ -520,19 +527,25 @@ def inner(x2: R.Tensor((n, m), "float32"), y2: R.Tensor((n, m), "float32")): sum_main = inner(x1, y1) return sum_main + n_main = T.dynamic("n") + m_main = T.dynamic("m") + n_main_inner = T.dynamic("n") + m_main_inner = T.dynamic("m") + @I.ir_module class Expected: @R.function def main( - x1: R.Tensor(("n", "m"), "float32"), y1: R.Tensor(("n", "m"), "float32") - ) -> R.Tensor(("n", "m"), "float32"): + x1: R.Tensor((n_main, m_main), "float32"), y1: R.Tensor((n_main, m_main), "float32") + ) -> R.Tensor((n_main, m_main), "float32"): sum_main = Expected.main_inner(x1, y1) return sum_main @R.function(private=True) def main_inner( - x2: R.Tensor(("n", "m"), "float32"), y2: R.Tensor(("n", "m"), "float32") - ) -> R.Tensor(("n", "m"), "float32"): + x2: R.Tensor((n_main_inner, m_main_inner), "float32"), + y2: R.Tensor((n_main_inner, m_main_inner), "float32"), + ) -> R.Tensor((n_main_inner, m_main_inner), "float32"): sum_inner = R.add(x2, y2) return sum_inner diff --git a/tests/python/relax/test_transform_lazy_transform_params.py b/tests/python/relax/test_transform_lazy_transform_params.py index 2b22ace81e72..3aaf6815b26e 100644 --- a/tests/python/relax/test_transform_lazy_transform_params.py +++ b/tests/python/relax/test_transform_lazy_transform_params.py @@ -403,6 +403,8 @@ def main_transform_params(setter: R.Any) -> R.Tuple: def test_lazy_transform_params_with_symbolic_vars(): + slice_index = T.dynamic("slice_index") + @I.ir_module class Before: @R.function @@ -410,7 +412,7 @@ def main_transform_params( params: R.Tuple( R.Tensor((16, 16), dtype="float32"), R.Shape( - ["slice_index"], + [slice_index], ), ), ): @@ -418,8 +420,6 @@ def main_transform_params( R.func_attr({"relax.force_pure": True}) cls = Before - slice_index = T.int64() - param = params[0] transformed = R.call_tir( cls.slice_buffer, @@ -440,14 +440,14 @@ def slice_buffer( vi = Ts.axis.remap("S", [i]) Output[vi] = Input[slice_index, vi] + slice_index = T.dynamic("slice_index") + @I.ir_module class Expected: @R.function(pure=False) - def main_transform_params(slice_shape_expr: R.Shape(["slice_index"])): + def main_transform_params(slice_shape_expr: R.Shape([slice_index])): cls = Expected - slice_index = T.int64() - param = R.call_packed("get_item", R.prim_value(0), ty_args=(R.Any,)) gv: R.Tensor((16, 16), dtype="float32") = R.match_cast( param, R.Tensor((16, 16), dtype="float32") @@ -480,14 +480,16 @@ def slice_buffer( def test_param_shape_symbolic(): + ic_transform_layout_IOHW_to_OIHW = T.dynamic("ic", "int32") + ic_main_transform_params = T.dynamic("ic") + @I.ir_module class Before: @Ts.prim_func def transform_layout_IOHW_to_OIHW(var_w1: T.handle, var_out: T.handle): - ic = T.int32() - w1 = T.match_buffer(var_w1, (ic, 16, 3, 3), "float32") - out = T.match_buffer(var_out, (16, ic, 3, 3), "float32") - for ax0, ax1, ax2, ax3 in T.grid(16, ic, 3, 3): + w1 = T.match_buffer(var_w1, (ic_transform_layout_IOHW_to_OIHW, 16, 3, 3), "float32") + out = T.match_buffer(var_out, (16, ic_transform_layout_IOHW_to_OIHW, 3, 3), "float32") + for ax0, ax1, ax2, ax3 in T.grid(16, ic_transform_layout_IOHW_to_OIHW, 3, 3): with Ts.sblock("layout_transform"): o, i, h, w = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) Ts.reads(w1[i, o, h, w]) @@ -497,37 +499,39 @@ def transform_layout_IOHW_to_OIHW(var_w1: T.handle, var_out: T.handle): @R.function def main_transform_params( params: R.Tuple( - R.Tensor((3, "ic", 3, 3), dtype="float32"), + R.Tensor((3, ic_main_transform_params, 3, 3), dtype="float32"), R.Tensor((16, 16, 3, 3), dtype="float32"), ), ) -> R.Tuple( - R.Tensor((16, 16, 3, 3), dtype="float32"), R.Tensor(("ic", 3, 3, 3), dtype="float32") + R.Tensor((16, 16, 3, 3), dtype="float32"), + R.Tensor((ic_main_transform_params, 3, 3, 3), dtype="float32"), ): - ic = T.int64() # we expect ToNonDataflow and RemovePurityTracking to be invoked first R.func_attr({"relax.force_pure": True}) cls = Before lv: R.Tensor((16, 16, 3, 3), dtype="float32") = params[1] - lv1: R.Tensor((3, ic, 3, 3), dtype="float32") = params[0] + lv1: R.Tensor((3, ic_main_transform_params, 3, 3), dtype="float32") = params[0] lv2 = R.call_tir( cls.transform_layout_IOHW_to_OIHW, (lv1,), - out_ty=R.Tensor((ic, 3, 3, 3), dtype="float32"), + out_ty=R.Tensor((ic_main_transform_params, 3, 3, 3), dtype="float32"), ) gv: R.Tuple( R.Tensor((16, 16, 3, 3), dtype="float32"), - R.Tensor((ic, 3, 3, 3), dtype="float32"), + R.Tensor((ic_main_transform_params, 3, 3, 3), dtype="float32"), ) = (lv, lv2) return gv + ic_transform_layout_IOHW_to_OIHW = T.dynamic("ic", "int32") + ic_main_transform_params = T.dynamic("ic") + @I.ir_module class Expected: @Ts.prim_func def transform_layout_IOHW_to_OIHW(var_w1: T.handle, var_out: T.handle): - ic = T.int32() - w1 = T.match_buffer(var_w1, (ic, 16, 3, 3), "float32") - out = T.match_buffer(var_out, (16, ic, 3, 3), "float32") - for ax0, ax1, ax2, ax3 in T.grid(16, ic, 3, 3): + w1 = T.match_buffer(var_w1, (ic_transform_layout_IOHW_to_OIHW, 16, 3, 3), "float32") + out = T.match_buffer(var_out, (16, ic_transform_layout_IOHW_to_OIHW, 3, 3), "float32") + for ax0, ax1, ax2, ax3 in T.grid(16, ic_transform_layout_IOHW_to_OIHW, 3, 3): with Ts.sblock("layout_transform"): o, i, h, w = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) Ts.reads(w1[i, o, h, w]) @@ -536,7 +540,6 @@ def transform_layout_IOHW_to_OIHW(var_w1: T.handle, var_out: T.handle): @R.function(pure=False) def main_transform_params() -> R.Tuple: - ic = T.int64() cls = Expected gv: R.Any = R.call_packed("get_item", R.prim_value(1), ty_args=(R.Any,)) gv1: R.Tensor((16, 16, 3, 3), dtype="float32") = R.match_cast( @@ -546,14 +549,14 @@ def main_transform_params() -> R.Tuple: _: R.Any = R.call_packed("set_item", R.prim_value(0), lv, ty_args=(R.Any,)) _1: R.Tuple = R.vm.kill_object(lv) gv2: R.Any = R.call_packed("get_item", R.prim_value(0), ty_args=(R.Any,)) - gv3: R.Tensor((3, ic, 3, 3), dtype="float32") = R.match_cast( - gv2, R.Tensor((3, ic, 3, 3), dtype="float32") + gv3: R.Tensor((3, ic_main_transform_params, 3, 3), dtype="float32") = R.match_cast( + gv2, R.Tensor((3, ic_main_transform_params, 3, 3), dtype="float32") ) - lv1: R.Tensor((3, ic, 3, 3), dtype="float32") = gv3 + lv1: R.Tensor((3, ic_main_transform_params, 3, 3), dtype="float32") = gv3 lv2 = R.call_tir( cls.transform_layout_IOHW_to_OIHW, (lv1,), - out_ty=R.Tensor((ic, 3, 3, 3), dtype="float32"), + out_ty=R.Tensor((ic_main_transform_params, 3, 3, 3), dtype="float32"), ) _2: R.Tuple = R.vm.kill_object(lv1) _3: R.Any = R.call_packed("set_item", R.prim_value(1), lv2, ty_args=(R.Any,)) @@ -618,16 +621,18 @@ def test_output(): target = "llvm" dev = tvm.cpu() + ic = T.dynamic("ic") + @I.ir_module class TransformModule: @R.function def transform_params( params: R.Tuple( - R.Tensor((3, "ic", 3, 3), dtype="float32"), + R.Tensor((3, ic, 3, 3), dtype="float32"), R.Tensor((16, 16, 3, 3), dtype="float32"), ), ) -> R.Tuple( - R.Tensor((16, 16, 3, 3), dtype="float32"), R.Tensor(("ic", 3, 3, 3), dtype="float32") + R.Tensor((16, 16, 3, 3), dtype="float32"), R.Tensor((ic, 3, 3, 3), dtype="float32") ): R.func_attr({"relax.force_pure": True}) param0 = params[0] @@ -790,16 +795,22 @@ def transform_params(fget_param: R.Callable([T.int64, R.Any], R.Any)): def test_get_item_callback_dynamic_shape(): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Before: @R.function def transform_params( - A: R.Tensor(["m", "n"], "float32"), B: R.Tensor(["m", "n"], "float32") - ) -> R.Tuple(R.Tensor(["m", "n"], "float32"), R.Tensor(["m", "n"], "float32")): + A: R.Tensor([m, n], "float32"), B: R.Tensor([m, n], "float32") + ) -> R.Tuple(R.Tensor([m, n], "float32"), R.Tensor([m, n], "float32")): C = R.multiply(A, R.const(2, "float32")) D = R.add(C, B) return (D, B) + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Expected: @R.function @@ -807,8 +818,6 @@ def transform_params( fget_param: R.Callable([T.int64, R.Any], R.Any), ) -> R.Tuple(R.Tensor(ndim=2, dtype="float32"), R.Tensor(ndim=2, dtype="float32")): R.func_attr({"num_input": 1}) - m = T.int64() - n = T.int64() A = fget_param(R.prim_value(0), R.str("A")) A = R.match_cast(A, R.Tensor([m, n], "float32")) diff --git a/tests/python/relax/test_transform_legalize_ops_binary.py b/tests/python/relax/test_transform_legalize_ops_binary.py index 833c8129edf4..840623ea9941 100644 --- a/tests/python/relax/test_transform_legalize_ops_binary.py +++ b/tests/python/relax/test_transform_legalize_ops_binary.py @@ -123,39 +123,41 @@ def add(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_add: T.B def test_add_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Add: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "float32"): gv: R.Tensor((a, b, c, d), "float32") = R.add(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_add = T.dynamic("a") + b_add = T.dynamic("b") + c_add = T.dynamic("c") + d_add = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.add, (x, y), R.Tensor((a, b, c, d), dtype="float32")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "float32"): + gv = R.call_tir(Expected.add, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) return gv @Ts.prim_func(private=True) def add(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_add: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_add = T.match_buffer(var_T_add, [a, b, c, d], dtype="float32") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_add, d_add], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_add, b_add, c_add, T.int64(1)], dtype="float32") + T_add = T.match_buffer(var_T_add, [a_add, b_add, c_add, d_add], dtype="float32") + for i0, i1, i2, i3 in T.grid(a_add, b_add, c_add, d_add): with Ts.sblock("T_add"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[T.int64(0), ax2, ax3], rxplaceholder_1[ax0, ax1, ax2, T.int64(0)]) @@ -300,39 +302,41 @@ def divide(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_divid def test_divide_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Divide: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "float32"): gv: R.Tensor((a, b, c, d), "float32") = R.divide(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_divide = T.dynamic("a") + b_divide = T.dynamic("b") + c_divide = T.dynamic("c") + d_divide = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.divide, (x, y), R.Tensor((a, b, c, d), dtype="float32")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "float32"): + gv = R.call_tir(Expected.divide, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) return gv @Ts.prim_func(private=True) def divide(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_divide: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_divide = T.match_buffer(var_T_divide, [a, b, c, d], dtype="float32") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_divide, d_divide], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_divide, b_divide, c_divide, T.int64(1)], dtype="float32") + T_divide = T.match_buffer(var_T_divide, [a_divide, b_divide, c_divide, d_divide], dtype="float32") + for i0, i1, i2, i3 in T.grid(a_divide, b_divide, c_divide, d_divide): with Ts.sblock("T_divide"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[T.int64(0), ax2, ax3], rxplaceholder_1[ax0, ax1, ax2, T.int64(0)]) @@ -477,39 +481,41 @@ def floor_divide(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T def test_floor_divide_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class FloorDivide: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "float32"): gv: R.Tensor((a, b, c, d), "float32") = R.floor_divide(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_floor_divide = T.dynamic("a") + b_floor_divide = T.dynamic("b") + c_floor_divide = T.dynamic("c") + d_floor_divide = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.floor_divide, (x, y), R.Tensor((a, b, c, d), dtype="float32")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "float32"): + gv = R.call_tir(Expected.floor_divide, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) return gv @Ts.prim_func(private=True) def floor_divide(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_floor_divide: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_floor_divide = T.match_buffer(var_T_floor_divide, [a, b, c, d], dtype="float32") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_floor_divide, d_floor_divide], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_floor_divide, b_floor_divide, c_floor_divide, T.int64(1)], dtype="float32") + T_floor_divide = T.match_buffer(var_T_floor_divide, [a_floor_divide, b_floor_divide, c_floor_divide, d_floor_divide], dtype="float32") + for i0, i1, i2, i3 in T.grid(a_floor_divide, b_floor_divide, c_floor_divide, d_floor_divide): with Ts.sblock("T_floor_divide"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[T.int64(0), ax2, ax3], rxplaceholder_1[ax0, ax1, ax2, T.int64(0)]) @@ -592,39 +598,41 @@ def multiply(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "floa def test_multiply_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Multiply: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "float32"): gv: R.Tensor((a, b, c, d), "float32") = R.multiply(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_multiply = T.dynamic("a") + b_multiply = T.dynamic("b") + c_multiply = T.dynamic("c") + d_multiply = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.multiply, (x, y), R.Tensor((a, b, c, d), dtype="float32")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "float32"): + gv = R.call_tir(Expected.multiply, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) return gv @Ts.prim_func(private=True) def multiply(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_multiply: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_multiply = T.match_buffer(var_T_multiply, [a, b, c, d], dtype="float32") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_multiply, d_multiply], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_multiply, b_multiply, c_multiply, T.int64(1)], dtype="float32") + T_multiply = T.match_buffer(var_T_multiply, [a_multiply, b_multiply, c_multiply, d_multiply], dtype="float32") + for i0, i1, i2, i3 in T.grid(a_multiply, b_multiply, c_multiply, d_multiply): with Ts.sblock("T_multiply"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[T.int64(0), ax2, ax3], rxplaceholder_1[ax0, ax1, ax2, T.int64(0)]) @@ -709,31 +717,37 @@ def main(x: R.Tensor((1, 2, 3), dtype="float32"), y: R.Tensor((4, 3, 2, 1), dtyp def test_power_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Power: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "float32"): gv: R.Tensor((a, b, c, d), "float32") = R.power(x, y) return gv + c_power = T.dynamic("c") + d_power = T.dynamic("d") + a_power = T.dynamic("a") + b_power = T.dynamic("b") + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) def power(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_power: T.handle): T.func_attr({"tirx.noalias": True}) - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (T.int64(1), c, d)) - a = T.int64() - b = T.int64() - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (a, b, c, T.int64(1))) - T_power = T.match_buffer(var_T_power, (a, b, c, d)) + rxplaceholder = T.match_buffer(var_rxplaceholder, (T.int64(1), c_power, d_power)) + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (a_power, b_power, c_power, T.int64(1))) + T_power = T.match_buffer(var_T_power, (a_power, b_power, c_power, d_power)) # with Ts.sblock("root"): - for ax0, ax1, ax2, ax3 in T.grid(a, b, c, d): + for ax0, ax1, ax2, ax3 in T.grid(a_power, b_power, c_power, d_power): with Ts.sblock("T_power"): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) Ts.reads(rxplaceholder[T.int64(0), v_ax2, v_ax3], rxplaceholder_1[v_ax0, v_ax1, v_ax2, T.int64(0)]) @@ -741,12 +755,8 @@ def power(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_powe T_power[v_ax0, v_ax1, v_ax2, v_ax3] = T.pow(rxplaceholder[T.int64(0), v_ax2, v_ax3], rxplaceholder_1[v_ax0, v_ax1, v_ax2, T.int64(0)]) @R.function - def main(x: R.Tensor((1, "c", "d"), dtype="float32"), y: R.Tensor(("a", "b", "c", 1), dtype="float32")) -> R.Tensor(("a", "b", "c", "d"), dtype="float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.power, (x, y), out_ty=R.Tensor((a, b, c, d), dtype="float32")) + def main(x: R.Tensor((1, c_main, d_main), dtype="float32"), y: R.Tensor((a_main, b_main, c_main, 1), dtype="float32")) -> R.Tensor((a_main, b_main, c_main, d_main), dtype="float32"): + gv = R.call_tir(Expected.power, (x, y), out_ty=R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) return gv # fmt: on @@ -827,30 +837,36 @@ def main(x: R.Tensor((1, 2, 3), dtype="float32"), y: R.Tensor((4, 3, 2, 1), dtyp def test_atan2_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Atan2: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "float32"): gv: R.Tensor((a, b, c, d), "float32") = R.atan2(x, y) return gv + c_atan2 = T.dynamic("c") + d_atan2 = T.dynamic("d") + a_atan2 = T.dynamic("a") + b_atan2 = T.dynamic("b") + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) def atan2(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_atan2: T.handle): T.func_attr({"tirx.noalias": True}) - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (T.int64(1), c, d)) - a = T.int64() - b = T.int64() - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (a, b, c, T.int64(1))) - T_atan2 = T.match_buffer(var_T_atan2, (a, b, c, d)) - for ax0, ax1, ax2, ax3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, (T.int64(1), c_atan2, d_atan2)) + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (a_atan2, b_atan2, c_atan2, T.int64(1))) + T_atan2 = T.match_buffer(var_T_atan2, (a_atan2, b_atan2, c_atan2, d_atan2)) + for ax0, ax1, ax2, ax3 in T.grid(a_atan2, b_atan2, c_atan2, d_atan2): with Ts.sblock("T_atan2"): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) Ts.reads(rxplaceholder[T.int64(0), v_ax2, v_ax3], rxplaceholder_1[v_ax0, v_ax1, v_ax2, T.int64(0)]) @@ -858,12 +874,8 @@ def atan2(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_atan T_atan2[v_ax0, v_ax1, v_ax2, v_ax3] = T.atan2(rxplaceholder[T.int64(0), v_ax2, v_ax3], rxplaceholder_1[v_ax0, v_ax1, v_ax2, T.int64(0)]) @R.function - def main(x: R.Tensor((1, "c", "d"), dtype="float32"), y: R.Tensor(("a", "b", "c", 1), dtype="float32")) -> R.Tensor(("a", "b", "c", "d"), dtype="float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.atan2, (x, y), out_ty=R.Tensor((a, b, c, d), dtype="float32")) + def main(x: R.Tensor((1, c_main, d_main), dtype="float32"), y: R.Tensor((a_main, b_main, c_main, 1), dtype="float32")) -> R.Tensor((a_main, b_main, c_main, d_main), dtype="float32"): + gv = R.call_tir(Expected.atan2, (x, y), out_ty=R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) return gv # fmt: on @@ -942,39 +954,41 @@ def subtract(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "floa def test_subtract_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Subtract: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "float32"): gv: R.Tensor((a, b, c, d), "float32") = R.subtract(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_subtract = T.dynamic("a") + b_subtract = T.dynamic("b") + c_subtract = T.dynamic("c") + d_subtract = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.subtract, (x, y), R.Tensor((a, b, c, d), dtype="float32")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "float32"): + gv = R.call_tir(Expected.subtract, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) return gv @Ts.prim_func(private=True) def subtract(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_subtract: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_subtract = T.match_buffer(var_T_subtract, [a, b, c, d], dtype="float32") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_subtract, d_subtract], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_subtract, b_subtract, c_subtract, T.int64(1)], dtype="float32") + T_subtract = T.match_buffer(var_T_subtract, [a_subtract, b_subtract, c_subtract, d_subtract], dtype="float32") + for i0, i1, i2, i3 in T.grid(a_subtract, b_subtract, c_subtract, d_subtract): with Ts.sblock("T_subtract"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[T.int64(0), ax2, ax3], rxplaceholder_1[ax0, ax1, ax2, T.int64(0)]) @@ -1122,39 +1136,41 @@ def equal(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_equal: def test_equal_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Equal: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "bool"): gv: R.Tensor((a, b, c, d), "bool") = R.equal(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_equal = T.dynamic("a") + b_equal = T.dynamic("b") + c_equal = T.dynamic("c") + d_equal = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.equal, (x, y), R.Tensor((a, b, c, d), dtype="bool")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "bool"): + gv = R.call_tir(Expected.equal, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="bool")) return gv @Ts.prim_func(private=True) def equal(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_equal: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_equal = T.match_buffer(var_T_equal, [a, b, c, d], dtype="bool") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_equal, d_equal], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_equal, b_equal, c_equal, T.int64(1)], dtype="float32") + T_equal = T.match_buffer(var_T_equal, [a_equal, b_equal, c_equal, d_equal], dtype="bool") + for i0, i1, i2, i3 in T.grid(a_equal, b_equal, c_equal, d_equal): with Ts.sblock("T_equal"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[T.int64(0), ax2, ax3], rxplaceholder_1[ax0, ax1, ax2, T.int64(0)]) @@ -1299,39 +1315,41 @@ def greater(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_grea def test_greater_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Greater: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "bool"): gv: R.Tensor((a, b, c, d), "bool") = R.greater(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_greater = T.dynamic("a") + b_greater = T.dynamic("b") + c_greater = T.dynamic("c") + d_greater = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.greater, (x, y), R.Tensor((a, b, c, d), dtype="bool")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "bool"): + gv = R.call_tir(Expected.greater, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="bool")) return gv @Ts.prim_func(private=True) def greater(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_greater: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_greater = T.match_buffer(var_T_greater, [a, b, c, d], dtype="bool") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_greater, d_greater], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_greater, b_greater, c_greater, T.int64(1)], dtype="float32") + T_greater = T.match_buffer(var_T_greater, [a_greater, b_greater, c_greater, d_greater], dtype="bool") + for i0, i1, i2, i3 in T.grid(a_greater, b_greater, c_greater, d_greater): with Ts.sblock("T_greater"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder_1[ax0, ax1, ax2, T.int64(0)], rxplaceholder[T.int64(0), ax2, ax3]) @@ -1414,39 +1432,41 @@ def greater_equal(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), def test_greater_equal_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class GreaterEqual: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "bool"): gv: R.Tensor((a, b, c, d), "bool") = R.greater_equal(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_greater_equal = T.dynamic("a") + b_greater_equal = T.dynamic("b") + c_greater_equal = T.dynamic("c") + d_greater_equal = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.greater_equal, (x, y), R.Tensor((a, b, c, d), dtype="bool")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "bool"): + gv = R.call_tir(Expected.greater_equal, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="bool")) return gv @Ts.prim_func(private=True) def greater_equal(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_greater_equal: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_greater_equal = T.match_buffer(var_T_greater_equal, [a, b, c, d], dtype="bool") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_greater_equal, d_greater_equal], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_greater_equal, b_greater_equal, c_greater_equal, T.int64(1)], dtype="float32") + T_greater_equal = T.match_buffer(var_T_greater_equal, [a_greater_equal, b_greater_equal, c_greater_equal, d_greater_equal], dtype="bool") + for i0, i1, i2, i3 in T.grid(a_greater_equal, b_greater_equal, c_greater_equal, d_greater_equal): with Ts.sblock("T_greater_equal"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder_1[ax0, ax1, ax2, T.int64(0)], rxplaceholder[T.int64(0), ax2, ax3]) @@ -1529,39 +1549,41 @@ def less(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32" def test_less_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Less: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "bool"): gv: R.Tensor((a, b, c, d), "bool") = R.less(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_less = T.dynamic("a") + b_less = T.dynamic("b") + c_less = T.dynamic("c") + d_less = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.less, (x, y), R.Tensor((a, b, c, d), dtype="bool")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "bool"): + gv = R.call_tir(Expected.less, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="bool")) return gv @Ts.prim_func(private=True) def less(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_less: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_less = T.match_buffer(var_T_less, [a, b, c, d], dtype="bool") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_less, d_less], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_less, b_less, c_less, T.int64(1)], dtype="float32") + T_less = T.match_buffer(var_T_less, [a_less, b_less, c_less, d_less], dtype="bool") + for i0, i1, i2, i3 in T.grid(a_less, b_less, c_less, d_less): with Ts.sblock("T_less"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[T.int64(0), ax2, ax3], rxplaceholder_1[ax0, ax1, ax2, T.int64(0)]) @@ -1706,39 +1728,41 @@ def less_equal(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_l def test_less_equal_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class LessEqual: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "bool"): gv: R.Tensor((a, b, c, d), "bool") = R.less_equal(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_less_equal = T.dynamic("a") + b_less_equal = T.dynamic("b") + c_less_equal = T.dynamic("c") + d_less_equal = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.less_equal, (x, y), R.Tensor((a, b, c, d), dtype="bool")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "bool"): + gv = R.call_tir(Expected.less_equal, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="bool")) return gv @Ts.prim_func(private=True) def less_equal(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_less_equal: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_less_equal = T.match_buffer(var_T_less_equal, [a, b, c, d], dtype="bool") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_less_equal, d_less_equal], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_less_equal, b_less_equal, c_less_equal, T.int64(1)], dtype="float32") + T_less_equal = T.match_buffer(var_T_less_equal, [a_less_equal, b_less_equal, c_less_equal, d_less_equal], dtype="bool") + for i0, i1, i2, i3 in T.grid(a_less_equal, b_less_equal, c_less_equal, d_less_equal): with Ts.sblock("T_less_equal"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[T.int64(0), ax2, ax3], rxplaceholder_1[ax0, ax1, ax2, T.int64(0)]) @@ -1821,39 +1845,41 @@ def not_equal(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "flo def test_not_equal_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class NotEqual: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "bool"): gv: R.Tensor((a, b, c, d), "bool") = R.not_equal(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_not_equal = T.dynamic("a") + b_not_equal = T.dynamic("b") + c_not_equal = T.dynamic("c") + d_not_equal = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "bool"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.not_equal, (x, y), R.Tensor((a, b, c, d), dtype="bool")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "bool"): + gv = R.call_tir(Expected.not_equal, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="bool")) return gv @Ts.prim_func(private=True) def not_equal(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_not_equal: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_not_equal = T.match_buffer(var_T_not_equal, [a, b, c, d], dtype="bool") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_not_equal, d_not_equal], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_not_equal, b_not_equal, c_not_equal, T.int64(1)], dtype="float32") + T_not_equal = T.match_buffer(var_T_not_equal, [a_not_equal, b_not_equal, c_not_equal, d_not_equal], dtype="bool") + for i0, i1, i2, i3 in T.grid(a_not_equal, b_not_equal, c_not_equal, d_not_equal): with Ts.sblock("T_not_equal"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[T.int64(0), ax2, ax3], rxplaceholder_1[ax0, ax1, ax2, T.int64(0)]) @@ -1999,39 +2025,41 @@ def maximum(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_maxi def test_maximum_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Maximum: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "float32"): gv: R.Tensor((a, b, c, d), "float32") = R.maximum(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_maximum = T.dynamic("a") + b_maximum = T.dynamic("b") + c_maximum = T.dynamic("c") + d_maximum = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.maximum, (x, y), R.Tensor((a, b, c, d), dtype="float32")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "float32"): + gv = R.call_tir(Expected.maximum, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) return gv @Ts.prim_func(private=True) def maximum(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_maximum: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_maximum = T.match_buffer(var_T_maximum, [a, b, c, d], dtype="float32") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_maximum, d_maximum], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_maximum, b_maximum, c_maximum, T.int64(1)], dtype="float32") + T_maximum = T.match_buffer(var_T_maximum, [a_maximum, b_maximum, c_maximum, d_maximum], dtype="float32") + for i0, i1, i2, i3 in T.grid(a_maximum, b_maximum, c_maximum, d_maximum): with Ts.sblock("T_maximum"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[T.int64(0), ax2, ax3], rxplaceholder_1[ax0, ax1, ax2, T.int64(0)]) @@ -2177,39 +2205,41 @@ def minimum(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_mini def test_minimum_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Minimum: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32")) -> R.Tensor((a, b, c, d), "float32"): gv: R.Tensor((a, b, c, d), "float32") = R.minimum(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_minimum = T.dynamic("a") + b_minimum = T.dynamic("b") + c_minimum = T.dynamic("c") + d_minimum = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((1, "c", "d"), "float32"), y: R.Tensor(("a", "b", "c", 1), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.minimum, (x, y), R.Tensor((a, b, c, d), dtype="float32")) + def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_main, c_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "float32"): + gv = R.call_tir(Expected.minimum, (x, y), R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) return gv @Ts.prim_func(private=True) def minimum(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_minimum: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c, d], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b, c, T.int64(1)], dtype="float32") - T_minimum = T.match_buffer(var_T_minimum, [a, b, c, d], dtype="float32") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(1), c_minimum, d_minimum], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_minimum, b_minimum, c_minimum, T.int64(1)], dtype="float32") + T_minimum = T.match_buffer(var_T_minimum, [a_minimum, b_minimum, c_minimum, d_minimum], dtype="float32") + for i0, i1, i2, i3 in T.grid(a_minimum, b_minimum, c_minimum, d_minimum): with Ts.sblock("T_minimum"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[T.int64(0), ax2, ax3], rxplaceholder_1[ax0, ax1, ax2, T.int64(0)]) diff --git a/tests/python/relax/test_transform_legalize_ops_create_datatype.py b/tests/python/relax/test_transform_legalize_ops_create_datatype.py index 09251917d78a..b71e0912ac21 100644 --- a/tests/python/relax/test_transform_legalize_ops_create_datatype.py +++ b/tests/python/relax/test_transform_legalize_ops_create_datatype.py @@ -121,31 +121,33 @@ def full(rxplaceholder: T.Buffer((), "int32"), T_full: T.Buffer((T.int64(2), T.i def test_full_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Full: @R.function - def main(dumb_param: R.Tensor(("m", "n")), v: R.Tensor((), "int32")) -> R.Tensor(("m", "n"), "int32"): - m = T.int64() - n = T.int64() + def main(dumb_param: R.Tensor((m, n)), v: R.Tensor((), "int32")) -> R.Tensor((m, n), "int32"): gv: R.Tensor((m, n), "int32") = R.full((m, n), v, dtype="int32") return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_full = T.dynamic("m") + n_full = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(dumb_param: R.Tensor(("m", "n")), v: R.Tensor((), "int32")) -> R.Tensor(("m", "n"), "int32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.full, (v,), R.Tensor((m, n), dtype="int32")) + def main(dumb_param: R.Tensor((m_main, n_main)), v: R.Tensor((), "int32")) -> R.Tensor((m_main, n_main), "int32"): + gv = R.call_tir(Expected.full, (v,), R.Tensor((m_main, n_main), dtype="int32")) return gv @Ts.prim_func(private=True) def full(rxplaceholder: T.Buffer((), "int32"), var_T_full: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int64() - n = T.int64() - T_full = T.match_buffer(var_T_full, [m, n], dtype="int32") - for i0, i1 in T.grid(m, n): + T_full = T.match_buffer(var_T_full, [m_full, n_full], dtype="int32") + for i0, i1 in T.grid(m_full, n_full): with Ts.sblock("T_full"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[()]) @@ -252,31 +254,33 @@ def full(rxplaceholder: T.Buffer((), "float32"), T_full: T.Buffer((T.int64(2), T def test_full_like_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class FullLike: @R.function - def main(x: R.Tensor(("m", "n"), "int32"), v: R.Tensor((), "float32")) -> R.Tensor(("m", "n"), "int32"): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n), "int32"), v: R.Tensor((), "float32")) -> R.Tensor((m, n), "int32"): gv: R.Tensor((m, n), "int32") = R.full_like(x, v) return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_full = T.dynamic("m") + n_full = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "int32"), v: R.Tensor((), "float32")) -> R.Tensor(("m", "n"), "int32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.full, (v,), R.Tensor((m, n), dtype="int32")) + def main(x: R.Tensor((m_main, n_main), "int32"), v: R.Tensor((), "float32")) -> R.Tensor((m_main, n_main), "int32"): + gv = R.call_tir(Expected.full, (v,), R.Tensor((m_main, n_main), dtype="int32")) return gv @Ts.prim_func(private=True) def full(rxplaceholder: T.Buffer((), "float32"), var_T_full: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int64() - n = T.int64() - T_full = T.match_buffer(var_T_full, [m, n], dtype="int32") - for i0, i1 in T.grid(m, n): + T_full = T.match_buffer(var_T_full, [m_full, n_full], dtype="int32") + for i0, i1 in T.grid(m_full, n_full): with Ts.sblock("T_full"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[()]) @@ -321,31 +325,33 @@ def ones(T_full: T.Buffer((T.int64(2), T.int64(3)), "float32")): def test_ones_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Ones: @R.function - def main(dumb_param: R.Tensor(("m", "n"))) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() + def main(dumb_param: R.Tensor((m, n))) -> R.Tensor((m, n), "float32"): gv: R.Tensor((m, n), "float32") = R.ones((m, n), "float32") return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_ones = T.dynamic("m") + n_ones = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(dumb_param: R.Tensor(("m", "n"))) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.ones, R.tuple(), R.Tensor((m, n), dtype="float32")) + def main(dumb_param: R.Tensor((m_main, n_main))) -> R.Tensor((m_main, n_main), "float32"): + gv = R.call_tir(Expected.ones, R.tuple(), R.Tensor((m_main, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def ones(var_T_full: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int64() - n = T.int64() - T_full = T.match_buffer(var_T_full, [m, n], dtype="float32") - for i0, i1 in T.grid(m, n): + T_full = T.match_buffer(var_T_full, [m_ones, n_ones], dtype="float32") + for i0, i1 in T.grid(m_ones, n_ones): with Ts.sblock("T_full"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads() @@ -390,31 +396,33 @@ def ones(T_full: T.Buffer((T.int64(2), T.int64(3)), "int32")): def test_ones_like_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class OnesLike: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((m, n), "float32"): gv: R.Tensor((m, n), "float32") = R.ones_like(x) return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_ones = T.dynamic("m") + n_ones = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.ones, R.tuple(), R.Tensor((m, n), dtype="float32")) + def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), "float32"): + gv = R.call_tir(Expected.ones, R.tuple(), R.Tensor((m_main, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def ones(var_T_full: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int64() - n = T.int64() - T_full = T.match_buffer(var_T_full, [m, n], dtype="float32") - for i0, i1 in T.grid(m, n): + T_full = T.match_buffer(var_T_full, [m_ones, n_ones], dtype="float32") + for i0, i1 in T.grid(m_ones, n_ones): with Ts.sblock("T_full"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads() @@ -459,31 +467,33 @@ def zeros(T_full: T.Buffer((T.int64(2), T.int64(3)), "float32")): def test_zeros_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Zeros: @R.function - def main(dumb_param: R.Tensor(("m", "n"))) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() + def main(dumb_param: R.Tensor((m, n))) -> R.Tensor((m, n), "float32"): gv: R.Tensor((m, n), "float32") = R.zeros((m, n), "float32") return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_zeros = T.dynamic("m") + n_zeros = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(dumb_param: R.Tensor(("m", "n"))) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.zeros, R.tuple(), R.Tensor((m, n), dtype="float32")) + def main(dumb_param: R.Tensor((m_main, n_main))) -> R.Tensor((m_main, n_main), "float32"): + gv = R.call_tir(Expected.zeros, R.tuple(), R.Tensor((m_main, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def zeros(var_T_full: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int64() - n = T.int64() - T_full = T.match_buffer(var_T_full, [m, n], dtype="float32") - for i0, i1 in T.grid(m, n): + T_full = T.match_buffer(var_T_full, [m_zeros, n_zeros], dtype="float32") + for i0, i1 in T.grid(m_zeros, n_zeros): with Ts.sblock("T_full"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads() @@ -528,31 +538,33 @@ def zeros(T_full: T.Buffer((T.int64(2), T.int64(3)), "int32")): def test_zeros_like_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class ZerosLike: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((m, n), "float32"): gv: R.Tensor((m, n), "float32") = R.zeros_like(x) return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_zeros = T.dynamic("m") + n_zeros = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.zeros, R.tuple(), R.Tensor((m, n), dtype="float32")) + def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), "float32"): + gv = R.call_tir(Expected.zeros, R.tuple(), R.Tensor((m_main, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def zeros(var_T_full: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int64() - n = T.int64() - T_full = T.match_buffer(var_T_full, [m, n], dtype="float32") - for i0, i1 in T.grid(m, n): + T_full = T.match_buffer(var_T_full, [m_zeros, n_zeros], dtype="float32") + for i0, i1 in T.grid(m_zeros, n_zeros): with Ts.sblock("T_full"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads() @@ -587,20 +599,22 @@ def main(): def test_arange_symbolic(): # fmt: off + n = T.dynamic("n") + @tvm.script.ir_module class Arange: @R.function - def main(x: R.Tensor(["n"], "float32")): - n = T.int64() + def main(x: R.Tensor([n], "float32")): gv = R.arange(1, R.prim_value(n), 2) return gv + n = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(["n"], "float32")): + def main(x: R.Tensor([n], "float32")): cls = Expected - n = T.int64() gv = R.call_tir(cls.arange, (n,), out_ty=R.Tensor((n // 2,), dtype="int64")) return gv @@ -651,19 +665,23 @@ def shape_to_tensor(shape_to_tensor: T.Buffer((T.int64(3),), "int64")): def test_shape_to_tensor_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class ShapeToTensor: @R.function - def main(x: R.Tensor(("m", "n"), "float32")): + def main(x: R.Tensor((m, n), "float32")): gv = R.shape_to_tensor(R.shape_of(x)) return gv + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor((2,), "int64"): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((2,), "int64"): cls = Expected gv: R.Shape([m, n]) = R.shape_of(x) gv_1 = R.call_tir(cls.shape_to_tensor, (m, n), out_ty=R.Tensor((2,), dtype="int64")) @@ -684,18 +702,21 @@ def shape_to_tensor(m: T.int64, n: T.int64, shape_to_tensor: T.Buffer((T.int64(2 def test_shape_to_tensor_mixed(): # fmt: off + m = T.dynamic("m") + @tvm.script.ir_module class ShapeToTensor: @R.function - def main(x: R.Tensor(("m", 3), "float32")): + def main(x: R.Tensor((m, 3), "float32")): gv = R.shape_to_tensor(R.shape_of(x)) return gv + m = T.dynamic("m") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", 3), "float32")) -> R.Tensor((2,), "int64"): - m = T.int64() + def main(x: R.Tensor((m, 3), "float32")) -> R.Tensor((2,), "int64"): cls = Expected gv: R.Shape([m, 3]) = R.shape_of(x) gv_1 = R.call_tir(cls.shape_to_tensor, (m,), out_ty=R.Tensor((2,), dtype="int64")) @@ -768,35 +789,37 @@ def tril(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32" def test_tril_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + k = T.dynamic("k") + @tvm.script.ir_module class Tril: @R.function - def main(x: R.Tensor(("m", "n", "k"), "int8")) -> R.Tensor(("m", "n", "k"), "int8"): - m = T.int64() - n = T.int64() - k = T.int64() + def main(x: R.Tensor((m, n, k), "int8")) -> R.Tensor((m, n, k), "int8"): gv: R.Tensor((m, n, k), "int8") = R.tril(x, k=-2) return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + k_main = T.dynamic("k") + k_tril = T.dynamic("k") + m_tril = T.dynamic("m") + n_tril = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n", "k"), "int8")) -> R.Tensor(("m", "n", "k"), "int8"): - m = T.int64() - n = T.int64() - k = T.int64() - gv = R.call_tir(Expected.tril, (x,), R.Tensor((m, n, k), dtype="int8")) + def main(x: R.Tensor((m_main, n_main, k_main), "int8")) -> R.Tensor((m_main, n_main, k_main), "int8"): + gv = R.call_tir(Expected.tril, (x,), R.Tensor((m_main, n_main, k_main), dtype="int8")) return gv @Ts.prim_func(private=True) def tril(var_rxplaceholder: T.handle, var_trilu: T.handle): T.func_attr({"tirx.noalias": True}) - k = T.int64() - m = T.int64() - n = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [m, n, k], dtype="int8") - trilu = T.match_buffer(var_trilu, [m, n, k], dtype="int8") - for i0, i1, i2 in T.grid(m, n, k): + rxplaceholder = T.match_buffer(var_rxplaceholder, [m_tril, n_tril, k_tril], dtype="int8") + trilu = T.match_buffer(var_trilu, [m_tril, n_tril, k_tril], dtype="int8") + for i0, i1, i2 in T.grid(m_tril, n_tril, k_tril): with Ts.sblock("trilu"): i0_1, i1_1, i2_1 = Ts.axis.remap("SSS", [i0, i1, i2]) Ts.reads(rxplaceholder[i0_1, i1_1, i2_1]) @@ -841,35 +864,37 @@ def triu(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32" def test_triu_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + k = T.dynamic("k") + @tvm.script.ir_module class Triu: @R.function - def main(x: R.Tensor(("m", "n", "k"), "int8")) -> R.Tensor(("m", "n", "k"), "int8"): - m = T.int64() - n = T.int64() - k = T.int64() + def main(x: R.Tensor((m, n, k), "int8")) -> R.Tensor((m, n, k), "int8"): gv: R.Tensor((m, n, k), "int8") = R.triu(x, k=-2) return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + k_main = T.dynamic("k") + k_triu = T.dynamic("k") + m_triu = T.dynamic("m") + n_triu = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n", "k"), "int8")) -> R.Tensor(("m", "n", "k"), "int8"): - m = T.int64() - n = T.int64() - k = T.int64() - gv = R.call_tir(Expected.triu, (x,), R.Tensor((m, n, k), dtype="int8")) + def main(x: R.Tensor((m_main, n_main, k_main), "int8")) -> R.Tensor((m_main, n_main, k_main), "int8"): + gv = R.call_tir(Expected.triu, (x,), R.Tensor((m_main, n_main, k_main), dtype="int8")) return gv @Ts.prim_func(private=True) def triu(var_rxplaceholder: T.handle, var_trilu: T.handle): T.func_attr({"tirx.noalias": True}) - k = T.int64() - m = T.int64() - n = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [m, n, k], dtype="int8") - trilu = T.match_buffer(var_trilu, [m, n, k], dtype="int8") - for i0, i1, i2 in T.grid(m, n, k): + rxplaceholder = T.match_buffer(var_rxplaceholder, [m_triu, n_triu, k_triu], dtype="int8") + trilu = T.match_buffer(var_trilu, [m_triu, n_triu, k_triu], dtype="int8") + for i0, i1, i2 in T.grid(m_triu, n_triu, k_triu): with Ts.sblock("trilu"): i0_1, i1_1, i2_1 = Ts.axis.remap("SSS", [i0, i1, i2]) Ts.reads(rxplaceholder[i0_1, i1_1, i2_1]) @@ -938,32 +963,34 @@ def main() -> R.Tensor((), "int32"): def test_astype_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Astype: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "int32"): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((m, n), "int32"): gv: R.Tensor((m, n), "int32") = R.astype(x, "int32") return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_cast = T.dynamic("m") + n_cast = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "int32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.cast, (x,), R.Tensor((m, n), dtype="int32")) + def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), "int32"): + gv = R.call_tir(Expected.cast, (x,), R.Tensor((m_main, n_main), dtype="int32")) return gv @Ts.prim_func(private=True) def cast(var_rxplaceholder: T.handle, var_compute: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int64() - n = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [m, n], dtype="float32") - compute = T.match_buffer(var_compute, [m, n], dtype="int32") - for i0, i1 in T.grid(m, n): + rxplaceholder = T.match_buffer(var_rxplaceholder, [m_cast, n_cast], dtype="float32") + compute = T.match_buffer(var_compute, [m_cast, n_cast], dtype="int32") + for i0, i1 in T.grid(m_cast, n_cast): with Ts.sblock("compute"): i0_1, i1_1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[i0_1, i1_1]) diff --git a/tests/python/relax/test_transform_legalize_ops_distributed.py b/tests/python/relax/test_transform_legalize_ops_distributed.py index 17b2f9fd8126..6727bb6ebd80 100644 --- a/tests/python/relax/test_transform_legalize_ops_distributed.py +++ b/tests/python/relax/test_transform_legalize_ops_distributed.py @@ -35,6 +35,8 @@ def main(x: R.Tensor((10, 10), "float32")) -> R.Tensor((10, 5), "float32"): gv0 = R.dist.redistribute_replica_to_shard(x, num_workers=2, axis=1) return gv0 + worker_id = T.dynamic("worker_id") + @I.ir_module class Expected: @Ts.prim_func(private=True) @@ -50,7 +52,6 @@ def strided_slice(A: T.Buffer((T.int64(10), T.int64(10)), "float32"), worker_id: @R.function def main(x: R.Tensor((10, 10), dtype="float32")) -> R.Tensor((10, 5), dtype="float32"): - worker_id = T.int64() cls = Expected gv: R.Shape(ndim=-1) = R.call_pure_packed("runtime.disco.worker_id", ty_args=(R.Shape(ndim=-1),)) gv1: R.Shape([worker_id]) = R.match_cast(gv, R.Shape([worker_id])) diff --git a/tests/python/relax/test_transform_legalize_ops_grad.py b/tests/python/relax/test_transform_legalize_ops_grad.py index 99284243c6b1..2e4c77c6adda 100644 --- a/tests/python/relax/test_transform_legalize_ops_grad.py +++ b/tests/python/relax/test_transform_legalize_ops_grad.py @@ -197,9 +197,9 @@ def nll_loss_backward(rxplaceholder: T.Buffer((), "float32"), rxplaceholder_1: T Ts.reads(T_broadcast_to[()], all_weights[()]) Ts.writes(T_divide[()]) T_divide[()] = T_broadcast_to[()] / all_weights[()] - for i in range(T.int64(4)): + for i_index in range(T.int64(4)): with Ts.sblock("pred_grad"): - v_i = Ts.axis.spatial(T.int64(4), i) + v_i = Ts.axis.spatial(T.int64(4), i_index) Ts.reads(rxplaceholder_2[()], all_weights[()], T_divide[()]) Ts.writes(pred_grad[v_i]) pred_grad[v_i] = T.Select(v_i == rxplaceholder_2[()], all_weights[()] * T.float32(-1) * T_divide[()], T.float32(0)) @@ -319,8 +319,8 @@ def take_backward(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, va rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, (T.int64(2),), "int32", offset_factor=1) with Ts.sblock("take_backward"): T.attr(0, "pragma_scope", "seq") - for i in range(T.int64(60)): - out_buf[i // T.int64(5) // T.int64(4), i // T.int64(5) % T.int64(4), i % T.int64(5)] = T.float32(0) + for i_index in range(T.int64(60)): + out_buf[i_index // T.int64(5) // T.int64(4), i_index // T.int64(5) % T.int64(4), i_index % T.int64(5)] = T.float32(0) for parallel, serial in T.grid(T.int64(15), T.int64(2)): out_buf[(parallel // T.int64(5) * T.int64(5) * T.int64(4) + T.Cast("int64", rxplaceholder_2[serial]) * T.int64(5) + parallel % T.int64(5)) // T.int64(5) // T.int64(4), (parallel // T.int64(5) * T.int64(5) * T.int64(4) + T.Cast("int64", rxplaceholder_2[serial]) * T.int64(5) + parallel % T.int64(5)) // T.int64(5) % T.int64(4), (parallel // T.int64(5) * T.int64(5) * T.int64(4) + T.Cast("int64", rxplaceholder_2[serial]) * T.int64(5) + parallel % T.int64(5)) % T.int64(5)] = out_buf[(parallel // T.int64(5) * T.int64(5) * T.int64(4) + T.Cast("int64", rxplaceholder_2[serial]) * T.int64(5) + parallel % T.int64(5)) // T.int64(5) // T.int64(4), (parallel // T.int64(5) * T.int64(5) * T.int64(4) + T.Cast("int64", rxplaceholder_2[serial]) * T.int64(5) + parallel % T.int64(5)) // T.int64(5) % T.int64(4), (parallel // T.int64(5) * T.int64(5) * T.int64(4) + T.Cast("int64", rxplaceholder_2[serial]) * T.int64(5) + parallel % T.int64(5)) % T.int64(5)] + rxplaceholder[(parallel // T.int64(5) * T.int64(5) * T.int64(2) + serial * T.int64(5) + parallel % T.int64(5)) // T.int64(5) // T.int64(2), (parallel // T.int64(5) * T.int64(5) * T.int64(2) + serial * T.int64(5) + parallel % T.int64(5)) // T.int64(5) % T.int64(2), (parallel // T.int64(5) * T.int64(5) * T.int64(2) + serial * T.int64(5) + parallel % T.int64(5)) % T.int64(5)] @@ -337,40 +337,44 @@ def main(output_grad: R.Tensor((3, 2, 5), dtype="float32"), x: R.Tensor((3, 4, 5 def test_take_backward_symbolic(): # fmt: off + m = T.dynamic("m") + i = T.dynamic("i") + n = T.dynamic("n") + @tvm.script.ir_module class TakeBackward: @R.function - def main(output_grad: R.Tensor(("m", "i"), "float32"), x: R.Tensor(("m", "n"), "float32"), indices: R.Tensor(("i",), "int32")): - m = T.int64() - i = T.int64() + def main(output_grad: R.Tensor((m, i), "float32"), x: R.Tensor((m, n), "float32"), indices: R.Tensor((i,), "int32")): gv = R.grad.take_backward(output_grad, x, indices, axis=1) return gv + m_take_backward = T.dynamic("m") + i_take_backward = T.dynamic("i") + n_take_backward = T.dynamic("n") + m_main = T.dynamic("m") + n_main = T.dynamic("n") + i_main = T.dynamic("i") + @I.ir_module class Expected: @Ts.prim_func(private=True) def take_backward(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_rxplaceholder_2: T.handle, var_take_backward: T.handle): T.func_attr({"tirx.noalias": True}) - m, i = T.int64(), T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (m, i), offset_factor=1) - n = T.int64() - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (m, n), offset_factor=1) - rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, (i,), "int32", offset_factor=1) - out_buf = T.match_buffer(var_take_backward, (m, n)) + rxplaceholder = T.match_buffer(var_rxplaceholder, (m_take_backward, i_take_backward), offset_factor=1) + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (m_take_backward, n_take_backward), offset_factor=1) + rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, (i_take_backward,), "int32", offset_factor=1) + out_buf = T.match_buffer(var_take_backward, (m_take_backward, n_take_backward)) with Ts.sblock("take_backward"): T.attr(0, "pragma_scope", "seq") - for i_1 in range(m * n): - out_buf[i_1 // n, i_1 % n] = T.float32(0) - for parallel, serial in T.grid(m, i): - out_buf[(parallel * n + T.Cast("int64", rxplaceholder_2[serial])) // n, (parallel * n + T.Cast("int64", rxplaceholder_2[serial])) % n] = out_buf[(parallel * n + T.Cast("int64", rxplaceholder_2[serial])) // n, (parallel * n + T.Cast("int64", rxplaceholder_2[serial])) % n] + rxplaceholder[(parallel * i + serial) // i, (parallel * i + serial) % i] + for i_1 in range(m_take_backward * n_take_backward): + out_buf[i_1 // n_take_backward, i_1 % n_take_backward] = T.float32(0) + for parallel, serial in T.grid(m_take_backward, i_take_backward): + out_buf[(parallel * n_take_backward + T.Cast("int64", rxplaceholder_2[serial])) // n_take_backward, (parallel * n_take_backward + T.Cast("int64", rxplaceholder_2[serial])) % n_take_backward] = out_buf[(parallel * n_take_backward + T.Cast("int64", rxplaceholder_2[serial])) // n_take_backward, (parallel * n_take_backward + T.Cast("int64", rxplaceholder_2[serial])) % n_take_backward] + rxplaceholder[(parallel * i_take_backward + serial) // i_take_backward, (parallel * i_take_backward + serial) % i_take_backward] @R.function - def main(output_grad: R.Tensor(("m", "i"), dtype="float32"), x: R.Tensor(("m", "n"), dtype="float32"), indices: R.Tensor(("i",), dtype="int32")) -> R.Tensor(("m", "n"), dtype="float32"): - m = T.int64() - n = T.int64() - i = T.int64() + def main(output_grad: R.Tensor((m_main, i_main), dtype="float32"), x: R.Tensor((m_main, n_main), dtype="float32"), indices: R.Tensor((i_main,), dtype="int32")) -> R.Tensor((m_main, n_main), dtype="float32"): cls = Expected - gv = R.call_tir(cls.take_backward, (output_grad, x, indices), out_ty=R.Tensor((m, n), dtype="float32")) + gv = R.call_tir(cls.take_backward, (output_grad, x, indices), out_ty=R.Tensor((m_main, n_main), dtype="float32")) return gv # fmt: on diff --git a/tests/python/relax/test_transform_legalize_ops_image.py b/tests/python/relax/test_transform_legalize_ops_image.py index f29bef4e6cb3..c4b802573989 100644 --- a/tests/python/relax/test_transform_legalize_ops_image.py +++ b/tests/python/relax/test_transform_legalize_ops_image.py @@ -58,45 +58,51 @@ def resize2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(8), T.int64(8), T.int6 def test_image_resize2d_symbolic(): # fmt: off + n = T.dynamic("n") + c = T.dynamic("c") + oh = T.dynamic("oh") + ow = T.dynamic("ow") + h = T.dynamic("h") + w = T.dynamic("w") + @tvm.script.ir_module class Resize2D: @R.function - def main(dumb_param: R.Tensor(("oh", "ow")), x: R.Tensor(("n", "c", "h", "w", 16), "float32")) -> R.Tensor(("n", "c", "oh", "ow", 16), "float32"): - n = T.int64() - c = T.int64() - oh = T.int64() - ow = T.int64() + def main(dumb_param: R.Tensor((oh, ow)), x: R.Tensor((n, c, h, w, 16), "float32")) -> R.Tensor((n, c, oh, ow, 16), "float32"): gv: R.Tensor((n, c, oh, ow, 16), "float32") = R.image.resize2d(x, size=(oh, ow), layout="NCHW16c", method="nearest_neighbor", coordinate_transformation_mode="asymmetric") return gv + n_main = T.dynamic("n") + c_main = T.dynamic("c") + oh_main = T.dynamic("oh") + ow_main = T.dynamic("ow") + h_main = T.dynamic("h") + w_main = T.dynamic("w") + c_resize2d = T.dynamic("c") + h_resize2d = T.dynamic("h") + n_resize2d = T.dynamic("n") + oh_resize2d = T.dynamic("oh") + ow_resize2d = T.dynamic("ow") + w_resize2d = T.dynamic("w") + @tvm.script.ir_module class Expected: @R.function - def main(dumb_param: R.Tensor(("oh", "ow")), x: R.Tensor(("n", "c", "h", "w", 16), "float32")) -> R.Tensor(("n", "c", "oh", "ow", 16), "float32"): - n = T.int64() - c = T.int64() - oh = T.int64() - ow = T.int64() - gv = R.call_tir(Expected.resize2d, (x,), R.Tensor((n, c, oh, ow, 16), dtype="float32")) + def main(dumb_param: R.Tensor((oh_main, ow_main)), x: R.Tensor((n_main, c_main, h_main, w_main, 16), "float32")) -> R.Tensor((n_main, c_main, oh_main, ow_main, 16), "float32"): + gv = R.call_tir(Expected.resize2d, (x,), R.Tensor((n_main, c_main, oh_main, ow_main, 16), dtype="float32")) return gv @Ts.prim_func(private=True) def resize2d(var_rxplaceholder: T.handle, var_resize: T.handle): T.func_attr({"tirx.noalias": True}) - c = T.int64() - h = T.int64() - n = T.int64() - oh = T.int64() - ow = T.int64() - w = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [n, c, h, w, T.int64(16)], dtype="float32") - resize = T.match_buffer(var_resize, [n, c, oh, ow, T.int64(16)], dtype="float32") - for i0, i1, i2, i3, i4 in T.grid(n, c, oh, ow, T.int64(16)): + rxplaceholder = T.match_buffer(var_rxplaceholder, [n_resize2d, c_resize2d, h_resize2d, w_resize2d, T.int64(16)], dtype="float32") + resize = T.match_buffer(var_resize, [n_resize2d, c_resize2d, oh_resize2d, ow_resize2d, T.int64(16)], dtype="float32") + for i0, i1, i2, i3, i4 in T.grid(n_resize2d, c_resize2d, oh_resize2d, ow_resize2d, T.int64(16)): with Ts.sblock("resize"): i0_1, i1_1, i2_1, i3_1, i4_1 = Ts.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) - Ts.reads(rxplaceholder[i0_1, i1_1, T.int64(0) : T.max(h, T.int64(1)), T.int64(0) : T.max(w, T.int64(1)), i4_1]) + Ts.reads(rxplaceholder[i0_1, i1_1, T.int64(0) : T.max(h_resize2d, T.int64(1)), T.int64(0) : T.max(w_resize2d, T.int64(1)), i4_1]) Ts.writes(resize[i0_1, i1_1, i2_1, i3_1, i4_1]) - resize[i0_1, i1_1, i2_1, i3_1, i4_1] = rxplaceholder[i0_1, i1_1, T.max(T.min(T.Cast("int64", T.round(T.Cast("float32", h) / T.Cast("float32", oh) * T.Cast("float32", i2_1), dtype="float32")), h - T.int64(1)), T.int64(0)), T.max(T.min(T.Cast("int64", T.round(T.Cast("float32", w) / T.Cast("float32", ow) * T.Cast("float32", i3_1), dtype="float32")), w - T.int64(1)), T.int64(0)), i4_1] + resize[i0_1, i1_1, i2_1, i3_1, i4_1] = rxplaceholder[i0_1, i1_1, T.max(T.min(T.Cast("int64", T.round(T.Cast("float32", h_resize2d) / T.Cast("float32", oh_resize2d) * T.Cast("float32", i2_1), dtype="float32")), h_resize2d - T.int64(1)), T.int64(0)), T.max(T.min(T.Cast("int64", T.round(T.Cast("float32", w_resize2d) / T.Cast("float32", ow_resize2d) * T.Cast("float32", i3_1), dtype="float32")), w_resize2d - T.int64(1)), T.int64(0)), i4_1] # fmt: on mod = LegalizeOps()(Resize2D) @@ -127,9 +133,9 @@ def affine_grid(var_theta: T.handle, var_compute: T.handle): with Ts.sblock("root"): Ts.reads() Ts.writes() - for n, dim, i0, i1 in T.grid(T.int64(2), T.int64(2), T.int64(16), T.int64(16)): + for n_index, dim, i0, i1 in T.grid(T.int64(2), T.int64(2), T.int64(16), T.int64(16)): with Ts.sblock("compute"): - v_n, v_dim, v_i0, v_i1 = Ts.axis.remap("SSSS", [n, dim, i0, i1]) + v_n, v_dim, v_i0, v_i1 = Ts.axis.remap("SSSS", [n_index, dim, i0, i1]) Ts.reads(theta[v_n, v_dim, T.int64(0):T.int64(3)]) Ts.writes(compute[v_n, v_dim, v_i0, v_i1]) compute[v_n, v_dim, v_i0, v_i1] = theta[v_n, v_dim, T.int64(2)] + theta[v_n, v_dim, T.int64(1)] * (T.float32(-1.0) + T.Cast("float32", v_i0) * T.float32(0.13333332666666667)) + theta[v_n, v_dim, T.int64(0)] * (T.float32(-1.0) + T.Cast("float32", v_i1) * T.float32(0.13333332666666667)) diff --git a/tests/python/relax/test_transform_legalize_ops_index_linear_algebra.py b/tests/python/relax/test_transform_legalize_ops_index_linear_algebra.py index 30d229e3214e..3779c857b556 100644 --- a/tests/python/relax/test_transform_legalize_ops_index_linear_algebra.py +++ b/tests/python/relax/test_transform_legalize_ops_index_linear_algebra.py @@ -123,34 +123,38 @@ def take(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32" def test_take_symbolic(): # fmt: off + m = T.dynamic("m") + i = T.dynamic("i") + n = T.dynamic("n") + @tvm.script.ir_module class Take: @R.function - def main(x: R.Tensor(("m", "n"), "float32"), indices: R.Tensor(("i",), "int64")) -> R.Tensor(("m", "i"), "float32"): - m = T.int64() - i = T.int64() + def main(x: R.Tensor((m, n), "float32"), indices: R.Tensor((i,), "int64")) -> R.Tensor((m, i), "float32"): gv: R.Tensor((m, i), "float32") = R.take(x, indices, axis=1) return gv + m_main = T.dynamic("m") + i_main = T.dynamic("i") + n_main = T.dynamic("n") + i_take = T.dynamic("i") + m_take = T.dynamic("m") + n_take = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "float32"), indices: R.Tensor(("i",), "int64")) -> R.Tensor(("m", "i"), "float32"): - m = T.int64() - i = T.int64() - gv = R.call_tir(Expected.take, (x, indices), R.Tensor((m, i), dtype="float32")) + def main(x: R.Tensor((m_main, n_main), "float32"), indices: R.Tensor((i_main,), "int64")) -> R.Tensor((m_main, i_main), "float32"): + gv = R.call_tir(Expected.take, (x, indices), R.Tensor((m_main, i_main), dtype="float32")) return gv @Ts.prim_func(private=True) def take(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_take: T.handle): T.func_attr({"tirx.noalias": True}) - i = T.int64() - m = T.int64() - n = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [m, n], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [i], dtype="int64") - T_take = T.match_buffer(var_T_take, [m, i], dtype="float32") - for i0, i1 in T.grid(m, i): + rxplaceholder = T.match_buffer(var_rxplaceholder, [m_take, n_take], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [i_take], dtype="int64") + T_take = T.match_buffer(var_T_take, [m_take, i_take], dtype="float32") + for i0, i1 in T.grid(m_take, i_take): with Ts.sblock("T_take"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[ax0, rxplaceholder_1[ax1]], rxplaceholder_1[ax1]) @@ -164,33 +168,36 @@ def take(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_T_take: def test_take_symbolic_prim_value(): # fmt: off + n = T.dynamic("n") + @tvm.script.ir_module class Take: @R.function - def main(x: R.Tensor((2, "n", 4), "float32")) -> R.Tensor((2, 4), "float32"): - n = T.int64() + def main(x: R.Tensor((2, n, 4), "float32")) -> R.Tensor((2, 4), "float32"): gv: R.Tensor((2, 4), "float32") = R.take(x, R.prim_value(n-1), axis=1) return gv + n_main = T.dynamic("n") + n_take = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((2, "n", 4), "float32")) -> R.Tensor((2, 4), "float32"): + def main(x: R.Tensor((2, n_main, 4), "float32")) -> R.Tensor((2, 4), "float32"): gv = R.call_tir(Expected.take, (x,), R.Tensor((2, 4), dtype="float32")) return gv @Ts.prim_func(private=True) def take(x_handle: T.handle, T_take: T.Buffer((T.int64(2), T.int64(4)), "float32")): - n = T.int64() - rxplaceholder = T.match_buffer(x_handle, (T.int64(2), n, T.int64(4)), "float32") + rxplaceholder = T.match_buffer(x_handle, (T.int64(2), n_take, T.int64(4)), "float32") T.func_attr({"tirx.noalias": True}) for i0, i2 in T.grid(T.int64(2), T.int64(4)): with Ts.sblock("T_take"): ax0, ax2 = Ts.axis.remap("SS", [i0, i2]) - Ts.reads(rxplaceholder[ax0, n-1, ax2]) + Ts.reads(rxplaceholder[ax0, n_take-1, ax2]) Ts.writes(T_take[ax0, ax2]) - T_take[ax0, ax2] = rxplaceholder[ax0, n-1, ax2] + T_take[ax0, ax2] = rxplaceholder[ax0, n_take-1, ax2] # fmt: on mod = LegalizeOps()(Take) @@ -293,24 +300,30 @@ def strided_slice(rxplaceholder: T.Buffer((T.int64(8), T.int64(9), T.int64(10)), def test_strided_slice_symbolic_sliced_axis(): # fmt: off + n = T.dynamic("n") + m = T.dynamic("m") + @tvm.script.ir_module class StridedSlice: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor((2, "n"), "float32"): - n = T.int64() + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((2, n), "float32"): gv: R.Tensor((3, n), "float32") = R.strided_slice(x, axes=[0], begin=[1], end=[8], strides=[3], assume_inbound=True) return gv + m_strided_slice = T.dynamic("m") + n_strided_slice = T.dynamic("n") + n_main = T.dynamic("n") + m_main = T.dynamic("m") + @I.ir_module class Expected: @Ts.prim_func(private=True) def strided_slice(var_A: T.handle, var_T_dynamic_strided_slice_with_axes: T.handle): T.func_attr({"tirx.noalias": True}) - m, n = T.int64(), T.int64() - A = T.match_buffer(var_A, (m, n)) - T_dynamic_strided_slice_with_axes = T.match_buffer(var_T_dynamic_strided_slice_with_axes, (T.int64(3), n)) + A = T.match_buffer(var_A, (m_strided_slice, n_strided_slice)) + T_dynamic_strided_slice_with_axes = T.match_buffer(var_T_dynamic_strided_slice_with_axes, (T.int64(3), n_strided_slice)) # with Ts.sblock("root"): - for ax0, ax1 in T.grid(T.int64(3), n): + for ax0, ax1 in T.grid(T.int64(3), n_strided_slice): with Ts.sblock("T_dynamic_strided_slice_with_axes"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(A[v_ax0 * T.int64(3) + T.int64(1), v_ax1]) @@ -318,11 +331,9 @@ def strided_slice(var_A: T.handle, var_T_dynamic_strided_slice_with_axes: T.hand T_dynamic_strided_slice_with_axes[v_ax0, v_ax1] = A[v_ax0 * T.int64(3) + T.int64(1), v_ax1] @R.function - def main(x: R.Tensor(("m", "n"), dtype="float32")) -> R.Tensor((3, "n"), dtype="float32"): - n = T.int64() - m = T.int64() + def main(x: R.Tensor((m_main, n_main), dtype="float32")) -> R.Tensor((3, n_main), dtype="float32"): cls = Expected - gv = R.call_tir(cls.strided_slice, (x,), out_ty=R.Tensor((3, n), dtype="float32")) + gv = R.call_tir(cls.strided_slice, (x,), out_ty=R.Tensor((3, n_main), dtype="float32")) return gv # fmt: on @@ -332,29 +343,31 @@ def main(x: R.Tensor(("m", "n"), dtype="float32")) -> R.Tensor((3, "n"), dtype=" def test_strided_slice_symbolic(): # fmt: off + n = T.dynamic("n") + @tvm.script.ir_module class StridedSlice: @R.function - def main(x: R.Tensor((10, "n"), "float32")) -> R.Tensor((3, "n"), "float32"): - n = T.int64() + def main(x: R.Tensor((10, n), "float32")) -> R.Tensor((3, n), "float32"): gv: R.Tensor((3, n), "float32") = R.strided_slice(x, axes=[0], begin=[1], end=[8], strides=[3]) return gv + n_main = T.dynamic("n") + n_strided_slice = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((10, "n"), dtype="float32")) -> R.Tensor((3, "n"), dtype="float32"): - n = T.int64() - gv = R.call_tir(Expected.strided_slice, (x,), R.Tensor((3, n), dtype="float32")) + def main(x: R.Tensor((10, n_main), dtype="float32")) -> R.Tensor((3, n_main), dtype="float32"): + gv = R.call_tir(Expected.strided_slice, (x,), R.Tensor((3, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def strided_slice(var_rxplaceholder: T.handle, var_T_strided_slice_with_axes: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(10), n], dtype="float32") - T_strided_slice_with_axes = T.match_buffer(var_T_strided_slice_with_axes, [T.int64(3), n], dtype="float32") - for i0, i1 in T.grid(T.int64(3), n): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(10), n_strided_slice], dtype="float32") + T_strided_slice_with_axes = T.match_buffer(var_T_strided_slice_with_axes, [T.int64(3), n_strided_slice], dtype="float32") + for i0, i1 in T.grid(T.int64(3), n_strided_slice): with Ts.sblock("T_strided_slice_with_axes"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[ax0 * T.int64(3) + T.int64(1), ax1]) @@ -368,29 +381,31 @@ def strided_slice(var_rxplaceholder: T.handle, var_T_strided_slice_with_axes: T. def test_strided_slice_symbolic_bound(): # fmt: off + n = T.dynamic("n") + @tvm.script.ir_module class StridedSlice: @R.function - def main(x: R.Tensor((10, "n"), "float32")) -> R.Tensor((3, "n"), "float32"): - n = T.int64() + def main(x: R.Tensor((10, n), "float32")) -> R.Tensor((3, n), "float32"): gv: R.Tensor((3, n), "float32") = R.strided_slice(x, axes=[0, 1], begin=[1, 0], end=[8, n], strides=[3, 1]) return gv + n_main = T.dynamic("n") + n_strided_slice = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((10, "n"), dtype="float32")) -> R.Tensor((3, "n"), dtype="float32"): - n = T.int64() - gv = R.call_tir(Expected.strided_slice, (x,), R.Tensor((3, n), dtype="float32")) + def main(x: R.Tensor((10, n_main), dtype="float32")) -> R.Tensor((3, n_main), dtype="float32"): + gv = R.call_tir(Expected.strided_slice, (x,), R.Tensor((3, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def strided_slice(var_rxplaceholder: T.handle, var_T_strided_slice_with_axes: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(10), n], dtype="float32") - T_strided_slice_with_axes = T.match_buffer(var_T_strided_slice_with_axes, [T.int64(3), n], dtype="float32") - for i0, i1 in T.grid(T.int64(3), n): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(10), n_strided_slice], dtype="float32") + T_strided_slice_with_axes = T.match_buffer(var_T_strided_slice_with_axes, [T.int64(3), n_strided_slice], dtype="float32") + for i0, i1 in T.grid(T.int64(3), n_strided_slice): with Ts.sblock("T_strided_slice_with_axes"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[ax0 * T.int64(3) + T.int64(1), ax1]) @@ -400,29 +415,31 @@ def strided_slice(var_rxplaceholder: T.handle, var_T_strided_slice_with_axes: T. def test_strided_slice_non_unit_stride(): # fmt: off + n = T.dynamic("n") + @tvm.script.ir_module class StridedSlice: @R.function - def main(x: R.Tensor((10, "n"), "float32")) -> R.Tensor((3, "n"), "float32"): - n = T.int64() + def main(x: R.Tensor((10, n), "float32")) -> R.Tensor((3, n), "float32"): gv: R.Tensor((3, n), "float32") = R.strided_slice(x, axes=[0, 1], begin=[1, 0], end=[8, n], strides=[3, 1]) return gv + n_main = T.dynamic("n") + n_strided_slice = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor((10, "n"), dtype="float32")) -> R.Tensor((3, "n"), dtype="float32"): - n = T.int64() - gv = R.call_tir(Expected.strided_slice, (x,), R.Tensor((3, n), dtype="float32")) + def main(x: R.Tensor((10, n_main), dtype="float32")) -> R.Tensor((3, n_main), dtype="float32"): + gv = R.call_tir(Expected.strided_slice, (x,), R.Tensor((3, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def strided_slice(var_rxplaceholder: T.handle, var_T_strided_slice_with_axes: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(10), n], dtype="float32") - T_strided_slice_with_axes = T.match_buffer(var_T_strided_slice_with_axes, [T.int64(3), n], dtype="float32") - for i0, i1 in T.grid(T.int64(3), n): + rxplaceholder = T.match_buffer(var_rxplaceholder, [T.int64(10), n_strided_slice], dtype="float32") + T_strided_slice_with_axes = T.match_buffer(var_T_strided_slice_with_axes, [T.int64(3), n_strided_slice], dtype="float32") + for i0, i1 in T.grid(T.int64(3), n_strided_slice): with Ts.sblock("T_strided_slice_with_axes"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[ax0 * T.int64(3) + T.int64(1), ax1]) @@ -438,6 +455,15 @@ class DynamicStridedSlice: def main(x: R.Tensor((8, 9, 10, 10), "float32"), begin: R.Tensor((4,),"int64"), end: R.Tensor((4,),"int64"), strides: R.Tensor((4,),"int64")) -> R.Tensor("float32", ndim=4): gv: R.Tensor("float32", ndim=4) = R.dynamic_strided_slice(x, begin, end, strides) return gv + s_dynamic_strided_slice = T.dynamic("s") + s_1_dynamic_strided_slice = T.dynamic("s_1") + s_2_dynamic_strided_slice = T.dynamic("s_2") + s_3_dynamic_strided_slice = T.dynamic("s_3") + s_main = T.dynamic("s") + s_1_main = T.dynamic("s_1") + s_2_main = T.dynamic("s_2") + s_3_main = T.dynamic("s_3") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) @@ -451,12 +477,11 @@ def dynamic_strided_slice( var_T_strided_slice_dynamic: T.handle, ): T.func_attr({"tirx.noalias": True}) - s, s_1, s_2, s_3 = T.int64(), T.int64(), T.int64(), T.int64() T_strided_slice_dynamic = T.match_buffer( - var_T_strided_slice_dynamic, (s, s_1, s_2, s_3) + var_T_strided_slice_dynamic, (s_dynamic_strided_slice, s_1_dynamic_strided_slice, s_2_dynamic_strided_slice, s_3_dynamic_strided_slice) ) # with Ts.sblock("root"): - for ax0, ax1, ax2, ax3 in T.grid(s, s_1, s_2, s_3): + for ax0, ax1, ax2, ax3 in T.grid(s_dynamic_strided_slice, s_1_dynamic_strided_slice, s_2_dynamic_strided_slice, s_3_dynamic_strided_slice): with Ts.sblock("T_strided_slice_dynamic"): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) Ts.reads( @@ -549,9 +574,9 @@ def shape_func( ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): - for i in range(T.int64(4)): + for i_index in range(T.int64(4)): with Ts.sblock("T_shape_func_strided_slice_dynamic"): - v_i = Ts.axis.spatial(T.int64(4), i) + v_i = Ts.axis.spatial(T.int64(4), i_index) Ts.reads( rxplaceholder_3[v_i], rxplaceholder_1[v_i], rxplaceholder_2[v_i] ) @@ -747,23 +772,19 @@ def main( end: R.Tensor((4,), dtype="int64"), strides: R.Tensor((4,), dtype="int64"), ) -> R.Tensor(dtype="float32", ndim=4): - s = T.int64() - s_1 = T.int64() - s_2 = T.int64() - s_3 = T.int64() gv = R.call_tir( Expected.shape_func, (x, begin, end, strides), out_ty=R.Tensor((4,), dtype="int64"), ) gv1: R.Shape(ndim=4) = R.tensor_to_shape(gv) - gv2: R.Shape([s, s_1, s_2, s_3]) = R.match_cast( - gv1, R.Shape([s, s_1, s_2, s_3]) + gv2: R.Shape([s_main, s_1_main, s_2_main, s_3_main]) = R.match_cast( + gv1, R.Shape([s_main, s_1_main, s_2_main, s_3_main]) ) gv_1 = R.call_tir( Expected.dynamic_strided_slice, (x, begin, end, strides), - out_ty=R.Tensor((s, s_1, s_2, s_3), dtype="float32"), + out_ty=R.Tensor((s_main, s_1_main, s_2_main, s_3_main), dtype="float32"), ) return gv_1 # fmt: on @@ -773,13 +794,22 @@ def main( def test_dynamic_strided_slice_symbolic(): # fmt: off + n = T.dynamic("n") + @tvm.script.ir_module class DynamicStridedSlice: @R.function - def main(x: R.Tensor((10, "n"), "float32"), begin:R.Tensor((2,), "int64"), end:R.Tensor((2,), "int64"), strides:R.Tensor((2,), "int64")) -> R.Tensor("float32", ndim=2): - n = T.int64() + def main(x: R.Tensor((10, n), "float32"), begin:R.Tensor((2,), "int64"), end:R.Tensor((2,), "int64"), strides:R.Tensor((2,), "int64")) -> R.Tensor("float32", ndim=2): gv: R.Tensor("float32", ndim=2) = R.dynamic_strided_slice(x, begin, end, strides) return gv + n_dynamic_strided_slice = T.dynamic("n") + s_dynamic_strided_slice = T.dynamic("s") + s_1_dynamic_strided_slice = T.dynamic("s_1") + n_shape_func = T.dynamic("n") + n_main = T.dynamic("n") + s_main = T.dynamic("s") + s_1_main = T.dynamic("s_1") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) @@ -791,18 +821,16 @@ def dynamic_strided_slice( var_T_strided_slice_dynamic: T.handle, ): T.func_attr({"tirx.noalias": True}) - n = T.int64() - rxplaceholder_3 = T.match_buffer(var_rxplaceholder, (T.int64(10), n)) - s, s_1 = T.int64(), T.int64() - T_strided_slice_dynamic = T.match_buffer(var_T_strided_slice_dynamic, (s, s_1)) + rxplaceholder_3 = T.match_buffer(var_rxplaceholder, (T.int64(10), n_dynamic_strided_slice)) + T_strided_slice_dynamic = T.match_buffer(var_T_strided_slice_dynamic, (s_dynamic_strided_slice, s_1_dynamic_strided_slice)) # with Ts.sblock("root"): - for ax0, ax1 in T.grid(s, s_1): + for ax0, ax1 in T.grid(s_dynamic_strided_slice, s_1_dynamic_strided_slice): with Ts.sblock("T_strided_slice_dynamic"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads( rxplaceholder_3[ T.int64(0) : T.int64(10), - T.int64(0) : n, + T.int64(0) : n_dynamic_strided_slice, ], rxplaceholder[T.int64(0) : T.int64(2)], rxplaceholder_2[T.int64(0) : T.int64(2)], @@ -829,7 +857,7 @@ def dynamic_strided_slice( T.max( T.if_then_else( rxplaceholder[T.int64(1)] < T.int64(0), - rxplaceholder[T.int64(1)] + n, + rxplaceholder[T.int64(1)] + n_dynamic_strided_slice, rxplaceholder[T.int64(1)], ), T.if_then_else( @@ -837,7 +865,7 @@ def dynamic_strided_slice( ), ), T.if_then_else( - rxplaceholder_2[T.int64(1)] < T.int64(0), n - T.int64(1), n + rxplaceholder_2[T.int64(1)] < T.int64(0), n_dynamic_strided_slice - T.int64(1), n_dynamic_strided_slice ), ) + v_ax1 * rxplaceholder_2[T.int64(1)], @@ -852,12 +880,11 @@ def shape_func( T_shape_func_strided_slice_dynamic: T.Buffer((T.int64(2),), "int64"), ): T.func_attr({"tirx.noalias": True}) - n = T.int64() - rxplaceholder_3 = T.match_buffer(var_rxplaceholder, (T.int64(10), n)) + rxplaceholder_3 = T.match_buffer(var_rxplaceholder, (T.int64(10), n_shape_func)) # with Ts.sblock("root"): - for i in range(T.int64(2)): + for i_index in range(T.int64(2)): with Ts.sblock("T_shape_func_strided_slice_dynamic"): - v_i = Ts.axis.spatial(T.int64(2), i) + v_i = Ts.axis.spatial(T.int64(2), i_index) Ts.reads(rxplaceholder_2[v_i], rxplaceholder[v_i], rxplaceholder_1[v_i]) Ts.writes(T_shape_func_strided_slice_dynamic[v_i]) T_shape_func_strided_slice_dynamic[v_i] = T.Select( @@ -870,7 +897,7 @@ def shape_func( rxplaceholder[v_i] + T.Select( v_i == T.int64(1), - n, + n_shape_func, T.Select( v_i == T.int64(0), T.int64(10), T.int64(-1) ), @@ -881,7 +908,7 @@ def shape_func( ), T.Select( v_i == T.int64(1), - n, + n_shape_func, T.Select(v_i == T.int64(0), T.int64(10), T.int64(-1)), ) - T.int64(1), @@ -893,7 +920,7 @@ def shape_func( rxplaceholder_1[v_i] + T.Select( v_i == T.int64(1), - n, + n_shape_func, T.Select( v_i == T.int64(0), T.int64(10), T.int64(-1) ), @@ -904,7 +931,7 @@ def shape_func( ), T.Select( v_i == T.int64(1), - n, + n_shape_func, T.Select(v_i == T.int64(0), T.int64(10), T.int64(-1)), ) - T.int64(1), @@ -921,7 +948,7 @@ def shape_func( rxplaceholder_1[v_i] + T.Select( v_i == T.int64(1), - n, + n_shape_func, T.Select( v_i == T.int64(0), T.int64(10), T.int64(-1) ), @@ -932,7 +959,7 @@ def shape_func( ), T.Select( v_i == T.int64(1), - n, + n_shape_func, T.Select(v_i == T.int64(0), T.int64(10), T.int64(-1)), ), ) @@ -944,7 +971,7 @@ def shape_func( rxplaceholder[v_i] + T.Select( v_i == T.int64(1), - n, + n_shape_func, T.Select( v_i == T.int64(0), T.int64(10), T.int64(-1) ), @@ -955,7 +982,7 @@ def shape_func( ), T.Select( v_i == T.int64(1), - n, + n_shape_func, T.Select(v_i == T.int64(0), T.int64(10), T.int64(-1)), ), ) @@ -966,25 +993,22 @@ def shape_func( @R.function def main( - x: R.Tensor((10, "n"), dtype="float32"), + x: R.Tensor((10, n_main), dtype="float32"), begin: R.Tensor((2,), dtype="int64"), end: R.Tensor((2,), dtype="int64"), strides: R.Tensor((2,), dtype="int64"), ) -> R.Tensor(dtype="float32", ndim=2): - n = T.int64() - s = T.int64() - s_1 = T.int64() gv = R.call_tir( Expected.shape_func, (x, begin, end, strides), out_ty=R.Tensor((2,), dtype="int64"), ) gv1: R.Shape(ndim=2) = R.tensor_to_shape(gv) - gv2: R.Shape([s, s_1]) = R.match_cast(gv1, R.Shape([s, s_1])) + gv2: R.Shape([s_main, s_1_main]) = R.match_cast(gv1, R.Shape([s_main, s_1_main])) gv_1 = R.call_tir( Expected.dynamic_strided_slice, (x, begin, end, strides), - out_ty=R.Tensor((s, s_1), dtype="float32"), + out_ty=R.Tensor((s_main, s_1_main), dtype="float32"), ) return gv_1 # fmt: on @@ -1130,43 +1154,47 @@ def matmul(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64( def test_matmul_4_5_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + m = T.dynamic("m") + n = T.dynamic("n") + k = T.dynamic("k") + @tvm.script.ir_module class Matmul: @R.function - def main(x: R.Tensor(("b", 1, "m", "k"), "float32"), y: R.Tensor(("a", 1, "c", "k", "n"), "float32")) -> R.Tensor(("a", "b", "c", "m", "n"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - m = T.int64() - n = T.int64() + def main(x: R.Tensor((b, 1, m, k), "float32"), y: R.Tensor((a, 1, c, k, n), "float32")) -> R.Tensor((a, b, c, m, n), "float32"): gv: R.Tensor((a, b, c, m, n), "float32") = R.matmul(x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + m_main = T.dynamic("m") + n_main = T.dynamic("n") + k_main = T.dynamic("k") + a_matmul = T.dynamic("a") + b_matmul = T.dynamic("b") + c_matmul = T.dynamic("c") + k_matmul = T.dynamic("k") + m_matmul = T.dynamic("m") + n_matmul = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("b", 1, "m", "k"), "float32"), y: R.Tensor(("a", 1, "c", "k", "n"), "float32")) -> R.Tensor(("a", "b", "c", "m", "n"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.matmul, (x, y), R.Tensor((a, b, c, m, n), dtype="float32")) + def main(x: R.Tensor((b_main, 1, m_main, k_main), "float32"), y: R.Tensor((a_main, 1, c_main, k_main, n_main), "float32")) -> R.Tensor((a_main, b_main, c_main, m_main, n_main), "float32"): + gv = R.call_tir(Expected.matmul, (x, y), R.Tensor((a_main, b_main, c_main, m_main, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def matmul(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_matmul: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - k = T.int64() - m = T.int64() - n = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [b, T.int64(1), m, k], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, T.int64(1), c, k, n], dtype="float32") - matmul = T.match_buffer(var_matmul, [a, b, c, m, n], dtype="float32") - for i0, i1, i2, i3, i4, i5 in T.grid(a, b, c, m, n, k): + rxplaceholder = T.match_buffer(var_rxplaceholder, [b_matmul, T.int64(1), m_matmul, k_matmul], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_matmul, T.int64(1), c_matmul, k_matmul, n_matmul], dtype="float32") + matmul = T.match_buffer(var_matmul, [a_matmul, b_matmul, c_matmul, m_matmul, n_matmul], dtype="float32") + for i0, i1, i2, i3, i4, i5 in T.grid(a_matmul, b_matmul, c_matmul, m_matmul, n_matmul, k_matmul): with Ts.sblock("matmul"): i0_1, i1_1, i2_1, i3_1, i4_1, k_1 = Ts.axis.remap("SSSSSR", [i0, i1, i2, i3, i4, i5]) Ts.reads(rxplaceholder[i1_1, T.int64(0), i3_1, k_1], rxplaceholder_1[i0_1, T.int64(0), i2_1, k_1, i4_1]) @@ -1195,9 +1223,9 @@ class Expected: def matmul(A: T.Buffer((T.int64(1), T.int64(1), T.int64(4), T.int64(5)), "float32"), B: T.Buffer((T.int64(1), T.int64(1), T.int64(5), T.int64(7)), "float32"), matmul_1: T.Buffer((T.int64(1), T.int64(1), T.int64(4), T.int64(7)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): - for i0, i1, i2, i3, k in T.grid(T.int64(1), T.int64(1), T.int64(4), T.int64(7), T.int64(5)): + for i0, i1, i2, i3, k_index in T.grid(T.int64(1), T.int64(1), T.int64(4), T.int64(7), T.int64(5)): with Ts.sblock("matmul"): - v_i0, v_i1, v_i2, v_i3, v_k = Ts.axis.remap("SSSSR", [i0, i1, i2, i3, k]) + v_i0, v_i1, v_i2, v_i3, v_k = Ts.axis.remap("SSSSR", [i0, i1, i2, i3, k_index]) Ts.reads(A[v_i0, v_i1, v_i2, v_k], B[v_i0, v_i1, v_k, v_i3]) Ts.writes(matmul_1[v_i0, v_i1, v_i2, v_i3]) with Ts.init(): @@ -1276,25 +1304,33 @@ def einsum( def test_einsum_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + @I.ir_module class Einsum: @R.function - def main(x: R.Tensor(("a", "b"), "float32"), y: R.Tensor(("b", "c"), "float32")): + def main(x: R.Tensor((a, b), "float32"), y: R.Tensor((b, c), "float32")): gv = R.einsum((x, y), subscripts="ij,jk->ik") return gv + a_main = T.dynamic("a") + c_main = T.dynamic("c") + b_main = T.dynamic("b") + a_einsum = T.dynamic("a") + b_einsum = T.dynamic("b") + c_einsum = T.dynamic("c") + @I.ir_module class Expected: @R.function def main( - x: R.Tensor(("a", "b"), dtype="float32"), - y: R.Tensor(("b", "c"), dtype="float32"), - ) -> R.Tensor(("a", "c"), dtype="float32"): - a = T.int64() - c = T.int64() - b = T.int64() + x: R.Tensor((a_main, b_main), dtype="float32"), + y: R.Tensor((b_main, c_main), dtype="float32"), + ) -> R.Tensor((a_main, c_main), dtype="float32"): cls = Expected - gv = R.call_tir(cls.einsum, (x, y), out_ty=R.Tensor((a, c), dtype="float32")) + gv = R.call_tir(cls.einsum, (x, y), out_ty=R.Tensor((a_main, c_main), dtype="float32")) return gv @Ts.prim_func(private=True) @@ -1304,12 +1340,10 @@ def einsum( var_T_einsum: T.handle, ): T.func_attr({"tirx.noalias": True}) - a, b = T.int64(), T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (a, b)) - c = T.int64() - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (b, c)) - T_einsum = T.match_buffer(var_T_einsum, (a, c)) - for ax0, ax1, j in T.grid(a, c, b): + rxplaceholder = T.match_buffer(var_rxplaceholder, (a_einsum, b_einsum)) + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (b_einsum, c_einsum)) + T_einsum = T.match_buffer(var_T_einsum, (a_einsum, c_einsum)) + for ax0, ax1, j in T.grid(a_einsum, c_einsum, b_einsum): with Ts.sblock("T_einsum"): v_ax0, v_ax1, v_j = Ts.axis.remap("SSR", [ax0, ax1, j]) Ts.reads(rxplaceholder[v_ax0, v_j], rxplaceholder_1[v_j, v_ax1]) diff --git a/tests/python/relax/test_transform_legalize_ops_manipulate.py b/tests/python/relax/test_transform_legalize_ops_manipulate.py index 40d4e12044c9..d2a4b8f2605d 100644 --- a/tests/python/relax/test_transform_legalize_ops_manipulate.py +++ b/tests/python/relax/test_transform_legalize_ops_manipulate.py @@ -62,38 +62,40 @@ def broadcast_to(rxplaceholder: T.Buffer((T.int64(2), T.int64(1), T.int64(3)), " def test_broadcast_to_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class BroadcastTo: @R.function - def main(dumb_param: R.Tensor(("a", "c")), x: R.Tensor(("b", 1, "d"), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(dumb_param: R.Tensor((a, c)), x: R.Tensor((b, 1, d), "float32")) -> R.Tensor((a, b, c, d), "float32"): gv: R.Tensor((a, b, c, d), "float32") = R.broadcast_to(x, (a, b, c, d)) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_broadcast_to = T.dynamic("a") + b_broadcast_to = T.dynamic("b") + c_broadcast_to = T.dynamic("c") + d_broadcast_to = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(dumb_param: R.Tensor(("a", "c")), x: R.Tensor(("b", 1, "d"), "float32")) -> R.Tensor(("a", "b", "c", "d"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.broadcast_to, (x,), R.Tensor((a, b, c, d), dtype="float32")) + def main(dumb_param: R.Tensor((a_main, c_main)), x: R.Tensor((b_main, 1, d_main), "float32")) -> R.Tensor((a_main, b_main, c_main, d_main), "float32"): + gv = R.call_tir(Expected.broadcast_to, (x,), R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) return gv @Ts.prim_func(private=True) def broadcast_to(var_rxplaceholder: T.handle, var_T_broadcast_to: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [b, T.int64(1), d], dtype="float32") - T_broadcast_to = T.match_buffer(var_T_broadcast_to, [a, b, c, d], dtype="float32") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [b_broadcast_to, T.int64(1), d_broadcast_to], dtype="float32") + T_broadcast_to = T.match_buffer(var_T_broadcast_to, [a_broadcast_to, b_broadcast_to, c_broadcast_to, d_broadcast_to], dtype="float32") + for i0, i1, i2, i3 in T.grid(a_broadcast_to, b_broadcast_to, c_broadcast_to, d_broadcast_to): with Ts.sblock("T_broadcast_to"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[ax1, T.int64(0), ax3]) @@ -171,48 +173,50 @@ def concatenate(rxplaceholder: T.Buffer((T.int64(3), T.int64(4)), "float32"), rx def test_concat_input_tuple_var_symbolic(): # fmt: off + a = T.dynamic("a") + b0 = T.dynamic("b0") + b1 = T.dynamic("b1") + b2 = T.dynamic("b2") + @tvm.script.ir_module class Concat: @R.function - def main(t: R.Tuple(R.Tensor(("a", "b0"), "float32"), R.Tensor(("a", "b1"), "float32"), R.Tensor(("a", "b2"), "float32"))) -> R.Tensor(("a", "b0 + b1 + b2"), "float32"): - a = T.int64() - b0 = T.int64() - b1 = T.int64() - b2 = T.int64() + def main(t: R.Tuple(R.Tensor((a, b0), "float32"), R.Tensor((a, b1), "float32"), R.Tensor((a, b2), "float32"))) -> R.Tensor((a, b0 + b1 + b2), "float32"): gv: R.Tensor((a, b0 + b1 + b2), "float32") = R.concat(t, axis=1) return gv + a_main = T.dynamic("a") + b0_main = T.dynamic("b0") + b1_main = T.dynamic("b1") + b2_main = T.dynamic("b2") + a_concatenate = T.dynamic("a") + b0_concatenate = T.dynamic("b0") + b1_concatenate = T.dynamic("b1") + b2_concatenate = T.dynamic("b2") + @tvm.script.ir_module class Expected: @R.function - def main(t: R.Tuple(R.Tensor(("a", "b0"), "float32"), R.Tensor(("a", "b1"), "float32"), R.Tensor(("a", "b2"), "float32"))) -> R.Tensor(("a", "b0 + b1 + b2"), "float32"): - a = T.int64() - b0 = T.int64() - b1 = T.int64() - b2 = T.int64() - gv: R.Tensor((a, b0), dtype="float32") = t[0] - gv1: R.Tensor((a, b1), dtype="float32") = t[1] - gv2: R.Tensor((a, b2), dtype="float32") = t[2] - gv3 = R.call_tir(Expected.concatenate, (gv, gv1, gv2), R.Tensor((a, ((b0 + b1) + b2)), dtype="float32")) + def main(t: R.Tuple(R.Tensor((a_main, b0_main), "float32"), R.Tensor((a_main, b1_main), "float32"), R.Tensor((a_main, b2_main), "float32"))) -> R.Tensor((a_main, b0_main + b1_main + b2_main), "float32"): + gv: R.Tensor((a_main, b0_main), dtype="float32") = t[0] + gv1: R.Tensor((a_main, b1_main), dtype="float32") = t[1] + gv2: R.Tensor((a_main, b2_main), dtype="float32") = t[2] + gv3 = R.call_tir(Expected.concatenate, (gv, gv1, gv2), R.Tensor((a_main, ((b0_main + b1_main) + b2_main)), dtype="float32")) return gv3 @Ts.prim_func(private=True) def concatenate(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_rxplaceholder_2: T.handle, var_T_concat: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b0 = T.int64() - b1 = T.int64() - b2 = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b0], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a, b1], dtype="float32") - rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, [a, b2], dtype="float32") - T_concat = T.match_buffer(var_T_concat, [a, b0 + b1 + b2], dtype="float32") - for i0, i1 in T.grid(a, b0 + b1 + b2): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_concatenate, b0_concatenate], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [a_concatenate, b1_concatenate], dtype="float32") + rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, [a_concatenate, b2_concatenate], dtype="float32") + T_concat = T.match_buffer(var_T_concat, [a_concatenate, b0_concatenate + b1_concatenate + b2_concatenate], dtype="float32") + for i0, i1 in T.grid(a_concatenate, b0_concatenate + b1_concatenate + b2_concatenate): with Ts.sblock("T_concat"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) - Ts.reads(rxplaceholder_2[ax0, ax1 - b0 - b1], rxplaceholder_1[ax0, ax1 - b0], rxplaceholder[ax0, ax1]) + Ts.reads(rxplaceholder_2[ax0, ax1 - b0_concatenate - b1_concatenate], rxplaceholder_1[ax0, ax1 - b0_concatenate], rxplaceholder[ax0, ax1]) Ts.writes(T_concat[ax0, ax1]) - T_concat[ax0, ax1] = T.if_then_else(T.int64(0) <= ax1 - b0 - b1, rxplaceholder_2[ax0, ax1 - b0 - b1], T.if_then_else(T.int64(0) <= ax1 - b0, rxplaceholder_1[ax0, ax1 - b0], rxplaceholder[ax0, ax1])) + T_concat[ax0, ax1] = T.if_then_else(T.int64(0) <= ax1 - b0_concatenate - b1_concatenate, rxplaceholder_2[ax0, ax1 - b0_concatenate - b1_concatenate], T.if_then_else(T.int64(0) <= ax1 - b0_concatenate, rxplaceholder_1[ax0, ax1 - b0_concatenate], rxplaceholder[ax0, ax1])) # fmt: on mod = LegalizeOps()(Concat) @@ -252,35 +256,37 @@ def expand_dims(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "f def test_expand_dims_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + @tvm.script.ir_module class ExpandDims: @R.function - def main(x: R.Tensor(("a", "b", "c"), "float32")) -> R.Tensor(("a", 1, "b", 1, "c", 1), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() + def main(x: R.Tensor((a, b, c), "float32")) -> R.Tensor((a, 1, b, 1, c, 1), "float32"): gv: R.Tensor((a, 1, b, 1, c, 1), "float32") = R.expand_dims(x, axis=[1, 3, 5]) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + a_expand_dims = T.dynamic("a") + b_expand_dims = T.dynamic("b") + c_expand_dims = T.dynamic("c") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c"), "float32")) -> R.Tensor(("a", 1, "b", 1, "c", 1), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - gv = R.call_tir(Expected.expand_dims, (x,), R.Tensor((a, 1, b, 1, c, 1), dtype="float32")) + def main(x: R.Tensor((a_main, b_main, c_main), "float32")) -> R.Tensor((a_main, 1, b_main, 1, c_main, 1), "float32"): + gv = R.call_tir(Expected.expand_dims, (x,), R.Tensor((a_main, 1, b_main, 1, c_main, 1), dtype="float32")) return gv @Ts.prim_func(private=True) def expand_dims(var_rxplaceholder: T.handle, var_expand_dims: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, c], dtype="float32") - expand_dims = T.match_buffer(var_expand_dims, [a, T.int64(1), b, T.int64(1), c, T.int64(1)], dtype="float32") - for i0, i1, i2, i3, i4, i5 in T.grid(a, T.int64(1), b, T.int64(1), c, T.int64(1)): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_expand_dims, b_expand_dims, c_expand_dims], dtype="float32") + expand_dims = T.match_buffer(var_expand_dims, [a_expand_dims, T.int64(1), b_expand_dims, T.int64(1), c_expand_dims, T.int64(1)], dtype="float32") + for i0, i1, i2, i3, i4, i5 in T.grid(a_expand_dims, T.int64(1), b_expand_dims, T.int64(1), c_expand_dims, T.int64(1)): with Ts.sblock("expand_dims"): i0_1, i1_1, i2_1, i3_1, i4_1, i5_1 = Ts.axis.remap("SSSSSS", [i0, i1, i2, i3, i4, i5]) Ts.reads(rxplaceholder[i0_1, i2_1, i4_1]) @@ -356,40 +362,42 @@ def reshape(rxplaceholder: T.Buffer((), "float32"), T_reshape: T.Buffer(T.int64( def test_flatten_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + @tvm.script.ir_module class Flatten: @R.function - def main(x: R.Tensor(("a", "b", "c"), "float32")) -> R.Tensor(("a * b * c",), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() + def main(x: R.Tensor((a, b, c), "float32")) -> R.Tensor((a * b * c,), "float32"): gv: R.Tensor((a * b * c,), "float32") = R.flatten(x) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + a_reshape = T.dynamic("a") + b_reshape = T.dynamic("b") + c_reshape = T.dynamic("c") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c"), "float32")) -> R.Tensor(("a * b * c",), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - gv = R.call_tir(Expected.reshape, (x,), R.Tensor((((a * b) * c),), dtype="float32")) + def main(x: R.Tensor((a_main, b_main, c_main), "float32")) -> R.Tensor((a_main * b_main * c_main,), "float32"): + gv = R.call_tir(Expected.reshape, (x,), R.Tensor((((a_main * b_main) * c_main),), dtype="float32")) return gv @Ts.prim_func(private=True) def reshape(var_rxplaceholder: T.handle, var_T_reshape: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, c], dtype="float32") - T_reshape = T.match_buffer(var_T_reshape, [a * b * c], dtype="float32") - for i0 in T.serial(a * b * c): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_reshape, b_reshape, c_reshape], dtype="float32") + T_reshape = T.match_buffer(var_T_reshape, [a_reshape * b_reshape * c_reshape], dtype="float32") + for i0 in T.serial(a_reshape * b_reshape * c_reshape): with Ts.sblock("T_reshape"): - ax0 = Ts.axis.spatial(a * b * c, i0) - Ts.reads(rxplaceholder[ax0 // c // b % a, ax0 // c % b, ax0 % c]) + ax0 = Ts.axis.spatial(a_reshape * b_reshape * c_reshape, i0) + Ts.reads(rxplaceholder[ax0 // c_reshape // b_reshape % a_reshape, ax0 // c_reshape % b_reshape, ax0 % c_reshape]) Ts.writes(T_reshape[ax0]) - T_reshape[ax0] = rxplaceholder[ax0 // c // b % a, ax0 // c % b, ax0 % c] + T_reshape[ax0] = rxplaceholder[ax0 // c_reshape // b_reshape % a_reshape, ax0 // c_reshape % b_reshape, ax0 % c_reshape] # fmt: on mod = LegalizeOps()(Flatten) @@ -429,38 +437,40 @@ def transpose(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3), T.int def test_permute_dims_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class PermuteDims: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor(("b", "d", "c", "a"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((b, d, c, a), "float32"): gv: R.Tensor((b, d, c, a), "float32") = R.permute_dims(x, axes=[1, -1, 2, -4]) return gv + b_main = T.dynamic("b") + d_main = T.dynamic("d") + c_main = T.dynamic("c") + a_main = T.dynamic("a") + a_transpose = T.dynamic("a") + b_transpose = T.dynamic("b") + c_transpose = T.dynamic("c") + d_transpose = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), dtype="float32")) -> R.Tensor(("b", "d", "c", "a"), dtype="float32"): - b = T.int64() - d = T.int64() - c = T.int64() - a = T.int64() - gv = R.call_tir(Expected.transpose, (x,), R.Tensor((b, d, c, a), dtype="float32")) + def main(x: R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) -> R.Tensor((b_main, d_main, c_main, a_main), dtype="float32"): + gv = R.call_tir(Expected.transpose, (x,), R.Tensor((b_main, d_main, c_main, a_main), dtype="float32")) return gv @Ts.prim_func(private=True) def transpose(var_rxplaceholder: T.handle, var_T_transpose: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, c, d], dtype="float32") - T_transpose = T.match_buffer(var_T_transpose, [b, d, c, a], dtype="float32") - for i0, i1, i2, i3 in T.grid(b, d, c, a): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_transpose, b_transpose, c_transpose, d_transpose], dtype="float32") + T_transpose = T.match_buffer(var_T_transpose, [b_transpose, d_transpose, c_transpose, a_transpose], dtype="float32") + for i0, i1, i2, i3 in T.grid(b_transpose, d_transpose, c_transpose, a_transpose): with Ts.sblock("T_transpose"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[ax3, ax0, ax2, ax1]) @@ -554,133 +564,148 @@ def main(x: R.Tensor((1, 2, 3, 4), dtype="float32")) -> R.Tensor((8, 3), dtype=" def test_reshape_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + @tvm.script.ir_module class Reshape: @R.function - def main(x: R.Tensor(("a", "b"), "float32")) -> R.Tensor(("a // 2", "b * 2"), "float32"): - a = T.int64() - b = T.int64() + def main(x: R.Tensor((a, b), "float32")) -> R.Tensor((a // 2, b * 2), "float32"): gv: R.Tensor((a // 2, b * 2), "float32") = R.reshape(x, (a // 2, b * 2)) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + a_reshape = T.dynamic("a") + b_reshape = T.dynamic("b") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b"), "float32")) -> R.Tensor(("a // 2", "b * 2"), "float32"): - a = T.int64() - b = T.int64() - gv = R.call_tir(Expected.reshape, (x,), R.Tensor(((a // 2), (b * 2)), dtype="float32")) + def main(x: R.Tensor((a_main, b_main), "float32")) -> R.Tensor((a_main // 2, b_main * 2), "float32"): + gv = R.call_tir(Expected.reshape, (x,), R.Tensor(((a_main // 2), (b_main * 2)), dtype="float32")) return gv @Ts.prim_func(private=True) def reshape(var_rxplaceholder: T.handle, var_T_reshape: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b], dtype="float32") - T_reshape = T.match_buffer(var_T_reshape, [a // T.int64(2), b * T.int64(2)], dtype="float32") - for i0, i1 in T.grid(a // T.int64(2), b * T.int64(2)): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_reshape, b_reshape], dtype="float32") + T_reshape = T.match_buffer(var_T_reshape, [a_reshape // T.int64(2), b_reshape * T.int64(2)], dtype="float32") + for i0, i1 in T.grid(a_reshape // T.int64(2), b_reshape * T.int64(2)): with Ts.sblock("T_reshape"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) - Ts.reads(rxplaceholder[(ax0 * b * T.int64(2) + ax1) // b % a, (ax0 * b * T.int64(2) + ax1) % b]) + Ts.reads(rxplaceholder[(ax0 * b_reshape * T.int64(2) + ax1) // b_reshape % a_reshape, (ax0 * b_reshape * T.int64(2) + ax1) % b_reshape]) Ts.writes(T_reshape[ax0, ax1]) - T_reshape[ax0, ax1] = rxplaceholder[(ax0 * b * T.int64(2) + ax1) // b % a, (ax0 * b * T.int64(2) + ax1) % b] + T_reshape[ax0, ax1] = rxplaceholder[(ax0 * b_reshape * T.int64(2) + ax1) // b_reshape % a_reshape, (ax0 * b_reshape * T.int64(2) + ax1) % b_reshape] # fmt: on mod = LegalizeOps()(Reshape) tvm.ir.assert_structural_equal(mod, Expected) # ShapeExpr might be produced by shape computation + a = T.dynamic("a") + b = T.dynamic("b") + @tvm.script.ir_module class Reshape2: @R.function - def main(x: R.Tensor(("a", "b"), "float32")) -> R.Tensor(("a // 2", "b * 2"), "float32"): - a = T.int64() - b = T.int64() + def main(x: R.Tensor((a, b), "float32")) -> R.Tensor((a // 2, b * 2), "float32"): lv: R.Shape((a // 2, b * 2)) = R.shape((a // 2, b * 2)) gv: R.Tensor((a // 2, b * 2), "float32") = R.reshape(x, lv) return gv # After lowering, redundant var might be removed by later dead code elimination + a_main = T.dynamic("a") + b_main = T.dynamic("b") + a_reshape = T.dynamic("a") + b_reshape = T.dynamic("b") + @tvm.script.ir_module class Expected2: @R.function - def main(x: R.Tensor(("a", "b"), "float32")) -> R.Tensor(("a // 2", "b * 2"), "float32"): - a = T.int64() - b = T.int64() - lv: R.Shape((a // 2, b * 2)) = R.shape((a // 2, b * 2)) - gv = R.call_tir(Expected2.reshape, (x,), R.Tensor(((a // 2), (b * 2)), dtype="float32")) + def main(x: R.Tensor((a_main, b_main), "float32")) -> R.Tensor( + (a_main // 2, b_main * 2), "float32" + ): + lv: R.Shape((a_main // 2, b_main * 2)) = R.shape((a_main // 2, b_main * 2)) + gv = R.call_tir( + Expected2.reshape, (x,), R.Tensor(((a_main // 2), (b_main * 2)), dtype="float32") + ) return gv @Ts.prim_func(private=True) def reshape(var_rxplaceholder: T.handle, var_T_reshape: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b], dtype="float32") + rxplaceholder = T.match_buffer( + var_rxplaceholder, [a_reshape, b_reshape], dtype="float32" + ) T_reshape = T.match_buffer( - var_T_reshape, [a // T.int64(2), b * T.int64(2)], dtype="float32" + var_T_reshape, [a_reshape // T.int64(2), b_reshape * T.int64(2)], dtype="float32" ) - for i0, i1 in T.grid(a // T.int64(2), b * T.int64(2)): + for i0, i1 in T.grid(a_reshape // T.int64(2), b_reshape * T.int64(2)): with Ts.sblock("T_reshape"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads( rxplaceholder[ - (ax0 * b * T.int64(2) + ax1) // b % a, - (ax0 * b * T.int64(2) + ax1) % b, + (ax0 * b_reshape * T.int64(2) + ax1) // b_reshape % a_reshape, + (ax0 * b_reshape * T.int64(2) + ax1) % b_reshape, ] ) Ts.writes(T_reshape[ax0, ax1]) T_reshape[ax0, ax1] = rxplaceholder[ - (ax0 * b * T.int64(2) + ax1) // b % a, (ax0 * b * T.int64(2) + ax1) % b + (ax0 * b_reshape * T.int64(2) + ax1) // b_reshape % a_reshape, + (ax0 * b_reshape * T.int64(2) + ax1) % b_reshape, ] mod2 = LegalizeOps()(Reshape2) tvm.ir.assert_structural_equal(mod2, Expected2) # ShapeExpr might be produced by shape computation + a = T.dynamic("a") + b = T.dynamic("b") + @I.ir_module class Reshape3: @R.function - def main(x: R.Tensor((10, "b"), "float32")) -> R.Tensor((5, "b * 2"), "float32"): - a = T.int64() - b = T.int64() + def main(x: R.Tensor((10, b), "float32")) -> R.Tensor((5, b * 2), "float32"): lv: R.Shape((5, b * 2)) = R.shape((5, b * 2)) gv: R.Tensor((5, b * 2), "float32") = R.reshape(x, lv) return gv # After lowering, redundant var might be removed by later dead code elimination + b_reshape = T.dynamic("b") + b_main = T.dynamic("b") + @I.ir_module class Expected3: @Ts.prim_func(private=True) def reshape(var_rxplaceholder: T.handle, var_T_reshape: T.handle): T.func_attr({"tirx.noalias": True}) - b = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (T.int64(10), b)) - T_reshape = T.match_buffer(var_T_reshape, (T.int64(5), b * T.int64(2))) + rxplaceholder = T.match_buffer(var_rxplaceholder, (T.int64(10), b_reshape)) + T_reshape = T.match_buffer(var_T_reshape, (T.int64(5), b_reshape * T.int64(2))) # with Ts.sblock("root"): - for ax0, ax1 in T.grid(T.int64(5), b * T.int64(2)): + for ax0, ax1 in T.grid(T.int64(5), b_reshape * T.int64(2)): with Ts.sblock("T_reshape"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads( rxplaceholder[ - (v_ax0 * b * T.int64(2) + v_ax1) // b % T.int64(10), - (v_ax0 * b * T.int64(2) + v_ax1) % b, + (v_ax0 * b_reshape * T.int64(2) + v_ax1) // b_reshape % T.int64(10), + (v_ax0 * b_reshape * T.int64(2) + v_ax1) % b_reshape, ] ) Ts.writes(T_reshape[v_ax0, v_ax1]) T_reshape[v_ax0, v_ax1] = rxplaceholder[ - (v_ax0 * b * T.int64(2) + v_ax1) // b % T.int64(10), - (v_ax0 * b * T.int64(2) + v_ax1) % b, + (v_ax0 * b_reshape * T.int64(2) + v_ax1) // b_reshape % T.int64(10), + (v_ax0 * b_reshape * T.int64(2) + v_ax1) % b_reshape, ] @R.function def main( - x: R.Tensor((10, "b"), dtype="float32"), - ) -> R.Tensor((5, "b * 2"), dtype="float32"): - b = T.int64() - lv: R.Shape([5, b * 2]) = R.shape([5, b * 2]) - gv = R.call_tir(Expected3.reshape, (x,), out_ty=R.Tensor((5, b * 2), dtype="float32")) + x: R.Tensor((10, b_main), dtype="float32"), + ) -> R.Tensor((5, b_main * 2), dtype="float32"): + lv: R.Shape([5, b_main * 2]) = R.shape([5, b_main * 2]) + gv = R.call_tir( + Expected3.reshape, (x,), out_ty=R.Tensor((5, b_main * 2), dtype="float32") + ) return gv mod3 = LegalizeOps()(Reshape3) @@ -706,6 +731,11 @@ def main( out_mod = relax.transform.LegalizeOps()(mod) # fmt: off + M_main = T.dynamic("M") + N_main = T.dynamic("N") + M_reshape = T.dynamic("M") + N_reshape = T.dynamic("N") + @I.ir_module class Expected: @R.function @@ -713,12 +743,10 @@ def main( x: R.Tensor([2], dtype="int64"), y: R.Tensor([16],dtype="float32"), ) -> R.Tensor(ndim=2, dtype="float32"): - M = T.int64() - N = T.int64() gv = R.call_pure_packed("vm.builtin.tensor_to_shape", x, ty_args=(R.Shape(ndim=2),)) - _ = R.match_cast(gv, R.Shape([M,N])) - _ = R.shape([M,N]) - gv_1 = R.call_tir(Expected.reshape, (y,), out_ty=R.Tensor([M,N], dtype="float32")) + _ = R.match_cast(gv, R.Shape([M_main,N_main])) + _ = R.shape([M_main,N_main]) + gv_1 = R.call_tir(Expected.reshape, (y,), out_ty=R.Tensor([M_main,N_main], dtype="float32")) return gv_1 @Ts.prim_func(private=True) @@ -727,15 +755,13 @@ def reshape( var_T_reshape: T.handle, ): T.func_attr({"tirx.noalias": True}) - M = T.int64() - N = T.int64() - T_reshape = T.match_buffer(var_T_reshape, [M,N], "float32") - for i,j in T.grid(M,N): + T_reshape = T.match_buffer(var_T_reshape, [M_reshape,N_reshape], "float32") + for i,j in T.grid(M_reshape,N_reshape): with Ts.sblock("T_reshape"): vi,vj = Ts.axis.remap('SS',[i,j]) - Ts.reads(rxplaceholder[(vi*N + vj) % 16]) + Ts.reads(rxplaceholder[(vi*N_reshape + vj) % 16]) Ts.writes(T_reshape[vi,vj]) - T_reshape[vi,vj] = rxplaceholder[(vi*N + vj) % 16] + T_reshape[vi,vj] = rxplaceholder[(vi*N_reshape + vj) % 16] # fmt: on tvm.ir.assert_structural_equal(out_mod, Expected) @@ -867,45 +893,47 @@ def split(rxplaceholder: T.Buffer((T.int64(2), T.int64(10), T.int64(4)), "float3 def test_split_by_indices_n_section_divisible_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Split: @R.function - def main(dumb_param: R.Tensor(("n",)), x: R.Tensor(("m", "n * 3"), "float32")) -> R.Tuple([R.Tensor(("m", "n"), "float32"), R.Tensor(("m", "n"), "float32"), R.Tensor(("m", "n"), "float32")]): - m = T.int64() - n = T.int64() + def main(dumb_param: R.Tensor((n,)), x: R.Tensor((m, n * 3), "float32")) -> R.Tuple([R.Tensor((m, n), "float32"), R.Tensor((m, n), "float32"), R.Tensor((m, n), "float32")]): gv: R.Tuple([R.Tensor((m, n), "float32"), R.Tensor((m, n), "float32"), R.Tensor((m, n), "float32")]) = R.split(x, 3, axis=1) return gv + m_main = T.dynamic("m") + n = T.dynamic("n") + m_split = T.dynamic("m") + @tvm.script.ir_module class Expected: @R.function - def main(dumb_param: R.Tensor(("n",)), x: R.Tensor(("m", "(n * 3)"), "float32")) -> R.Tuple(R.Tensor(("m", "((n * 3) // 3)"), "float32"), R.Tensor(("m", "((((n * 3) // 3) * 2) - ((n * 3) // 3))"), "float32"), R.Tensor(("m", "((n * 3) - (((n * 3) // 3) * 2))"), "float32")): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.split, (x, n), [R.Tensor((m, ((n * 3 + 3 - 1) // 3)), "float32"), R.Tensor((m, ((((n * 3 + 3 - 1) // 3) * 2) - ((n * 3 + 3 - 1) // 3))), "float32"), R.Tensor((m, ((n * 3) - (((n * 3 + 3 - 1) // 3) * 2))), "float32")]) + def main(dumb_param: R.Tensor((n,)), x: R.Tensor((m_main, n * 3), "float32")) -> R.Tuple(R.Tensor((m_main, n * 3 // 3), "float32"), R.Tensor((m_main, n * 3 // 3 * 2 - n * 3 // 3), "float32"), R.Tensor((m_main, n * 3 - n * 3 // 3 * 2), "float32")): + gv = R.call_tir(Expected.split, (x, n), [R.Tensor((m_main, ((n * 3 + 3 - 1) // 3)), "float32"), R.Tensor((m_main, ((((n * 3 + 3 - 1) // 3) * 2) - ((n * 3 + 3 - 1) // 3))), "float32"), R.Tensor((m_main, ((n * 3) - (((n * 3 + 3 - 1) // 3) * 2))), "float32")]) return gv @Ts.prim_func(private=True) def split(var_rxplaceholder: T.handle, n: T.int64, var_T_split_sections: T.handle, var_T_split_sections_1: T.handle, var_T_split_sections_2: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [m, n * T.int64(3)], dtype="float32") - T_split_sections = T.match_buffer(var_T_split_sections, [m, (n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3)], dtype="float32") - T_split_sections_1 = T.match_buffer(var_T_split_sections_1, [m, (n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3) * T.int64(2) - (n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3)], dtype="float32") - T_split_sections_2 = T.match_buffer(var_T_split_sections_2, [m, n * T.int64(3) - (n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3) * T.int64(2)], dtype="float32") - for i0, i1 in T.grid(m, n): + rxplaceholder = T.match_buffer(var_rxplaceholder, [m_split, n * T.int64(3)], dtype="float32") + T_split_sections = T.match_buffer(var_T_split_sections, [m_split, (n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3)], dtype="float32") + T_split_sections_1 = T.match_buffer(var_T_split_sections_1, [m_split, (n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3) * T.int64(2) - (n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3)], dtype="float32") + T_split_sections_2 = T.match_buffer(var_T_split_sections_2, [m_split, n * T.int64(3) - (n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3) * T.int64(2)], dtype="float32") + for i0, i1 in T.grid(m_split, n): with Ts.sblock("T_split_sections"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[ax0, ax1]) Ts.writes(T_split_sections[ax0, ax1]) T_split_sections[ax0, ax1] = rxplaceholder[ax0, ax1] - for i0, i1 in T.grid(m, n): + for i0, i1 in T.grid(m_split, n): with Ts.sblock("T_split_sections_1"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[ax0, ax1 + n]) Ts.writes(T_split_sections_1[ax0, ax1]) T_split_sections_1[ax0, ax1] = rxplaceholder[ax0, ax1 + n] - for i0, i1 in T.grid(m, n): + for i0, i1 in T.grid(m_split, n): with Ts.sblock("T_split_sections_2"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[ax0, n * T.int64(2) + ax1]) @@ -981,32 +1009,34 @@ def squeeze(rxplaceholder: T.Buffer((T.int64(2), T.int64(1), T.int64(3), T.int64 def test_squeeze_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + @tvm.script.ir_module class Squeeze: @R.function - def main(x: R.Tensor(("a", 1, "b", 1), "float32")) -> R.Tensor(("a", "b", 1), "float32"): - a = T.int64() - b = T.int64() + def main(x: R.Tensor((a, 1, b, 1), "float32")) -> R.Tensor((a, b, 1), "float32"): gv: R.Tensor((a, b, 1), "float32") = R.squeeze(x, [1]) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + a_squeeze = T.dynamic("a") + b_squeeze = T.dynamic("b") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", 1, "b", 1), "float32")) -> R.Tensor(("a", "b", 1), "float32"): - a = T.int64() - b = T.int64() - gv = R.call_tir(Expected.squeeze, (x,), R.Tensor((a, b, 1), dtype="float32")) + def main(x: R.Tensor((a_main, 1, b_main, 1), "float32")) -> R.Tensor((a_main, b_main, 1), "float32"): + gv = R.call_tir(Expected.squeeze, (x,), R.Tensor((a_main, b_main, 1), dtype="float32")) return gv @Ts.prim_func(private=True) def squeeze(var_rxplaceholder: T.handle, var_T_squeeze: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, T.int64(1), b, T.int64(1)], dtype="float32") - T_squeeze = T.match_buffer(var_T_squeeze, [a, b, T.int64(1)], dtype="float32") - for i0, i1, i2 in T.grid(a, b, T.int64(1)): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_squeeze, T.int64(1), b_squeeze, T.int64(1)], dtype="float32") + T_squeeze = T.match_buffer(var_T_squeeze, [a_squeeze, b_squeeze, T.int64(1)], dtype="float32") + for i0, i1, i2 in T.grid(a_squeeze, b_squeeze, T.int64(1)): with Ts.sblock("T_squeeze"): ax0, ax1, ax2 = Ts.axis.remap("SSS", [i0, i1, i2]) Ts.reads(rxplaceholder[ax0, T.int64(0), ax1, ax2]) @@ -1175,25 +1205,33 @@ def repeat( def test_repeat_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + @I.ir_module class Repeat: @R.function - def main(x: R.Tensor(("a", "b", "c"), "float32")): + def main(x: R.Tensor((a, b, c), "float32")): gv = R.repeat(x, 2, 0) return gv + a_repeat = T.dynamic("a") + b_repeat = T.dynamic("b") + c_repeat = T.dynamic("c") + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + @I.ir_module class Expected: @Ts.prim_func(private=True) def repeat(var_rxplaceholder: T.handle, var_T_repeat: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (a, b, c)) - T_repeat = T.match_buffer(var_T_repeat, (T.int64(2) * a, b, c)) + rxplaceholder = T.match_buffer(var_rxplaceholder, (a_repeat, b_repeat, c_repeat)) + T_repeat = T.match_buffer(var_T_repeat, (T.int64(2) * a_repeat, b_repeat, c_repeat)) # with Ts.sblock("root"): - for ax0, ax1, ax2 in T.grid(a * T.int64(2), b, c): + for ax0, ax1, ax2 in T.grid(a_repeat * T.int64(2), b_repeat, c_repeat): with Ts.sblock("T_repeat"): v_ax0, v_ax1, v_ax2 = Ts.axis.remap("SSS", [ax0, ax1, ax2]) Ts.reads(rxplaceholder[v_ax0 // T.int64(2), v_ax1, v_ax2]) @@ -1201,11 +1239,8 @@ def repeat(var_rxplaceholder: T.handle, var_T_repeat: T.handle): T_repeat[v_ax0, v_ax1, v_ax2] = rxplaceholder[v_ax0 // T.int64(2), v_ax1, v_ax2] @R.function - def main(x: R.Tensor(("a", "b", "c"), dtype="float32")) -> R.Tensor(("2 * a", "b", "c"), dtype="float32"): - a = T.int64() - b = T.int64() - c = T.int64() - gv = R.call_tir(Expected.repeat, (x,), out_ty=R.Tensor((2 * a, b, c), dtype="float32")) + def main(x: R.Tensor((a_main, b_main, c_main), dtype="float32")) -> R.Tensor((2 * a_main, b_main, c_main), dtype="float32"): + gv = R.call_tir(Expected.repeat, (x,), out_ty=R.Tensor((2 * a_main, b_main, c_main), dtype="float32")) return gv # fmt: on @@ -1247,37 +1282,42 @@ def main(x: R.Tensor((3, 2, 3), dtype="float32")) -> R.Tensor((2, 3, 4, 9), dtyp def test_tile_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + @I.ir_module class Tile: @R.function - def main(x: R.Tensor(("a", "b", "c"), "float32")): + def main(x: R.Tensor((a, b, c), "float32")): gv = R.tile(x, (2, 1, 2, 3)) return gv + a_tile = T.dynamic("a") + b_tile = T.dynamic("b") + c_tile = T.dynamic("c") + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + @I.ir_module class Expected: @Ts.prim_func(private=True) def tile(var_rxplaceholder: T.handle, var_T_tile: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (a, b, c)) - T_tile = T.match_buffer(var_T_tile, (T.int64(2), a, b * T.int64(2), c * T.int64(3))) + rxplaceholder = T.match_buffer(var_rxplaceholder, (a_tile, b_tile, c_tile)) + T_tile = T.match_buffer(var_T_tile, (T.int64(2), a_tile, b_tile * T.int64(2), c_tile * T.int64(3))) # with Ts.sblock("root"): - for ax0, ax1, ax2, ax3 in T.grid(T.int64(2), a, b * T.int64(2), c * T.int64(3)): + for ax0, ax1, ax2, ax3 in T.grid(T.int64(2), a_tile, b_tile * T.int64(2), c_tile * T.int64(3)): with Ts.sblock("T_tile"): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) - Ts.reads(rxplaceholder[v_ax1 % a, v_ax2 % b, v_ax3 % c]) + Ts.reads(rxplaceholder[v_ax1 % a_tile, v_ax2 % b_tile, v_ax3 % c_tile]) Ts.writes(T_tile[v_ax0, v_ax1, v_ax2, v_ax3]) - T_tile[v_ax0, v_ax1, v_ax2, v_ax3] = rxplaceholder[v_ax1 % a, v_ax2 % b, v_ax3 % c] + T_tile[v_ax0, v_ax1, v_ax2, v_ax3] = rxplaceholder[v_ax1 % a_tile, v_ax2 % b_tile, v_ax3 % c_tile] @R.function - def main(x: R.Tensor(("a", "b", "c"), dtype="float32")) -> R.Tensor((2, "a", "b * 2", "c * 3"), dtype="float32"): - a = T.int64() - b = T.int64() - c = T.int64() - gv = R.call_tir(Expected.tile, (x,), out_ty=R.Tensor((2, a, b * 2, c * 3), dtype="float32")) + def main(x: R.Tensor((a_main, b_main, c_main), dtype="float32")) -> R.Tensor((2, a_main, b_main * 2, c_main * 3), dtype="float32"): + gv = R.call_tir(Expected.tile, (x,), out_ty=R.Tensor((2, a_main, b_main * 2, c_main * 3), dtype="float32")) return gv # fmt: on mod = LegalizeOps()(Tile) @@ -1324,38 +1364,43 @@ def flip( def test_flip_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + @I.ir_module class Flip: @R.function - def main(x: R.Tensor(("a", "b"), "float32")): + def main(x: R.Tensor((a, b), "float32")): gv = R.flip(x, axis=1) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + a_flip = T.dynamic("a") + b_flip = T.dynamic("b") + @I.ir_module class Expected: @R.function def main( - x: R.Tensor(("a", "b"), dtype="float32") - ) -> R.Tensor(("a", "b"), dtype="float32"): - a = T.int64() - b = T.int64() + x: R.Tensor((a_main, b_main), dtype="float32") + ) -> R.Tensor((a_main, b_main), dtype="float32"): cls = Expected - gv = R.call_tir(cls.flip, (x,), out_ty=R.Tensor((a, b), dtype="float32")) + gv = R.call_tir(cls.flip, (x,), out_ty=R.Tensor((a_main, b_main), dtype="float32")) return gv @Ts.prim_func(private=True) def flip(var_rxplaceholder: T.handle, var_T_reverse_sequence: T.handle): T.func_attr({"tirx.noalias": True}) - a, b = T.int64(), T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (a, b)) - T_reverse_sequence = T.match_buffer(var_T_reverse_sequence, (a, b)) - for ax0, ax1 in T.grid(a, b): + rxplaceholder = T.match_buffer(var_rxplaceholder, (a_flip, b_flip)) + T_reverse_sequence = T.match_buffer(var_T_reverse_sequence, (a_flip, b_flip)) + for ax0, ax1 in T.grid(a_flip, b_flip): with Ts.sblock("T_reverse_sequence"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) - Ts.reads(rxplaceholder[v_ax0, b - v_ax1 - T.int64(1)]) + Ts.reads(rxplaceholder[v_ax0, b_flip - v_ax1 - T.int64(1)]) Ts.writes(T_reverse_sequence[v_ax0, v_ax1]) T_reverse_sequence[v_ax0, v_ax1] = rxplaceholder[ - v_ax0, b - v_ax1 - T.int64(1) + v_ax0, b_flip - v_ax1 - T.int64(1) ] # fmt: on @@ -1519,12 +1564,26 @@ def main( def test_scatter_elements_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class ScatterElements: @R.function - def main(x: R.Tensor(("a", "b"), "float32"), indices:R.Tensor(("m", "n"), "int64"), updates:R.Tensor(("m","n"), "float32")): + def main(x: R.Tensor((a, b), "float32"), indices:R.Tensor((m, n), "int64"), updates:R.Tensor((m,n), "float32")): gv = R.scatter_elements(x, indices, updates, axis=1) return gv + a_scatter_elements = T.dynamic("a") + b_scatter_elements = T.dynamic("b") + m_scatter_elements = T.dynamic("m") + n_scatter_elements = T.dynamic("n") + a_main = T.dynamic("a") + b_main = T.dynamic("b") + m_main = T.dynamic("m") + n_main = T.dynamic("n") + @I.ir_module class Expected: @Ts.prim_func(private=True) @@ -1535,71 +1594,65 @@ def scatter_elements( var_scatter_elements_generic: T.handle, ): T.func_attr({"tirx.noalias": True}) - a, b = T.int64(), T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (a, b), offset_factor=1) - m, n = T.int64(), T.int64() + rxplaceholder = T.match_buffer(var_rxplaceholder, (a_scatter_elements, b_scatter_elements), offset_factor=1) rxplaceholder_1 = T.match_buffer( - var_rxplaceholder_1, (m, n), "int64", offset_factor=1 + var_rxplaceholder_1, (m_scatter_elements, n_scatter_elements), "int64", offset_factor=1 ) - rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, (m, n), offset_factor=1) - out_buf = T.match_buffer(var_scatter_elements_generic, (a, b)) + rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, (m_scatter_elements, n_scatter_elements), offset_factor=1) + out_buf = T.match_buffer(var_scatter_elements_generic, (a_scatter_elements, b_scatter_elements)) with Ts.sblock("scatter_elements_generic"): T.attr(0, "pragma_scope", "seq") - for i in T.parallel(a * b): - out_buf[i // b, i % b] = rxplaceholder[i // b, i % b] - for fused in T.parallel(m): - for k in range(n): + for i in T.parallel(a_scatter_elements * b_scatter_elements): + out_buf[i // b_scatter_elements, i % b_scatter_elements] = rxplaceholder[i // b_scatter_elements, i % b_scatter_elements] + for fused in T.parallel(m_scatter_elements): + for k in range(n_scatter_elements): out_buf[ ( - fused * b + fused * b_scatter_elements + ( rxplaceholder_1[ - (fused * n + k) // n, (fused * n + k) % n + (fused * n_scatter_elements + k) // n_scatter_elements, (fused * n_scatter_elements + k) % n_scatter_elements ] + T.Cast( "int64", rxplaceholder_1[ - (fused * n + k) // n, (fused * n + k) % n + (fused * n_scatter_elements + k) // n_scatter_elements, (fused * n_scatter_elements + k) % n_scatter_elements ] < T.int64(0), ) - * b + * b_scatter_elements ) ) - // b, + // b_scatter_elements, ( - fused * b + fused * b_scatter_elements + ( rxplaceholder_1[ - (fused * n + k) // n, (fused * n + k) % n + (fused * n_scatter_elements + k) // n_scatter_elements, (fused * n_scatter_elements + k) % n_scatter_elements ] + T.Cast( "int64", rxplaceholder_1[ - (fused * n + k) // n, (fused * n + k) % n + (fused * n_scatter_elements + k) // n_scatter_elements, (fused * n_scatter_elements + k) % n_scatter_elements ] < T.int64(0), ) - * b + * b_scatter_elements ) ) - % b, - ] = rxplaceholder_2[(fused * n + k) // n, (fused * n + k) % n] + % b_scatter_elements, + ] = rxplaceholder_2[(fused * n_scatter_elements + k) // n_scatter_elements, (fused * n_scatter_elements + k) % n_scatter_elements] @R.function def main( - x: R.Tensor(("a", "b"), dtype="float32"), - indices: R.Tensor(("m", "n"), dtype="int64"), - updates: R.Tensor(("m", "n"), dtype="float32"), - ) -> R.Tensor(("a", "b"), dtype="float32"): - a = T.int64() - b = T.int64() - m = T.int64() - n = T.int64() + x: R.Tensor((a_main, b_main), dtype="float32"), + indices: R.Tensor((m_main, n_main), dtype="int64"), + updates: R.Tensor((m_main, n_main), dtype="float32"), + ) -> R.Tensor((a_main, b_main), dtype="float32"): gv = R.call_tir( Expected.scatter_elements, (x, indices, updates), - out_ty=R.Tensor((a, b), dtype="float32"), + out_ty=R.Tensor((a_main, b_main), dtype="float32"), ) return gv # fmt: on @@ -1714,38 +1767,45 @@ def test_layout_transform_symbolic(): pad_value = 2 # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + @I.ir_module class LayoutTransform: @R.function - def main(x: R.Tensor(("a", "b", "c"), "float32")): + def main(x: R.Tensor((a, b, c), "float32")): gv = R.layout_transform( x, index_map=transformation, pad_value=pad_value ) return gv + a_te_layout_transform_with_pad = T.dynamic("a") + b_te_layout_transform_with_pad = T.dynamic("b") + c_te_layout_transform_with_pad = T.dynamic("c") + a_main = T.dynamic("a") + c_main = T.dynamic("c") + b_main = T.dynamic("b") + @I.ir_module class Expected: @Ts.prim_func(private=True) def te_layout_transform_with_pad(var_A: T.handle, var_te_layout_transform_with_pad: T.handle): T.func_attr({"tirx.noalias": True}) - a, b, c = T.int64(), T.int64(), T.int64() - A = T.match_buffer(var_A, (a, b, c)) - te_layout_transform_with_pad_1 = T.match_buffer(var_te_layout_transform_with_pad, (a, c, (b - b % T.int64(-3)) // T.int64(3), T.int64(3))) + A = T.match_buffer(var_A, (a_te_layout_transform_with_pad, b_te_layout_transform_with_pad, c_te_layout_transform_with_pad)) + te_layout_transform_with_pad_1 = T.match_buffer(var_te_layout_transform_with_pad, (a_te_layout_transform_with_pad, c_te_layout_transform_with_pad, (b_te_layout_transform_with_pad - b_te_layout_transform_with_pad % T.int64(-3)) // T.int64(3), T.int64(3))) # with Ts.sblock("root"): - for axis0, axis1, axis2, axis3 in T.grid(a, c, (b - b % T.int64(-3)) // T.int64(3), T.int64(3)): + for axis0, axis1, axis2, axis3 in T.grid(a_te_layout_transform_with_pad, c_te_layout_transform_with_pad, (b_te_layout_transform_with_pad - b_te_layout_transform_with_pad % T.int64(-3)) // T.int64(3), T.int64(3)): with Ts.sblock("te_layout_transform_with_pad_with_pad"): v_axis0, v_axis1, v_axis2, v_axis3 = Ts.axis.remap("SSSS", [axis0, axis1, axis2, axis3]) Ts.reads(A[v_axis0, v_axis2 * T.int64(3) + v_axis3, v_axis1]) Ts.writes(te_layout_transform_with_pad_1[v_axis0, v_axis1, v_axis2, v_axis3]) - te_layout_transform_with_pad_1[v_axis0, v_axis1, v_axis2, v_axis3] = T.if_then_else(b % T.int64(-3) < T.int64(0) and v_axis2 == b // T.int64(3) and b % T.int64(3) <= v_axis3, T.float32(2), A[v_axis0, v_axis2 * T.int64(3) + v_axis3, v_axis1]) + te_layout_transform_with_pad_1[v_axis0, v_axis1, v_axis2, v_axis3] = T.if_then_else(b_te_layout_transform_with_pad % T.int64(-3) < T.int64(0) and v_axis2 == b_te_layout_transform_with_pad // T.int64(3) and b_te_layout_transform_with_pad % T.int64(3) <= v_axis3, T.float32(2), A[v_axis0, v_axis2 * T.int64(3) + v_axis3, v_axis1]) @R.function - def main(x: R.Tensor(("a", "b", "c"), dtype="float32")) -> R.Tensor(("a", "c", "(b - b % -3) // 3", 3), dtype="float32"): - a = T.int64() - c = T.int64() - b = T.int64() + def main(x: R.Tensor((a_main, b_main, c_main), dtype="float32")) -> R.Tensor((a_main, c_main, (b_main - b_main % -3) // 3, 3), dtype="float32"): cls = Expected - gv = R.call_tir(cls.te_layout_transform_with_pad, (x,), out_ty=R.Tensor((a, c, (b - b % -3) // 3, 3), dtype="float32")) + gv = R.call_tir(cls.te_layout_transform_with_pad, (x,), out_ty=R.Tensor((a_main, c_main, (b_main - b_main % -3) // 3, 3), dtype="float32")) return gv # fmt: on diff --git a/tests/python/relax/test_transform_legalize_ops_nn.py b/tests/python/relax/test_transform_legalize_ops_nn.py index 3f4a26642bb7..d16caee1928b 100644 --- a/tests/python/relax/test_transform_legalize_ops_nn.py +++ b/tests/python/relax/test_transform_legalize_ops_nn.py @@ -154,46 +154,52 @@ def conv1d(rxplaceholder: T.Buffer((T.int64(2), T.int64(28), T.int64(128)), "flo def test_conv1d_symbolic(): # fmt: off + n = T.dynamic("n") + w = T.dynamic("w") + f = T.dynamic("f") + kw = T.dynamic("kw") + c = T.dynamic("c") + @tvm.script.ir_module class Conv1d: @R.function - def main(x: R.Tensor(("n", "c", "w"), "float32"), kernel: R.Tensor(("f", "c", "kw"), "float32")) -> R.Tensor(("n", "f", "w - kw + 1"), "float32"): - n = T.int64() - w = T.int64() - f = T.int64() - kw = T.int64() + def main(x: R.Tensor((n, c, w), "float32"), kernel: R.Tensor((f, c, kw), "float32")) -> R.Tensor((n, f, w - kw + 1), "float32"): gv: R.Tensor((n, f, w - kw + 1), "float32") = R.nn.conv1d(x, kernel) return gv + n_main = T.dynamic("n") + f_main = T.dynamic("f") + w_main = T.dynamic("w") + kw_main = T.dynamic("kw") + c_main = T.dynamic("c") + n_conv1d = T.dynamic("n") + c_conv1d = T.dynamic("c") + w_conv1d = T.dynamic("w") + f_conv1d = T.dynamic("f") + kw_conv1d = T.dynamic("kw") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("n", "c", "w"), dtype="float32"), kernel: R.Tensor(("f", "c", "kw"), dtype="float32")) -> R.Tensor(("n", "f", "w - kw + 1"), dtype="float32"): - n = T.int64() - f = T.int64() - w = T.int64() - kw = T.int64() - c = T.int64() - gv = R.call_tir(Expected.conv1d, (x, kernel), out_ty=R.Tensor((n, f, w + 1 - kw), dtype="float32")) + def main(x: R.Tensor((n_main, c_main, w_main), dtype="float32"), kernel: R.Tensor((f_main, c_main, kw_main), dtype="float32")) -> R.Tensor((n_main, f_main, w_main - kw_main + 1), dtype="float32"): + gv = R.call_tir(Expected.conv1d, (x, kernel), out_ty=R.Tensor((n_main, f_main, w_main + 1 - kw_main), dtype="float32")) return gv @Ts.prim_func(private=True) def conv1d(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_conv1d_ncw: T.handle): T.func_attr({"tirx.noalias": True}) - n, c, w = T.int64(), T.int64(), T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (n, c, w)) - f, kw = T.int64(), T.int64() - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (f, c, kw)) - conv1d_ncw = T.match_buffer(var_conv1d_ncw, (n, f, w + T.int64(1) - kw)) + rxplaceholder = T.match_buffer(var_rxplaceholder, (n_conv1d, c_conv1d, w_conv1d)) + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (f_conv1d, c_conv1d, kw_conv1d)) + conv1d_ncw = T.match_buffer(var_conv1d_ncw, (n_conv1d, f_conv1d, w_conv1d + T.int64(1) - kw_conv1d)) # with Ts.sblock("root"): - pad_temp = Ts.sblock_alloc_buffer((n, c, w)) - for i0, i1, i2 in T.grid(n, c, w): + pad_temp = Ts.sblock_alloc_buffer((n_conv1d, c_conv1d, w_conv1d)) + for i0, i1, i2 in T.grid(n_conv1d, c_conv1d, w_conv1d): with Ts.sblock("pad_temp"): v_i0, v_i1, v_i2 = Ts.axis.remap("SSS", [i0, i1, i2]) Ts.reads(rxplaceholder[v_i0, v_i1, v_i2]) Ts.writes(pad_temp[v_i0, v_i1, v_i2]) pad_temp[v_i0, v_i1, v_i2] = rxplaceholder[v_i0, v_i1, v_i2] - for nn, ff, yy, rc, ry in T.grid(n, f, w + T.int64(1) - kw, c, kw): + for nn, ff, yy, rc, ry in T.grid(n_conv1d, f_conv1d, w_conv1d + T.int64(1) - kw_conv1d, c_conv1d, kw_conv1d): with Ts.sblock("conv1d_ncw"): v_nn, v_ff, v_yy, v_rc, v_ry = Ts.axis.remap("SSSRR", [nn, ff, yy, rc, ry]) Ts.reads(pad_temp[v_nn, v_rc, v_yy + v_ry], rxplaceholder_1[v_ff, v_rc, v_ry]) @@ -236,9 +242,9 @@ def conv1d_transpose(x: T.Buffer((T.int64(2), T.int64(128), T.int64(28)), "float with Ts.sblock("kernel"): v_o, v_i, v_w = Ts.axis.remap("SSS", [o, i, w_1]) kernel[v_o, v_i, v_w] = w[v_i, v_o, T.int64(2) - v_w] - for b, c, w_1, dc, dw in T.grid(T.int64(2), T.int64(128), T.int64(56), T.int64(16), T.int64(3)): + for b_index, c_index, w_1, dc, dw in T.grid(T.int64(2), T.int64(128), T.int64(56), T.int64(16), T.int64(3)): with Ts.sblock("compute"): - v_b, v_c, v_w, v_dc, v_dw = Ts.axis.remap("SSSRR", [b, c, w_1, dc, dw]) + v_b, v_c, v_w, v_dc, v_dw = Ts.axis.remap("SSSRR", [b_index, c_index, w_1, dc, dw]) with Ts.init(): compute[v_b, v_c, v_w] = T.float32(0.0) compute[v_b, v_c, v_w] = compute[v_b, v_c, v_w] + data_pad[v_b, v_c // T.int64(16) * T.int64(16) + v_dc, v_w + v_dw] * kernel[v_c % T.int64(16), v_c // T.int64(16) * T.int64(16) + v_dc, v_dw] @@ -376,53 +382,57 @@ def conv2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(28), T.int64(28), T.int6 def test_conv2d_symbolic(): # fmt: off + n = T.dynamic("n") + h = T.dynamic("h") + w = T.dynamic("w") + f = T.dynamic("f") + kh = T.dynamic("kh") + kw = T.dynamic("kw") + c = T.dynamic("c") + @tvm.script.ir_module class Conv2d: @R.function - def main(x: R.Tensor(("n", "c", "h", "w"), "float32"), kernel: R.Tensor(("f", "c", "kh", "kw"), "float32")) -> R.Tensor(("n", "f", "h - kh + 1", "w - kw + 1"), "float32"): - n = T.int64() - h = T.int64() - w = T.int64() - f = T.int64() - kh = T.int64() - kw = T.int64() + def main(x: R.Tensor((n, c, h, w), "float32"), kernel: R.Tensor((f, c, kh, kw), "float32")) -> R.Tensor((n, f, h - kh + 1, w - kw + 1), "float32"): gv: R.Tensor((n, f, h - kh + 1, w - kw + 1), "float32") = R.nn.conv2d(x, kernel) return gv + n_main = T.dynamic("n") + f_main = T.dynamic("f") + h_main = T.dynamic("h") + kh_main = T.dynamic("kh") + w_main = T.dynamic("w") + kw_main = T.dynamic("kw") + c_main = T.dynamic("c") + c_conv2d = T.dynamic("c") + f_conv2d = T.dynamic("f") + h_conv2d = T.dynamic("h") + kh_conv2d = T.dynamic("kh") + kw_conv2d = T.dynamic("kw") + n_conv2d = T.dynamic("n") + w_conv2d = T.dynamic("w") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("n", "c", "h", "w"), "float32"), kernel: R.Tensor(("f", "c", "kh", "kw"), "float32")) -> R.Tensor(("n", "f", "h - kh + 1", "w - kw + 1"), "float32"): - n = T.int64() - f = T.int64() - h = T.int64() - kh = T.int64() - w = T.int64() - kw = T.int64() - gv = R.call_tir(Expected.conv2d, (x, kernel), R.Tensor((n, f, h + 1 - kh, w + 1 - kw), dtype="float32")) + def main(x: R.Tensor((n_main, c_main, h_main, w_main), "float32"), kernel: R.Tensor((f_main, c_main, kh_main, kw_main), "float32")) -> R.Tensor((n_main, f_main, h_main - kh_main + 1, w_main - kw_main + 1), "float32"): + gv = R.call_tir(Expected.conv2d, (x, kernel), R.Tensor((n_main, f_main, h_main + 1 - kh_main, w_main + 1 - kw_main), dtype="float32")) return gv @Ts.prim_func(private=True) def conv2d(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_conv2d_nchw: T.handle): T.func_attr({"tirx.noalias": True}) - c = T.int64() - f = T.int64() - h = T.int64() - kh = T.int64() - kw = T.int64() - n = T.int64() - w = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [n, c, h, w], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [f, c, kh, kw], dtype="float32") - conv2d_nchw = T.match_buffer(var_conv2d_nchw, [n, f, h + T.int64(1) - kh, w + T.int64(1) - kw], dtype="float32") - pad_temp = Ts.sblock_alloc_buffer([n, c, h, w], dtype="float32") - for i0, i1, i2, i3 in T.grid(n, c, h, w): + rxplaceholder = T.match_buffer(var_rxplaceholder, [n_conv2d, c_conv2d, h_conv2d, w_conv2d], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [f_conv2d, c_conv2d, kh_conv2d, kw_conv2d], dtype="float32") + conv2d_nchw = T.match_buffer(var_conv2d_nchw, [n_conv2d, f_conv2d, h_conv2d + T.int64(1) - kh_conv2d, w_conv2d + T.int64(1) - kw_conv2d], dtype="float32") + pad_temp = Ts.sblock_alloc_buffer([n_conv2d, c_conv2d, h_conv2d, w_conv2d], dtype="float32") + for i0, i1, i2, i3 in T.grid(n_conv2d, c_conv2d, h_conv2d, w_conv2d): with Ts.sblock("pad_temp"): i0_1, i1_1, i2_1, i3_1 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[i0_1, i1_1, i2_1, i3_1]) Ts.writes(pad_temp[i0_1, i1_1, i2_1, i3_1]) pad_temp[i0_1, i1_1, i2_1, i3_1] = rxplaceholder[i0_1, i1_1, i2_1, i3_1] - for i0, i1, i2, i3, i4, i5, i6 in T.grid(n, f, h + T.int64(1) - kh, w + T.int64(1) - kw, c, kh, kw): + for i0, i1, i2, i3, i4, i5, i6 in T.grid(n_conv2d, f_conv2d, h_conv2d + T.int64(1) - kh_conv2d, w_conv2d + T.int64(1) - kw_conv2d, c_conv2d, kh_conv2d, kw_conv2d): with Ts.sblock("conv2d_nchw"): nn, ff, yy, xx, rc, ry, rx = Ts.axis.remap("SSSSRRR", [i0, i1, i2, i3, i4, i5, i6]) Ts.reads(pad_temp[nn, rc, yy + ry, xx + rx], rxplaceholder_1[ff, rc, ry, rx]) @@ -438,47 +448,55 @@ def conv2d(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_conv2 def test_conv2d_symbolic_group(): # fmt: off + n = T.dynamic("n") + f = T.dynamic("f") + c = T.dynamic("c") + c_div_8 = T.dynamic("c_div_8") + @tvm.script.ir_module class Conv2d: @R.function - def main(x: R.Tensor(("n", "c", 28, 28), "float32"), w: R.Tensor(("f", "c_div_8", 3, 3), "float32")) -> R.Tensor(("n", "f", 26, 26), "float32"): - n = T.int64() - f = T.int64() + def main(x: R.Tensor((n, c, 28, 28), "float32"), w: R.Tensor((f, c_div_8, 3, 3), "float32")) -> R.Tensor((n, f, 26, 26), "float32"): gv: R.Tensor((n, f, 26, 26), "float32") = R.nn.conv2d(x, w, groups=8) return gv + n_main = T.dynamic("n") + f_main = T.dynamic("f") + c_main = T.dynamic("c") + c_div_8_main = T.dynamic("c_div_8") + n_conv2d = T.dynamic("n") + c_conv2d = T.dynamic("c") + f_conv2d = T.dynamic("f") + c_div_8_conv2d = T.dynamic("c_div_8") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("n", "c", 28, 28), dtype="float32"), w: R.Tensor(("f", "c_div_8", 3, 3), dtype="float32")) -> R.Tensor(("n", "f", 26, 26), dtype="float32"): - n = T.int64() - f = T.int64() - gv = R.call_tir(Expected.conv2d, (x, w), out_ty=R.Tensor((n, f, 26, 26), dtype="float32")) + def main(x: R.Tensor((n_main, c_main, 28, 28), dtype="float32"), w: R.Tensor((f_main, c_div_8_main, 3, 3), dtype="float32")) -> R.Tensor((n_main, f_main, 26, 26), dtype="float32"): + gv = R.call_tir(Expected.conv2d, (x, w), out_ty=R.Tensor((n_main, f_main, 26, 26), dtype="float32")) return gv @Ts.prim_func(private=True) def conv2d(var_x: T.handle, var_w: T.handle, var_group_conv2d_nchw: T.handle): T.func_attr({"tirx.noalias": True}) - n, c = T.int64(), T.int64() - x = T.match_buffer(var_x, (n, c, T.int64(28), T.int64(28))) - f, c_div_8 = T.int64(), T.int64() - w = T.match_buffer(var_w, (f, c_div_8, T.int64(3), T.int64(3))) - group_conv2d_nchw = T.match_buffer(var_group_conv2d_nchw, (n, f, T.int64(26), T.int64(26))) - pad_temp = Ts.sblock_alloc_buffer((n, c, T.int64(28), T.int64(28))) - for i0, i1, i2, i3 in T.grid(n, c, T.int64(28), T.int64(28)): + x = T.match_buffer(var_x, (n_conv2d, c_conv2d, T.int64(28), T.int64(28))) + w = T.match_buffer(var_w, (f_conv2d, c_div_8_conv2d, T.int64(3), T.int64(3))) + group_conv2d_nchw = T.match_buffer(var_group_conv2d_nchw, (n_conv2d, f_conv2d, T.int64(26), T.int64(26))) + pad_temp = Ts.sblock_alloc_buffer((n_conv2d, c_conv2d, T.int64(28), T.int64(28))) + for i0, i1, i2, i3 in T.grid(n_conv2d, c_conv2d, T.int64(28), T.int64(28)): with Ts.sblock("pad_temp"): v_i0, v_i1, v_i2, v_i3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(x[v_i0, v_i1, v_i2, v_i3]) Ts.writes(pad_temp[v_i0, v_i1, v_i2, v_i3]) pad_temp[v_i0, v_i1, v_i2, v_i3] = x[v_i0, v_i1, v_i2, v_i3] - for nn, ff, yy, xx, rc, ry, rx in T.grid(n, f, T.int64(26), T.int64(26), c // T.int64(8), T.int64(3), T.int64(3)): + for nn, ff, yy, xx, rc, ry, rx in T.grid(n_conv2d, f_conv2d, T.int64(26), T.int64(26), c_conv2d // T.int64(8), T.int64(3), T.int64(3)): with Ts.sblock("group_conv2d_nchw"): v_nn, v_ff, v_yy, v_xx, v_rc, v_ry, v_rx = Ts.axis.remap("SSSSRRR", [nn, ff, yy, xx, rc, ry, rx]) - Ts.reads(pad_temp[v_nn, v_ff // (f // T.int64(8)) * (c // T.int64(8)) + v_rc, v_yy + v_ry, v_xx + v_rx], w[v_ff, v_rc, v_ry, v_rx]) + Ts.reads(pad_temp[v_nn, v_ff // (f_conv2d // T.int64(8)) * (c_conv2d // T.int64(8)) + v_rc, v_yy + v_ry, v_xx + v_rx], w[v_ff, v_rc, v_ry, v_rx]) Ts.writes(group_conv2d_nchw[v_nn, v_ff, v_yy, v_xx]) with Ts.init(): group_conv2d_nchw[v_nn, v_ff, v_yy, v_xx] = T.float32(0.0) - group_conv2d_nchw[v_nn, v_ff, v_yy, v_xx] = group_conv2d_nchw[v_nn, v_ff, v_yy, v_xx] + pad_temp[v_nn, v_ff // (f // T.int64(8)) * (c // T.int64(8)) + v_rc, v_yy + v_ry, v_xx + v_rx] * w[v_ff, v_rc, v_ry, v_rx] + group_conv2d_nchw[v_nn, v_ff, v_yy, v_xx] = group_conv2d_nchw[v_nn, v_ff, v_yy, v_xx] + pad_temp[v_nn, v_ff // (f_conv2d // T.int64(8)) * (c_conv2d // T.int64(8)) + v_rc, v_yy + v_ry, v_xx + v_rx] * w[v_ff, v_rc, v_ry, v_rx] # fmt: on mod = LegalizeOps()(Conv2d) @@ -520,15 +538,15 @@ def conv2d_transpose(rxplaceholder: T.Buffer((T.int64(2), T.int64(128), T.int64( Ts.reads(data_dilate[v_i0, v_i1, v_i2 - T.int64(1), v_i3 - T.int64(1)]) Ts.writes(data_pad[v_i0, v_i1, v_i2, v_i3]) data_pad[v_i0, v_i1, v_i2, v_i3] = T.if_then_else(T.int64(1) <= v_i2 and v_i2 < T.int64(56) and T.int64(1) <= v_i3 and v_i3 < T.int64(83), data_dilate[v_i0, v_i1, v_i2 - T.int64(1), v_i3 - T.int64(1)], T.float32(0)) - for i, o, h, w in T.grid(T.int64(16), T.int64(128), T.int64(3), T.int64(3)): + for i, o, h_index, w_index in T.grid(T.int64(16), T.int64(128), T.int64(3), T.int64(3)): with Ts.sblock("kernel_transform"): - v_i, v_o, v_h, v_w = Ts.axis.remap("SSSS", [i, o, h, w]) + v_i, v_o, v_h, v_w = Ts.axis.remap("SSSS", [i, o, h_index, w_index]) Ts.reads(rxplaceholder_1[v_o, v_i, T.int64(2) - v_h, T.int64(2) - v_w]) Ts.writes(kernel_transform[v_i, v_o, v_h, v_w]) kernel_transform[v_i, v_o, v_h, v_w] = rxplaceholder_1[v_o, v_i, T.int64(2) - v_h, T.int64(2) - v_w] - for b, c, h, w, dc, dh, dw in T.grid(T.int64(2), T.int64(128), T.int64(56), T.int64(84), T.int64(16), T.int64(3), T.int64(3)): + for b_index, c_index, h_index, w_index, dc, dh, dw in T.grid(T.int64(2), T.int64(128), T.int64(56), T.int64(84), T.int64(16), T.int64(3), T.int64(3)): with Ts.sblock("compute"): - v_b, v_c, v_h, v_w, v_dc, v_dh, v_dw = Ts.axis.remap("SSSSRRR", [b, c, h, w, dc, dh, dw]) + v_b, v_c, v_h, v_w, v_dc, v_dh, v_dw = Ts.axis.remap("SSSSRRR", [b_index, c_index, h_index, w_index, dc, dh, dw]) Ts.reads(data_pad[v_b, v_c // T.int64(16) * T.int64(16) + v_dc, v_h + v_dh, v_w + v_dw], kernel_transform[v_c % T.int64(16), v_c // T.int64(16) * T.int64(16) + v_dc, v_dh, v_dw]) Ts.writes(compute[v_b, v_c, v_h, v_w]) with Ts.init(): @@ -574,15 +592,15 @@ def conv3d_transpose(x: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(4) Ts.reads(data_dilate[v_i0, v_i1, v_i2 - T.int64(2), v_i3 - T.int64(2), v_i4 - T.int64(2)]) Ts.writes(data_pad[v_i0, v_i1, v_i2, v_i3, v_i4]) data_pad[v_i0, v_i1, v_i2, v_i3, v_i4] = T.if_then_else(T.int64(2) <= v_i2 and v_i2 < T.int64(6) and T.int64(2) <= v_i3 and v_i3 < T.int64(6) and T.int64(2) <= v_i4 and v_i4 < T.int64(6), data_dilate[v_i0, v_i1, v_i2 - T.int64(2), v_i3 - T.int64(2), v_i4 - T.int64(2)], T.float32(0.0)) - for o, i, d, h, w_1 in T.grid(T.int64(4), T.int64(3), T.int64(3), T.int64(3), T.int64(3)): + for o, i, d, h_index, w_1 in T.grid(T.int64(4), T.int64(3), T.int64(3), T.int64(3), T.int64(3)): with Ts.sblock("kernel_transform"): - v_o, v_i, v_d, v_h, v_w = Ts.axis.remap("SSSSS", [o, i, d, h, w_1]) + v_o, v_i, v_d, v_h, v_w = Ts.axis.remap("SSSSS", [o, i, d, h_index, w_1]) Ts.reads(w[v_i, v_o, T.int64(2) - v_d, T.int64(2) - v_h, T.int64(2) - v_w]) Ts.writes(kernel_transform[v_o, v_i, v_d, v_h, v_w]) kernel_transform[v_o, v_i, v_d, v_h, v_w] = w[v_i, v_o, T.int64(2) - v_d, T.int64(2) - v_h, T.int64(2) - v_w] - for b, c, d, h, w_1, dc, dd, dh, dw in T.grid(T.int64(2), T.int64(4), T.int64(6), T.int64(6), T.int64(6), T.int64(3), T.int64(3), T.int64(3), T.int64(3)): + for b_index, c_index, d, h_index, w_1, dc, dd, dh, dw in T.grid(T.int64(2), T.int64(4), T.int64(6), T.int64(6), T.int64(6), T.int64(3), T.int64(3), T.int64(3), T.int64(3)): with Ts.sblock("compute"): - v_b, v_c, v_d, v_h, v_w, v_dc, v_dd, v_dh, v_dw = Ts.axis.remap("SSSSSRRRR", [b, c, d, h, w_1, dc, dd, dh, dw]) + v_b, v_c, v_d, v_h, v_w, v_dc, v_dd, v_dh, v_dw = Ts.axis.remap("SSSSSRRRR", [b_index, c_index, d, h_index, w_1, dc, dd, dh, dw]) Ts.reads(data_pad[v_b, v_dc, v_d + v_dd, v_h + v_dh, v_w + v_dw], kernel_transform[v_c, v_dc, v_dd, v_dh, v_dw]) Ts.writes(compute[v_b, v_c, v_d, v_h, v_w]) with Ts.init(): @@ -628,15 +646,15 @@ def conv3d_transpose(x: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(4) Ts.reads(data_dilate[v_i0, v_i1, v_i2 - T.int64(2), v_i3 - T.int64(2), v_i4 - T.int64(2)]) Ts.writes(data_pad[v_i0, v_i1, v_i2, v_i3, v_i4]) data_pad[v_i0, v_i1, v_i2, v_i3, v_i4] = T.if_then_else(T.int64(2) <= v_i2 and v_i2 < T.int64(6) and T.int64(2) <= v_i3 and v_i3 < T.int64(6) and T.int64(2) <= v_i4 and v_i4 < T.int64(6), data_dilate[v_i0, v_i1, v_i2 - T.int64(2), v_i3 - T.int64(2), v_i4 - T.int64(2)], T.float32(0.0)) - for o, i, d, h, w_1 in T.grid(T.int64(4), T.int64(3), T.int64(3), T.int64(3), T.int64(3)): + for o, i, d, h_index, w_1 in T.grid(T.int64(4), T.int64(3), T.int64(3), T.int64(3), T.int64(3)): with Ts.sblock("kernel_transform"): - v_o, v_i, v_d, v_h, v_w = Ts.axis.remap("SSSSS", [o, i, d, h, w_1]) + v_o, v_i, v_d, v_h, v_w = Ts.axis.remap("SSSSS", [o, i, d, h_index, w_1]) Ts.reads(w[v_i, v_o, T.int64(2) - v_d, T.int64(2) - v_h, T.int64(2) - v_w]) Ts.writes(kernel_transform[v_o, v_i, v_d, v_h, v_w]) kernel_transform[v_o, v_i, v_d, v_h, v_w] = w[v_i, v_o, T.int64(2) - v_d, T.int64(2) - v_h, T.int64(2) - v_w] - for b, c, d, h, w_1, dc, dd, dh, dw in T.grid(T.int64(2), T.int64(4), T.int64(6), T.int64(6), T.int64(6), T.int64(3), T.int64(3), T.int64(3), T.int64(3)): + for b_index, c_index, d, h_index, w_1, dc, dd, dh, dw in T.grid(T.int64(2), T.int64(4), T.int64(6), T.int64(6), T.int64(6), T.int64(3), T.int64(3), T.int64(3), T.int64(3)): with Ts.sblock("compute"): - v_b, v_c, v_d, v_h, v_w, v_dc, v_dd, v_dh, v_dw = Ts.axis.remap("SSSSSRRRR", [b, c, d, h, w_1, dc, dd, dh, dw]) + v_b, v_c, v_d, v_h, v_w, v_dc, v_dd, v_dh, v_dw = Ts.axis.remap("SSSSSRRRR", [b_index, c_index, d, h_index, w_1, dc, dd, dh, dw]) Ts.reads(data_pad[v_b, v_dc, v_d + v_dd, v_h + v_dh, v_w + v_dw], kernel_transform[v_c, v_dc, v_dd, v_dh, v_dw]) Ts.writes(compute[v_b, v_c, v_d, v_h, v_w]) with Ts.init(): @@ -683,15 +701,15 @@ def conv2d_transpose(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(28 Ts.reads(data_dilate[v_i0, v_i1, v_i2 - T.int64(2), v_i3 - T.int64(2)]) Ts.writes(data_pad[v_i0, v_i1, v_i2, v_i3]) data_pad[v_i0, v_i1, v_i2, v_i3] = T.if_then_else(T.int64(2) <= v_i2 and v_i2 < T.int64(30) and T.int64(2) <= v_i3 and v_i3 < T.int64(30), data_dilate[v_i0, v_i1, v_i2 - T.int64(2), v_i3 - T.int64(2)], T.float32(0)) - for o, i, h, w in T.grid(T.int64(4), T.int64(3), T.int64(3), T.int64(3)): + for o, i, h_index, w_index in T.grid(T.int64(4), T.int64(3), T.int64(3), T.int64(3)): with Ts.sblock("kernel_transform"): - v_o, v_i, v_h, v_w = Ts.axis.remap("SSSS", [o, i, h, w]) + v_o, v_i, v_h, v_w = Ts.axis.remap("SSSS", [o, i, h_index, w_index]) Ts.reads(rxplaceholder_1[v_i, v_o, T.int64(2) - v_h, T.int64(2) - v_w]) Ts.writes(kernel_transform[v_o, v_i, v_h, v_w]) kernel_transform[v_o, v_i, v_h, v_w] = rxplaceholder_1[v_i, v_o, T.int64(2) - v_h, T.int64(2) - v_w] - for b, c, h, w, dc, dh, dw in T.grid(T.int64(2), T.int64(4), T.int64(30), T.int64(30), T.int64(3), T.int64(3), T.int64(3)): + for b_index, c_index, h_index, w_index, dc, dh, dw in T.grid(T.int64(2), T.int64(4), T.int64(30), T.int64(30), T.int64(3), T.int64(3), T.int64(3)): with Ts.sblock("compute"): - v_b, v_c, v_h, v_w, v_dc, v_dh, v_dw = Ts.axis.remap("SSSSRRR", [b, c, h, w, dc, dh, dw]) + v_b, v_c, v_h, v_w, v_dc, v_dh, v_dw = Ts.axis.remap("SSSSRRR", [b_index, c_index, h_index, w_index, dc, dh, dw]) Ts.reads(data_pad[v_b, v_dc, v_h + v_dh, v_w + v_dw], kernel_transform[v_c, v_dc, v_dh, v_dw]) Ts.writes(compute[v_b, v_c, v_h, v_w]) with Ts.init(): @@ -705,65 +723,74 @@ def conv2d_transpose(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(28 def test_conv2d_transpose_symbolic(): # fmt: off + n = T.dynamic("n") + c = T.dynamic("c") + h = T.dynamic("h") + w = T.dynamic("w") + f = T.dynamic("f") + kh = T.dynamic("kh") + kw = T.dynamic("kw") + @tvm.script.ir_module class Conv2dTranspose: @R.function - def main(x: R.Tensor(("n", "c", "h", "w"), "float32"), kernel: R.Tensor(("f", "c", "kh", "kw"), "float32")): + def main(x: R.Tensor((n, c, h, w), "float32"), kernel: R.Tensor((f, c, kh, kw), "float32")): gv = R.nn.conv2d_transpose(x, kernel, strides=(3, 3)) return gv + n_main = T.dynamic("n") + c_main = T.dynamic("c") + h_main = T.dynamic("h") + kh_main = T.dynamic("kh") + w_main = T.dynamic("w") + kw_main = T.dynamic("kw") + f_main = T.dynamic("f") + n_conv2d_transpose = T.dynamic("n") + c_conv2d_transpose = T.dynamic("c") + h_conv2d_transpose = T.dynamic("h") + w_conv2d_transpose = T.dynamic("w") + f_conv2d_transpose = T.dynamic("f") + kh_conv2d_transpose = T.dynamic("kh") + kw_conv2d_transpose = T.dynamic("kw") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(("n", "c", "h", "w"), dtype="float32"), kernel: R.Tensor(("f", "c", "kh", "kw"), dtype="float32")) -> R.Tensor(("n", "c", "h * 3 + kh - 3", "w * 3 + kw - 3"), dtype="float32"): - n = T.int64() - c = T.int64() - h = T.int64() - kh = T.int64() - w = T.int64() - kw = T.int64() - f = T.int64() - gv = R.call_tir(Expected.conv2d_transpose, (x, kernel), out_ty=R.Tensor((n, c, h * 3 + kh - 3, w * 3 + kw - 3), dtype="float32")) + def main(x: R.Tensor((n_main, c_main, h_main, w_main), dtype="float32"), kernel: R.Tensor((f_main, c_main, kh_main, kw_main), dtype="float32")) -> R.Tensor((n_main, c_main, h_main * 3 + kh_main - 3, w_main * 3 + kw_main - 3), dtype="float32"): + gv = R.call_tir(Expected.conv2d_transpose, (x, kernel), out_ty=R.Tensor((n_main, c_main, h_main * 3 + kh_main - 3, w_main * 3 + kw_main - 3), dtype="float32")) return gv @Ts.prim_func(private=True) def conv2d_transpose(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_compute: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() - c = T.int64() - h = T.int64() - w = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (n, c, h, w)) - f = T.int64() - kh = T.int64() - kw = T.int64() - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (f, c, kh, kw)) - compute = T.match_buffer(var_compute, (n, c, h * T.int64(3) + kh - T.int64(3), w * T.int64(3) + kw - T.int64(3))) + rxplaceholder = T.match_buffer(var_rxplaceholder, (n_conv2d_transpose, c_conv2d_transpose, h_conv2d_transpose, w_conv2d_transpose)) + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (f_conv2d_transpose, c_conv2d_transpose, kh_conv2d_transpose, kw_conv2d_transpose)) + compute = T.match_buffer(var_compute, (n_conv2d_transpose, c_conv2d_transpose, h_conv2d_transpose * T.int64(3) + kh_conv2d_transpose - T.int64(3), w_conv2d_transpose * T.int64(3) + kw_conv2d_transpose - T.int64(3))) # with Ts.sblock("root"): - data_dilate = Ts.sblock_alloc_buffer((n, c, h * T.int64(3) - T.int64(2), w * T.int64(3) - T.int64(2))) - data_pad = Ts.sblock_alloc_buffer((n, c, h * T.int64(3) + kh * T.int64(2) - T.int64(4), w * T.int64(3) + kw * T.int64(2) - T.int64(4))) - kernel_transform = Ts.sblock_alloc_buffer((c, c, kh, kw)) - for i0, i1, i2, i3 in T.grid(n, c, h * T.int64(3) - T.int64(2), w * T.int64(3) - T.int64(2)): + data_dilate = Ts.sblock_alloc_buffer((n_conv2d_transpose, c_conv2d_transpose, h_conv2d_transpose * T.int64(3) - T.int64(2), w_conv2d_transpose * T.int64(3) - T.int64(2))) + data_pad = Ts.sblock_alloc_buffer((n_conv2d_transpose, c_conv2d_transpose, h_conv2d_transpose * T.int64(3) + kh_conv2d_transpose * T.int64(2) - T.int64(4), w_conv2d_transpose * T.int64(3) + kw_conv2d_transpose * T.int64(2) - T.int64(4))) + kernel_transform = Ts.sblock_alloc_buffer((c_conv2d_transpose, c_conv2d_transpose, kh_conv2d_transpose, kw_conv2d_transpose)) + for i0, i1, i2, i3 in T.grid(n_conv2d_transpose, c_conv2d_transpose, h_conv2d_transpose * T.int64(3) - T.int64(2), w_conv2d_transpose * T.int64(3) - T.int64(2)): with Ts.sblock("data_dilate"): v_i0, v_i1, v_i2, v_i3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[v_i0, v_i1, v_i2 // T.int64(3), v_i3 // T.int64(3)]) Ts.writes(data_dilate[v_i0, v_i1, v_i2, v_i3]) data_dilate[v_i0, v_i1, v_i2, v_i3] = T.if_then_else(v_i2 % T.int64(3) == T.int64(0) and v_i3 % T.int64(3) == T.int64(0), rxplaceholder[v_i0, v_i1, v_i2 // T.int64(3), v_i3 // T.int64(3)], T.float32(0)) - for i0, i1, i2, i3 in T.grid(n, c, h * T.int64(3) + kh * T.int64(2) - T.int64(4), w * T.int64(3) + kw * T.int64(2) - T.int64(4)): + for i0, i1, i2, i3 in T.grid(n_conv2d_transpose, c_conv2d_transpose, h_conv2d_transpose * T.int64(3) + kh_conv2d_transpose * T.int64(2) - T.int64(4), w_conv2d_transpose * T.int64(3) + kw_conv2d_transpose * T.int64(2) - T.int64(4)): with Ts.sblock("data_pad"): v_i0, v_i1, v_i2, v_i3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) - Ts.reads(data_dilate[v_i0, v_i1, v_i2 + T.int64(1) - kh, v_i3 + T.int64(1) - kw]) + Ts.reads(data_dilate[v_i0, v_i1, v_i2 + T.int64(1) - kh_conv2d_transpose, v_i3 + T.int64(1) - kw_conv2d_transpose]) Ts.writes(data_pad[v_i0, v_i1, v_i2, v_i3]) - data_pad[v_i0, v_i1, v_i2, v_i3] = T.if_then_else(kh <= v_i2 + T.int64(1) and v_i2 + T.int64(3)< h * T.int64(3) + kh and kw <= v_i3 + T.int64(1) and v_i3 + T.int64(3) < w * T.int64(3) + kw , data_dilate[v_i0, v_i1, v_i2 + T.int64(1) - kh, v_i3 + T.int64(1) - kw], T.float32(0)) - for o, i, h_1, w_1 in T.grid(c, c, kh, kw): + data_pad[v_i0, v_i1, v_i2, v_i3] = T.if_then_else(kh_conv2d_transpose <= v_i2 + T.int64(1) and v_i2 + T.int64(3)< h_conv2d_transpose * T.int64(3) + kh_conv2d_transpose and kw_conv2d_transpose <= v_i3 + T.int64(1) and v_i3 + T.int64(3) < w_conv2d_transpose * T.int64(3) + kw_conv2d_transpose , data_dilate[v_i0, v_i1, v_i2 + T.int64(1) - kh_conv2d_transpose, v_i3 + T.int64(1) - kw_conv2d_transpose], T.float32(0)) + for o, i, h_1, w_1 in T.grid(c_conv2d_transpose, c_conv2d_transpose, kh_conv2d_transpose, kw_conv2d_transpose): with Ts.sblock("kernel_transform"): v_o, v_i, v_h, v_w = Ts.axis.remap("SSSS", [o, i, h_1, w_1]) - Ts.reads(rxplaceholder_1[v_i, v_o, kh - v_h - T.int64(1), kw - v_w - T.int64(1)]) + Ts.reads(rxplaceholder_1[v_i, v_o, kh_conv2d_transpose - v_h - T.int64(1), kw_conv2d_transpose - v_w - T.int64(1)]) Ts.writes(kernel_transform[v_o, v_i, v_h, v_w]) - kernel_transform[v_o, v_i, v_h, v_w] = rxplaceholder_1[v_i, v_o, kh - v_h - T.int64(1), kw - v_w - T.int64(1)] - for b, c_1, h_1, w_1, dc, dh, dw in T.grid(n, c, h * T.int64(3) + kh - T.int64(3), w * T.int64(3) + kw - T.int64(3), c, kh, kw): + kernel_transform[v_o, v_i, v_h, v_w] = rxplaceholder_1[v_i, v_o, kh_conv2d_transpose - v_h - T.int64(1), kw_conv2d_transpose - v_w - T.int64(1)] + for b_index, c_1, h_1, w_1, dc, dh, dw in T.grid(n_conv2d_transpose, c_conv2d_transpose, h_conv2d_transpose * T.int64(3) + kh_conv2d_transpose - T.int64(3), w_conv2d_transpose * T.int64(3) + kw_conv2d_transpose - T.int64(3), c_conv2d_transpose, kh_conv2d_transpose, kw_conv2d_transpose): with Ts.sblock("compute"): - v_b, v_c, v_h, v_w, v_dc, v_dh, v_dw = Ts.axis.remap("SSSSRRR", [b, c_1, h_1, w_1, dc, dh, dw]) + v_b, v_c, v_h, v_w, v_dc, v_dh, v_dw = Ts.axis.remap("SSSSRRR", [b_index, c_1, h_1, w_1, dc, dh, dw]) Ts.reads(data_pad[v_b, v_dc, v_h + v_dh, v_w + v_dw], kernel_transform[v_c, v_dc, v_dh, v_dw]) Ts.writes(compute[v_b, v_c, v_h, v_w]) with Ts.init(): @@ -805,13 +832,13 @@ def conv2d_transpose(x: T.Buffer((T.int64(1), T.int64(1), T.int64(3), T.int64(3) with Ts.sblock("kernel_dilate"): v_i0, v_i1, v_i2, v_i3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) kernel_dilate[v_i0, v_i1, v_i2, v_i3] = T.if_then_else(v_i2 % T.int64(2) == T.int64(0) and v_i3 % T.int64(2) == T.int64(0), w[v_i0, v_i1, v_i2 // T.int64(2), v_i3 // T.int64(2)], T.float32(0.0)) - for o, i, h, w_1 in T.grid(T.int64(1), T.int64(1), T.int64(3), T.int64(3)): + for o, i, h_index, w_1 in T.grid(T.int64(1), T.int64(1), T.int64(3), T.int64(3)): with Ts.sblock("kernel_transform"): - v_o, v_i, v_h, v_w = Ts.axis.remap("SSSS", [o, i, h, w_1]) + v_o, v_i, v_h, v_w = Ts.axis.remap("SSSS", [o, i, h_index, w_1]) kernel_transform[v_o, v_i, v_h, v_w] = kernel_dilate[v_i, v_o, T.int64(2) - v_h, T.int64(2) - v_w] - for b, c, h, w_1, dc, dh, dw in T.grid(T.int64(1), T.int64(1), T.int64(5), T.int64(5), T.int64(1), T.int64(3), T.int64(3)): + for b_index, c_index, h_index, w_1, dc, dh, dw in T.grid(T.int64(1), T.int64(1), T.int64(5), T.int64(5), T.int64(1), T.int64(3), T.int64(3)): with Ts.sblock("compute"): - v_b, v_c, v_h, v_w, v_dc, v_dh, v_dw = Ts.axis.remap("SSSSRRR", [b, c, h, w_1, dc, dh, dw]) + v_b, v_c, v_h, v_w, v_dc, v_dh, v_dw = Ts.axis.remap("SSSSRRR", [b_index, c_index, h_index, w_1, dc, dh, dw]) with Ts.init(): compute[v_b, v_c, v_h, v_w] = T.float32(0.0) compute[v_b, v_c, v_h, v_w] = compute[v_b, v_c, v_h, v_w] + data_pad[v_b, v_dc, v_h + v_dh, v_w + v_dw] * kernel_transform[v_c, v_dc, v_dh, v_dw] @@ -946,16 +973,17 @@ def max_pool2d(rxplaceholder: T.Buffer((T.int64(4), T.int64(6), T.int64(112), T. @pytest.mark.skip("TOPI pooling casts every shape value to i32.") def test_max_pool2d_symbolic(): # fmt: off + n = T.dynamic("n") + c = T.dynamic("c") + h = T.dynamic("h") + w = T.dynamic("w") + kh = T.dynamic("kh") + kw = T.dynamic("kw") + @tvm.script.ir_module class MaxPool2D: @R.function - def main(dumb_param: R.Tensor(("kh", "kw")), x: R.Tensor(("n", "c", "h", "w"), "float32")) -> R.Tensor(("n", "c", "h - kh + 1", "w - kw + 1"), "float32"): - n = T.int64() - c = T.int64() - h = T.int64() - w = T.int64() - kh = T.int64() - kw = T.int64() + def main(dumb_param: R.Tensor((kh, kw)), x: R.Tensor((n, c, h, w), "float32")) -> R.Tensor((n, c, h - kh + 1, w - kw + 1), "float32"): gv: R.Tensor((n, c, h - kh + 1, w - kw + 1), "float32") = R.nn.max_pool2d(x, pool_size=[kh, kw]) return gv @@ -1108,16 +1136,17 @@ def main(x: R.Tensor((4, 6, 112, 112), dtype="float32")) -> R.Tensor((4, 6, 38, @pytest.mark.skip("TOPI pooling casts every shape value to i32.") def test_avg_pool2d_symbolic(): # fmt: off + n = T.dynamic("n") + c = T.dynamic("c") + h = T.dynamic("h") + w = T.dynamic("w") + kh = T.dynamic("kh") + kw = T.dynamic("kw") + @tvm.script.ir_module class AvgPool2D: @R.function - def main(dumb_param: R.Tensor(("kh", "kw")), x: R.Tensor(("n", "c", "h", "w"), "float32")) -> R.Tensor(("n", "c", "h - kh + 1", "w - kw + 1"), "float32"): - n = T.int64() - c = T.int64() - h = T.int64() - w = T.int64() - kh = T.int64() - kw = T.int64() + def main(dumb_param: R.Tensor((kh, kw)), x: R.Tensor((n, c, h, w), "float32")) -> R.Tensor((n, c, h - kh + 1, w - kw + 1), "float32"): gv: R.Tensor((n, c, h - kh + 1, w - kw + 1), "float32") = R.nn.avg_pool2d(x, pool_size=[kh, kw]) return gv @@ -1212,14 +1241,17 @@ def adaptive_avg_pool2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(16), T.int6 @pytest.mark.skip("TOPI pooling casts every shape value to i32.") def test_adaptive_avg_pool2d_symbolic(): # fmt: off + n = T.dynamic("n") + c = T.dynamic("c") + oh = T.dynamic("oh") + ow = T.dynamic("ow") + h = T.dynamic("h") + w = T.dynamic("w") + @tvm.script.ir_module class AdaptiveAvgPool2D: @R.function - def main(dumb_param: R.Tensor(("oh", "ow")), x: R.Tensor(("n", "c", "h", "w"), "float32")) -> R.Tensor(("n", "c", "oh", "ow"), "float32"): - n = T.int64() - c = T.int64() - oh = T.int64() - ow = T.int64() + def main(dumb_param: R.Tensor((oh, ow)), x: R.Tensor((n, c, h, w), "float32")) -> R.Tensor((n, c, oh, ow), "float32"): gv: R.Tensor((n, c, oh, ow), "float32") = R.nn.adaptive_avg_pool2d(x, (oh, ow)) return gv # fmt: on @@ -1261,32 +1293,34 @@ def relu(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), compute: def test_relu_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Relu: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((m, n), "float32"): gv: R.Tensor((m, n), "float32") = R.nn.relu(x) return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_relu = T.dynamic("m") + n_relu = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.relu, (x,), R.Tensor((m, n), dtype="float32")) + def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), "float32"): + gv = R.call_tir(Expected.relu, (x,), R.Tensor((m_main, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def relu(var_rxplaceholder: T.handle, var_compute: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int64() - n = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [m, n], dtype="float32") - compute = T.match_buffer(var_compute, [m, n], dtype="float32") - for i0, i1 in T.grid(m, n): + rxplaceholder = T.match_buffer(var_rxplaceholder, [m_relu, n_relu], dtype="float32") + compute = T.match_buffer(var_compute, [m_relu, n_relu], dtype="float32") + for i0, i1 in T.grid(m_relu, n_relu): with Ts.sblock("compute"): i0_1, i1_1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[i0_1, i1_1]) @@ -1332,31 +1366,34 @@ def leaky_relu(x: T.Buffer((T.int64(2), T.int64(3)), "float32"), compute: T.Buff def test_leakyrelu_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class LeakyRelu: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((m, n), "float32"): gv: R.Tensor((m, n), "float32") = R.nn.leakyrelu(x, 0.03) return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_leaky_relu = T.dynamic("m") + n_leaky_relu = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.leaky_relu, (x, ), R.Tensor((m, n), dtype="float32")) + def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), "float32"): + gv = R.call_tir(Expected.leaky_relu, (x, ), R.Tensor((m_main, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def leaky_relu(var_x: T.handle, var_compute: T.handle): T.func_attr({"tirx.noalias": True}) - m, n = T.int64(), T.int64() - x = T.match_buffer(var_x, (m, n)) - compute = T.match_buffer(var_compute, (m, n)) - for i0, i1 in T.grid(m, n): + x = T.match_buffer(var_x, (m_leaky_relu, n_leaky_relu)) + compute = T.match_buffer(var_compute, (m_leaky_relu, n_leaky_relu)) + for i0, i1 in T.grid(m_leaky_relu, n_leaky_relu): with Ts.sblock("compute"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(x[v_i0, v_i1]) @@ -1389,9 +1426,9 @@ def prelu(x: T.Buffer((T.int64(2), T.int64(3)), "float32"), y: T.Buffer((T.int64 T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): slope_broadcasted = Ts.sblock_alloc_buffer((T.int64(3),)) - for c in range(T.int64(3)): + for c_index in range(T.int64(3)): with Ts.sblock("slope_broadcasted"): - v_c = Ts.axis.spatial(T.int64(3), c) + v_c = Ts.axis.spatial(T.int64(3), c_index) Ts.reads(y[T.int64(0)]) Ts.writes(slope_broadcasted[v_c]) slope_broadcasted[v_c] = y[T.int64(0)] @@ -1409,37 +1446,39 @@ def prelu(x: T.Buffer((T.int64(2), T.int64(3)), "float32"), y: T.Buffer((T.int64 def test_prelu_symbolic(): # fmt: off + m = T.dynamic("m") + @tvm.script.ir_module class PRelu: @R.function - def main(x: R.Tensor(("m", 7), "float32"), y: R.Tensor((1,), "float32")) -> R.Tensor(("m", 7), "float32"): - m = T.int64() + def main(x: R.Tensor((m, 7), "float32"), y: R.Tensor((1,), "float32")) -> R.Tensor((m, 7), "float32"): gv: R.Tensor((m, 7), "float32") = R.nn.prelu(x, y) return gv + m_main = T.dynamic("m") + m_prelu = T.dynamic("m") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", 7), dtype="float32"), y: R.Tensor((1,), dtype="float32")) -> R.Tensor(("m", 7), dtype="float32"): - m = T.int64() - gv = R.call_tir(Expected.prelu, (x, y), out_ty=R.Tensor((m, 7), dtype="float32")) + def main(x: R.Tensor((m_main, 7), dtype="float32"), y: R.Tensor((1,), dtype="float32")) -> R.Tensor((m_main, 7), dtype="float32"): + gv = R.call_tir(Expected.prelu, (x, y), out_ty=R.Tensor((m_main, 7), dtype="float32")) return gv @Ts.prim_func(private=True) def prelu(var_x: T.handle, y: T.Buffer((T.int64(1),), "float32"), var_compute: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int64() - x = T.match_buffer(var_x, (m, T.int64(7))) - compute = T.match_buffer(var_compute, (m, T.int64(7))) + x = T.match_buffer(var_x, (m_prelu, T.int64(7))) + compute = T.match_buffer(var_compute, (m_prelu, T.int64(7))) # with Ts.sblock("root"): slope_broadcasted = Ts.sblock_alloc_buffer((T.int64(7),)) - for c in range(T.int64(7)): + for c_index in range(T.int64(7)): with Ts.sblock("slope_broadcasted"): - v_c = Ts.axis.spatial(T.int64(7), c) + v_c = Ts.axis.spatial(T.int64(7), c_index) Ts.reads(y[T.int64(0)]) Ts.writes(slope_broadcasted[v_c]) slope_broadcasted[v_c] = y[T.int64(0)] - for i0, i1 in T.grid(m, T.int64(7)): + for i0, i1 in T.grid(m_prelu, T.int64(7)): with Ts.sblock("compute"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(x[v_i0, v_i1], slope_broadcasted[v_i1]) @@ -1512,59 +1551,62 @@ def gelu(x: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_multiply: T.Buffer( def test_gelu_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Gelu: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((m, n), "float32"): gv: R.Tensor((m, n), "float32") = R.nn.gelu(x) return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_gelu = T.dynamic("m") + n_gelu = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.gelu, (x,), R.Tensor((m, n), dtype="float32")) + def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), "float32"): + gv = R.call_tir(Expected.gelu, (x,), R.Tensor((m_main, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def gelu(var_x: T.handle, var_T_multiply: T.handle): T.func_attr({"tirx.noalias": True}) - m, n = T.int64(), T.int64() - x = T.match_buffer(var_x, (m, n)) - T_multiply = T.match_buffer(var_T_multiply, (m, n)) - T_multiply_1 = Ts.sblock_alloc_buffer((m, n)) - compute = Ts.sblock_alloc_buffer((m, n)) - T_multiply_2 = Ts.sblock_alloc_buffer((m, n)) - T_add = Ts.sblock_alloc_buffer((m, n)) - for ax0, ax1 in T.grid(m, n): + x = T.match_buffer(var_x, (m_gelu, n_gelu)) + T_multiply = T.match_buffer(var_T_multiply, (m_gelu, n_gelu)) + T_multiply_1 = Ts.sblock_alloc_buffer((m_gelu, n_gelu)) + compute = Ts.sblock_alloc_buffer((m_gelu, n_gelu)) + T_multiply_2 = Ts.sblock_alloc_buffer((m_gelu, n_gelu)) + T_add = Ts.sblock_alloc_buffer((m_gelu, n_gelu)) + for ax0, ax1 in T.grid(m_gelu, n_gelu): with Ts.sblock("T_multiply"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(x[v_ax0, v_ax1]) Ts.writes(T_multiply_1[v_ax0, v_ax1]) T_multiply_1[v_ax0, v_ax1] = x[v_ax0, v_ax1] * T.float32(0.70710678118654757) - for i0, i1 in T.grid(m, n): + for i0, i1 in T.grid(m_gelu, n_gelu): with Ts.sblock("compute"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(T_multiply_1[v_i0, v_i1]) Ts.writes(compute[v_i0, v_i1]) compute[v_i0, v_i1] = T.erf(T_multiply_1[v_i0, v_i1]) - for ax0, ax1 in T.grid(m, n): + for ax0, ax1 in T.grid(m_gelu, n_gelu): with Ts.sblock("T_multiply_1"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(compute[v_ax0, v_ax1]) Ts.writes(T_multiply_2[v_ax0, v_ax1]) T_multiply_2[v_ax0, v_ax1] = compute[v_ax0, v_ax1] * T.float32(0.5) - for ax0, ax1 in T.grid(m, n): + for ax0, ax1 in T.grid(m_gelu, n_gelu): with Ts.sblock("T_add"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(T_multiply_2[v_ax0, v_ax1]) Ts.writes(T_add[v_ax0, v_ax1]) T_add[v_ax0, v_ax1] = T.float32(0.5) + T_multiply_2[v_ax0, v_ax1] - for ax0, ax1 in T.grid(m, n): + for ax0, ax1 in T.grid(m_gelu, n_gelu): with Ts.sblock("T_multiply_2"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(x[v_ax0, v_ax1], T_add[v_ax0, v_ax1]) @@ -1664,88 +1706,91 @@ def gelu_tanh(A: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_multiply: T.Bu def test_gelu_tanh_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class GeluTanh: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((m, n), "float32"): gv: R.Tensor((m, n), "float32") = R.nn.gelu_tanh(x) return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_gelu_tanh = T.dynamic("m") + n_gelu_tanh = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), dtype="float32")) -> R.Tensor(("m", "n"), dtype="float32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.gelu_tanh, (x,), out_ty=R.Tensor((m, n), dtype="float32")) + def main(x: R.Tensor((m_main, n_main), dtype="float32")) -> R.Tensor((m_main, n_main), dtype="float32"): + gv = R.call_tir(Expected.gelu_tanh, (x,), out_ty=R.Tensor((m_main, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def gelu_tanh(var_A: T.handle, var_T_multiply: T.handle): T.func_attr({"tirx.noalias": True}) - m, n = T.int64(), T.int64() - A = T.match_buffer(var_A, (m, n)) - T_multiply = T.match_buffer(var_T_multiply, (m, n)) + A = T.match_buffer(var_A, (m_gelu_tanh, n_gelu_tanh)) + T_multiply = T.match_buffer(var_T_multiply, (m_gelu_tanh, n_gelu_tanh)) # with Ts.sblock("root"): - T_multiply_1 = Ts.sblock_alloc_buffer((m, n)) - T_multiply_2 = Ts.sblock_alloc_buffer((m, n)) - T_multiply_3 = Ts.sblock_alloc_buffer((m, n)) - T_multiply_4 = Ts.sblock_alloc_buffer((m, n)) - T_add = Ts.sblock_alloc_buffer((m, n)) - T_multiply_5 = Ts.sblock_alloc_buffer((m, n)) - compute = Ts.sblock_alloc_buffer((m, n)) - T_add_1 = Ts.sblock_alloc_buffer((m, n)) - for ax0, ax1 in T.grid(m, n): + T_multiply_1 = Ts.sblock_alloc_buffer((m_gelu_tanh, n_gelu_tanh)) + T_multiply_2 = Ts.sblock_alloc_buffer((m_gelu_tanh, n_gelu_tanh)) + T_multiply_3 = Ts.sblock_alloc_buffer((m_gelu_tanh, n_gelu_tanh)) + T_multiply_4 = Ts.sblock_alloc_buffer((m_gelu_tanh, n_gelu_tanh)) + T_add = Ts.sblock_alloc_buffer((m_gelu_tanh, n_gelu_tanh)) + T_multiply_5 = Ts.sblock_alloc_buffer((m_gelu_tanh, n_gelu_tanh)) + compute = Ts.sblock_alloc_buffer((m_gelu_tanh, n_gelu_tanh)) + T_add_1 = Ts.sblock_alloc_buffer((m_gelu_tanh, n_gelu_tanh)) + for ax0, ax1 in T.grid(m_gelu_tanh, n_gelu_tanh): with Ts.sblock("T_multiply"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(A[v_ax0, v_ax1]) Ts.writes(T_multiply_1[v_ax0, v_ax1]) T_multiply_1[v_ax0, v_ax1] = T.float32(0.5) * A[v_ax0, v_ax1] - for ax0, ax1 in T.grid(m, n): + for ax0, ax1 in T.grid(m_gelu_tanh, n_gelu_tanh): with Ts.sblock("T_multiply_1"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(A[v_ax0, v_ax1]) Ts.writes(T_multiply_2[v_ax0, v_ax1]) T_multiply_2[v_ax0, v_ax1] = T.float32(0.79788456080286541) * A[v_ax0, v_ax1] - for ax0, ax1 in T.grid(m, n): + for ax0, ax1 in T.grid(m_gelu_tanh, n_gelu_tanh): with Ts.sblock("T_multiply_2"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(A[v_ax0, v_ax1]) Ts.writes(T_multiply_3[v_ax0, v_ax1]) T_multiply_3[v_ax0, v_ax1] = T.float32(0.044714999999999998) * A[v_ax0, v_ax1] - for ax0, ax1 in T.grid(m, n): + for ax0, ax1 in T.grid(m_gelu_tanh, n_gelu_tanh): with Ts.sblock("T_multiply_3"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(T_multiply_3[v_ax0, v_ax1], A[v_ax0, v_ax1]) Ts.writes(T_multiply_4[v_ax0, v_ax1]) T_multiply_4[v_ax0, v_ax1] = T_multiply_3[v_ax0, v_ax1] * A[v_ax0, v_ax1] - for ax0, ax1 in T.grid(m, n): + for ax0, ax1 in T.grid(m_gelu_tanh, n_gelu_tanh): with Ts.sblock("T_add"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(T_multiply_4[v_ax0, v_ax1]) Ts.writes(T_add[v_ax0, v_ax1]) T_add[v_ax0, v_ax1] = T.float32(1) + T_multiply_4[v_ax0, v_ax1] - for ax0, ax1 in T.grid(m, n): + for ax0, ax1 in T.grid(m_gelu_tanh, n_gelu_tanh): with Ts.sblock("T_multiply_4"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(T_multiply_2[v_ax0, v_ax1], T_add[v_ax0, v_ax1]) Ts.writes(T_multiply_5[v_ax0, v_ax1]) T_multiply_5[v_ax0, v_ax1] = T_multiply_2[v_ax0, v_ax1] * T_add[v_ax0, v_ax1] - for i0, i1 in T.grid(m, n): + for i0, i1 in T.grid(m_gelu_tanh, n_gelu_tanh): with Ts.sblock("compute"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(T_multiply_5[v_i0, v_i1]) Ts.writes(compute[v_i0, v_i1]) compute[v_i0, v_i1] = T.tanh(T_multiply_5[v_i0, v_i1]) - for ax0, ax1 in T.grid(m, n): + for ax0, ax1 in T.grid(m_gelu_tanh, n_gelu_tanh): with Ts.sblock("T_add_1"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(compute[v_ax0, v_ax1]) Ts.writes(T_add_1[v_ax0, v_ax1]) T_add_1[v_ax0, v_ax1] = T.float32(1) + compute[v_ax0, v_ax1] - for ax0, ax1 in T.grid(m, n): + for ax0, ax1 in T.grid(m_gelu_tanh, n_gelu_tanh): with Ts.sblock("T_multiply_5"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(T_multiply_1[v_ax0, v_ax1], T_add_1[v_ax0, v_ax1]) @@ -1797,39 +1842,41 @@ def silu(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_multipl def test_silu_symbolic(): # fmt: off + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Silu: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((m, n), "float32"): gv: R.Tensor((m, n), "float32") = R.nn.silu(x) return gv + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_silu = T.dynamic("m") + n_silu = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor(("m", "n"), "float32"): - m = T.int64() - n = T.int64() - gv = R.call_tir(Expected.silu, (x,), R.Tensor((m, n), dtype="float32")) + def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), "float32"): + gv = R.call_tir(Expected.silu, (x,), R.Tensor((m_main, n_main), dtype="float32")) return gv @Ts.prim_func(private=True) def silu(var_rxplaceholder: T.handle, var_T_multiply: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int64() - n = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [m, n], dtype="float32") - T_multiply = T.match_buffer(var_T_multiply, [m, n], dtype="float32") - compute = Ts.sblock_alloc_buffer([m, n], dtype="float32") - for i0, i1 in T.grid(m, n): + rxplaceholder = T.match_buffer(var_rxplaceholder, [m_silu, n_silu], dtype="float32") + T_multiply = T.match_buffer(var_T_multiply, [m_silu, n_silu], dtype="float32") + compute = Ts.sblock_alloc_buffer([m_silu, n_silu], dtype="float32") + for i0, i1 in T.grid(m_silu, n_silu): with Ts.sblock("compute"): i0_1, i1_1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[i0_1, i1_1]) Ts.writes(compute[i0_1, i1_1]) compute[i0_1, i1_1] = T.sigmoid(rxplaceholder[i0_1, i1_1]) - for i0, i1 in T.grid(m, n): + for i0, i1 in T.grid(m_silu, n_silu): with Ts.sblock("T_multiply"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder[ax0, ax1], compute[ax0, ax1]) @@ -1900,38 +1947,40 @@ def softmax(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(16), T.int6 def test_softmax_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + @tvm.script.ir_module class Softmax: @R.function - def main(x: R.Tensor(("a", "b", "c"), "float32")) -> R.Tensor(("a", "b", "c"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() + def main(x: R.Tensor((a, b, c), "float32")) -> R.Tensor((a, b, c), "float32"): gv: R.Tensor((a, b, c), "float32") = R.nn.softmax(x) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + a_softmax = T.dynamic("a") + b_softmax = T.dynamic("b") + c_softmax = T.dynamic("c") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c"), "float32")) -> R.Tensor(("a", "b", "c"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - gv = R.call_tir(Expected.softmax, (x,), R.Tensor((a, b, c), dtype="float32")) + def main(x: R.Tensor((a_main, b_main, c_main), "float32")) -> R.Tensor((a_main, b_main, c_main), "float32"): + gv = R.call_tir(Expected.softmax, (x,), R.Tensor((a_main, b_main, c_main), dtype="float32")) return gv @Ts.prim_func(private=True) def softmax(var_rxplaceholder: T.handle, var_T_softmax_norm: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, c], dtype="float32") - T_softmax_norm = T.match_buffer(var_T_softmax_norm, [a, b, c], dtype="float32") - T_softmax_maxelem = Ts.sblock_alloc_buffer([a, b], dtype="float32") - T_softmax_exp = Ts.sblock_alloc_buffer([a, b, c], dtype="float32") - T_softmax_expsum = Ts.sblock_alloc_buffer([a, b], dtype="float32") - for i0, i1, i2 in T.grid(a, b, c): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_softmax, b_softmax, c_softmax], dtype="float32") + T_softmax_norm = T.match_buffer(var_T_softmax_norm, [a_softmax, b_softmax, c_softmax], dtype="float32") + T_softmax_maxelem = Ts.sblock_alloc_buffer([a_softmax, b_softmax], dtype="float32") + T_softmax_exp = Ts.sblock_alloc_buffer([a_softmax, b_softmax, c_softmax], dtype="float32") + T_softmax_expsum = Ts.sblock_alloc_buffer([a_softmax, b_softmax], dtype="float32") + for i0, i1, i2 in T.grid(a_softmax, b_softmax, c_softmax): with Ts.sblock("T_softmax_maxelem"): i0_1, i1_1, k = Ts.axis.remap("SSR", [i0, i1, i2]) Ts.reads(rxplaceholder[i0_1, i1_1, k]) @@ -1939,13 +1988,13 @@ def softmax(var_rxplaceholder: T.handle, var_T_softmax_norm: T.handle): with Ts.init(): T_softmax_maxelem[i0_1, i1_1] = T.float32(-3.4028234663852886e+38) T_softmax_maxelem[i0_1, i1_1] = T.max(T_softmax_maxelem[i0_1, i1_1], rxplaceholder[i0_1, i1_1, k]) - for i0, i1, i2 in T.grid(a, b, c): + for i0, i1, i2 in T.grid(a_softmax, b_softmax, c_softmax): with Ts.sblock("T_softmax_exp"): i0_2, i1_2, i2_1 = Ts.axis.remap("SSS", [i0, i1, i2]) Ts.reads(rxplaceholder[i0_2, i1_2, i2_1], T_softmax_maxelem[i0_2, i1_2]) Ts.writes(T_softmax_exp[i0_2, i1_2, i2_1]) T_softmax_exp[i0_2, i1_2, i2_1] = T.exp(rxplaceholder[i0_2, i1_2, i2_1] - T_softmax_maxelem[i0_2, i1_2], dtype="float32") - for i0_3, i1_3, i2 in T.grid(a, b, c): + for i0_3, i1_3, i2 in T.grid(a_softmax, b_softmax, c_softmax): with Ts.sblock("T_softmax_expsum"): i0_4, i1_4, k = Ts.axis.remap("SSR", [i0_3, i1_3, i2]) Ts.reads(T_softmax_exp[i0_4, i1_4, k]) @@ -1953,7 +2002,7 @@ def softmax(var_rxplaceholder: T.handle, var_T_softmax_norm: T.handle): with Ts.init(): T_softmax_expsum[i0_4, i1_4] = T.float32(0) T_softmax_expsum[i0_4, i1_4] = T_softmax_expsum[i0_4, i1_4] + T_softmax_exp[i0_4, i1_4, k] - for i0_5, i1_5, i2 in T.grid(a, b, c): + for i0_5, i1_5, i2 in T.grid(a_softmax, b_softmax, c_softmax): with Ts.sblock("T_softmax_norm"): i0_6, i1_6, i2_2 = Ts.axis.remap("SSS", [i0_5, i1_5, i2]) Ts.reads(T_softmax_exp[i0_6, i1_6, i2_2], T_softmax_expsum[i0_6, i1_6]) @@ -2018,38 +2067,40 @@ def log_softmax(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(16), T. def test_log_softmax_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + @tvm.script.ir_module class LogSoftmax: @R.function - def main(x: R.Tensor(("a", "b", "c"), "float32")) -> R.Tensor(("a", "b", "c"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() + def main(x: R.Tensor((a, b, c), "float32")) -> R.Tensor((a, b, c), "float32"): gv: R.Tensor((a, b, c), "float32") = R.nn.log_softmax(x) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + a_log_softmax = T.dynamic("a") + b_log_softmax = T.dynamic("b") + c_log_softmax = T.dynamic("c") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c"), dtype="float32")) -> R.Tensor(("a", "b", "c"), dtype="float32"): - a = T.int64() - b = T.int64() - c = T.int64() + def main(x: R.Tensor((a_main, b_main, c_main), dtype="float32")) -> R.Tensor((a_main, b_main, c_main), dtype="float32"): # block 0 - gv = R.call_tir(Expected.log_softmax, (x,), R.Tensor((a, b, c), dtype="float32")) + gv = R.call_tir(Expected.log_softmax, (x,), R.Tensor((a_main, b_main, c_main), dtype="float32")) return gv @Ts.prim_func(private=True) def log_softmax(var_rxplaceholder: T.handle, var_compute: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, c], dtype="float32") - compute = T.match_buffer(var_compute, [a, b, c], dtype="float32") - T_softmax_maxelem = Ts.sblock_alloc_buffer([a, b], dtype="float32") - compute_1 = Ts.sblock_alloc_buffer([a, b], dtype="float32") - for i0, i1, k in T.grid(a, b, c): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_log_softmax, b_log_softmax, c_log_softmax], dtype="float32") + compute = T.match_buffer(var_compute, [a_log_softmax, b_log_softmax, c_log_softmax], dtype="float32") + T_softmax_maxelem = Ts.sblock_alloc_buffer([a_log_softmax, b_log_softmax], dtype="float32") + compute_1 = Ts.sblock_alloc_buffer([a_log_softmax, b_log_softmax], dtype="float32") + for i0, i1, k in T.grid(a_log_softmax, b_log_softmax, c_log_softmax): with Ts.sblock("T_softmax_maxelem"): v_i0, v_i1, v_k = Ts.axis.remap("SSR", [i0, i1, k]) Ts.reads(rxplaceholder[v_i0, v_i1, v_k]) @@ -2057,7 +2108,7 @@ def log_softmax(var_rxplaceholder: T.handle, var_compute: T.handle): with Ts.init(): T_softmax_maxelem[v_i0, v_i1] = T.float32(-3.4028234663852886e38) T_softmax_maxelem[v_i0, v_i1] = T.max(T_softmax_maxelem[v_i0, v_i1], rxplaceholder[v_i0, v_i1, v_k]) - for i0, i1, k in T.grid(a, b, c): + for i0, i1, k in T.grid(a_log_softmax, b_log_softmax, c_log_softmax): with Ts.sblock("compute"): v_i0, v_i1, v_k = Ts.axis.remap("SSR", [i0, i1, k]) Ts.reads(rxplaceholder[v_i0, v_i1, v_k], T_softmax_maxelem[v_i0, v_i1]) @@ -2065,7 +2116,7 @@ def log_softmax(var_rxplaceholder: T.handle, var_compute: T.handle): with Ts.init(): compute_1[v_i0, v_i1] = T.float32(0) compute_1[v_i0, v_i1] = compute_1[v_i0, v_i1] + T.exp(rxplaceholder[v_i0, v_i1, v_k] - T_softmax_maxelem[v_i0, v_i1], dtype="float32") - for i0, i1, i2 in T.grid(a, b, c): + for i0, i1, i2 in T.grid(a_log_softmax, b_log_softmax, c_log_softmax): with Ts.sblock("compute_1"): v_i0, v_i1, v_i2 = Ts.axis.remap("SSS", [i0, i1, i2]) Ts.reads(rxplaceholder[v_i0, v_i1, v_i2], T_softmax_maxelem[v_i0, v_i1], compute_1[v_i0, v_i1],) @@ -2178,38 +2229,43 @@ def cross_entropy_with_logits(x: T.Buffer((T.int64(2), T.int64(3)), "float32"), def test_cross_entropy_with_logits_batch_symbolic(): # fmt: off + n = T.dynamic("n") + m = T.dynamic("m") + @tvm.script.ir_module class CrossEntropyWithLogits: @R.function - def main(x: R.Tensor(("n", "m"), "float32"), y: R.Tensor(("n", "m"), "float32")) -> R.Tensor(None, "float32", ndim=2): - n = T.int64() - m = T.int64() + def main(x: R.Tensor((n, m), "float32"), y: R.Tensor((n, m), "float32")) -> R.Tensor(None, "float32", ndim=2): gv: R.Tensor((), "float32") = R.nn.cross_entropy_with_logits(x, y) return gv + n_main = T.dynamic("n") + m_main = T.dynamic("m") + m_cross_entropy_with_logits = T.dynamic("m") + n_cross_entropy_with_logits = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("n", "m"), dtype="float32"), y: R.Tensor(("n", "m"), dtype="float32")): + def main(x: R.Tensor((n_main, m_main), dtype="float32"), y: R.Tensor((n_main, m_main), dtype="float32")): gv = R.call_tir(Expected.cross_entropy_with_logits, (x, y), R.Tensor((), dtype="float32")) return gv @Ts.prim_func(private=True) def cross_entropy_with_logits(var_x: T.handle, var_y: T.handle, T_divide: T.Buffer((), "float32")): T.func_attr({"tirx.noalias": True}) - m, n = T.int64(), T.int64() - x = T.match_buffer(var_x, (n, m)) - y = T.match_buffer(var_y, (n, m)) - T_multiply = Ts.sblock_alloc_buffer((n, m)) + x = T.match_buffer(var_x, (n_cross_entropy_with_logits, m_cross_entropy_with_logits)) + y = T.match_buffer(var_y, (n_cross_entropy_with_logits, m_cross_entropy_with_logits)) + T_multiply = Ts.sblock_alloc_buffer((n_cross_entropy_with_logits, m_cross_entropy_with_logits)) T_multiply_red = Ts.sblock_alloc_buffer(()) T_multiply_1 = Ts.sblock_alloc_buffer(()) - for ax0, ax1 in T.grid(n, m): + for ax0, ax1 in T.grid(n_cross_entropy_with_logits, m_cross_entropy_with_logits): with Ts.sblock("T_multiply"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(x[v_ax0, v_ax1], y[v_ax0, v_ax1]) Ts.writes(T_multiply[v_ax0, v_ax1]) T_multiply[v_ax0, v_ax1] = x[v_ax0, v_ax1] * y[v_ax0, v_ax1] - for k0, k1 in T.grid(n, m): + for k0, k1 in T.grid(n_cross_entropy_with_logits, m_cross_entropy_with_logits): with Ts.sblock("T_multiply_red"): v_k0, v_k1 = Ts.axis.remap("RR", [k0, k1]) Ts.reads(T_multiply[v_k0, v_k1]) @@ -2226,7 +2282,7 @@ def cross_entropy_with_logits(var_x: T.handle, var_y: T.handle, T_divide: T.Buff vi = Ts.axis.spatial(T.int64(1), T.int64(0)) Ts.reads(T_multiply_1[()]) Ts.writes(T_divide[()]) - T_divide[()] = T_multiply_1[()] / T.Cast("float32", n) + T_divide[()] = T_multiply_1[()] / T.Cast("float32", n_cross_entropy_with_logits) # fmt: on mod = LegalizeOps()(CrossEntropyWithLogits) @@ -2524,295 +2580,300 @@ def main(x: R.Tensor((2, 3, 28, 28), dtype="float32"), gamma: R.Tensor((3,), dty def test_batch_norm_symbolic(): # fmt: off + n = T.dynamic("n") + h = T.dynamic("h") + w = T.dynamic("w") + c = T.dynamic("c") + @tvm.script.ir_module class BatchNorm: @R.function - def main(x: R.Tensor(("n", "h", "w", "c"), "float32"), gamma: R.Tensor(("c",), "float32"), beta: R.Tensor(("c",), "float32"), moving_mean: R.Tensor(("c",), "float32"), moving_var: R.Tensor(("c",), "float32")) -> R.Tuple(R.Tensor(("n", "h", "w", "c"), "float32"), R.Tensor(("c",), "float32"), R.Tensor(("c",), "float32")): - n = T.int64() - h = T.int64() - w = T.int64() - c = T.int64() + def main(x: R.Tensor((n, h, w, c), "float32"), gamma: R.Tensor((c,), "float32"), beta: R.Tensor((c,), "float32"), moving_mean: R.Tensor((c,), "float32"), moving_var: R.Tensor((c,), "float32")) -> R.Tuple(R.Tensor((n, h, w, c), "float32"), R.Tensor((c,), "float32"), R.Tensor((c,), "float32")): gv: R.Tuple(R.Tensor((n, h, w, c), "float32"), R.Tensor((c,), "float32"), R.Tensor((c,), "float32")) = R.nn.batch_norm(x, gamma, beta, moving_mean, moving_var, axis=1) return gv + n_batch_norm = T.dynamic("n") + h_batch_norm = T.dynamic("h") + w_batch_norm = T.dynamic("w") + c_batch_norm = T.dynamic("c") + n_main = T.dynamic("n") + h_main = T.dynamic("h") + w_main = T.dynamic("w") + c_main = T.dynamic("c") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) def batch_norm(var_x: T.handle, var_gamma: T.handle, var_beta: T.handle, var_moving_mean: T.handle, var_moving_var: T.handle, var_T_add: T.handle, var_T_add_1: T.handle, var_T_add_2: T.handle): T.func_attr({"tirx.noalias": True}) - n, h, w, c = T.int64(), T.int64(), T.int64(), T.int64() - x = T.match_buffer(var_x, (n, h, w, c)) - gamma = T.match_buffer(var_gamma, (c,)) - beta = T.match_buffer(var_beta, (c,)) - moving_mean = T.match_buffer(var_moving_mean, (c,)) - moving_var = T.match_buffer(var_moving_var, (c,)) - T_add = T.match_buffer(var_T_add, (n, h, w, c)) - T_add_1 = T.match_buffer(var_T_add_1, (T.max(c, h),)) - T_add_2 = T.match_buffer(var_T_add_2, (T.max(c, h),)) + x = T.match_buffer(var_x, (n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)) + gamma = T.match_buffer(var_gamma, (c_batch_norm,)) + beta = T.match_buffer(var_beta, (c_batch_norm,)) + moving_mean = T.match_buffer(var_moving_mean, (c_batch_norm,)) + moving_var = T.match_buffer(var_moving_var, (c_batch_norm,)) + T_add = T.match_buffer(var_T_add, (n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)) + T_add_1 = T.match_buffer(var_T_add_1, (T.max(c_batch_norm, h_batch_norm),)) + T_add_2 = T.match_buffer(var_T_add_2, (T.max(c_batch_norm, h_batch_norm),)) with Ts.sblock("root"): Ts.reads() Ts.writes() - x_red = Ts.sblock_alloc_buffer((h,)) - T_divide = Ts.sblock_alloc_buffer((h,)) - T_reshape = Ts.sblock_alloc_buffer((T.int64(1), h, T.int64(1), T.int64(1))) - T_subtract = Ts.sblock_alloc_buffer((n, h, w, c)) - T_subtract_1 = Ts.sblock_alloc_buffer((n, h, w, c)) - T_subtract_2 = Ts.sblock_alloc_buffer((n, h, w, c)) - T_multiply = Ts.sblock_alloc_buffer((n, h, w, c)) - T_multiply_red = Ts.sblock_alloc_buffer((h,)) - T_divide_1 = Ts.sblock_alloc_buffer((h,)) - T_reshape_1 = Ts.sblock_alloc_buffer((T.int64(1), h, T.int64(1), T.int64(1))) - T_add_3 = Ts.sblock_alloc_buffer((T.int64(1), h, T.int64(1), T.int64(1))) - compute = Ts.sblock_alloc_buffer((T.int64(1), h, T.int64(1), T.int64(1))) - T_divide_2 = Ts.sblock_alloc_buffer((n, h, w, c)) - T_reshape_2 = Ts.sblock_alloc_buffer((T.int64(1), h, T.int64(1), T.int64(1))) - T_multiply_1 = Ts.sblock_alloc_buffer((n, h, w, c)) - T_reshape_3 = Ts.sblock_alloc_buffer((T.int64(1), h, T.int64(1), T.int64(1))) - T_multiply_2 = Ts.sblock_alloc_buffer((c,)) - T_multiply_3 = Ts.sblock_alloc_buffer((h,)) - T_multiply_4 = Ts.sblock_alloc_buffer((c,)) - T_multiply_5 = Ts.sblock_alloc_buffer((h,)) - for ax0 in range(h): - for k0 in range(n): - for k2 in range(w): - for k3 in range(c): + x_red = Ts.sblock_alloc_buffer((h_batch_norm,)) + T_divide = Ts.sblock_alloc_buffer((h_batch_norm,)) + T_reshape = Ts.sblock_alloc_buffer((T.int64(1), h_batch_norm, T.int64(1), T.int64(1))) + T_subtract = Ts.sblock_alloc_buffer((n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)) + T_subtract_1 = Ts.sblock_alloc_buffer((n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)) + T_subtract_2 = Ts.sblock_alloc_buffer((n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)) + T_multiply = Ts.sblock_alloc_buffer((n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)) + T_multiply_red = Ts.sblock_alloc_buffer((h_batch_norm,)) + T_divide_1 = Ts.sblock_alloc_buffer((h_batch_norm,)) + T_reshape_1 = Ts.sblock_alloc_buffer((T.int64(1), h_batch_norm, T.int64(1), T.int64(1))) + T_add_3 = Ts.sblock_alloc_buffer((T.int64(1), h_batch_norm, T.int64(1), T.int64(1))) + compute = Ts.sblock_alloc_buffer((T.int64(1), h_batch_norm, T.int64(1), T.int64(1))) + T_divide_2 = Ts.sblock_alloc_buffer((n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)) + T_reshape_2 = Ts.sblock_alloc_buffer((T.int64(1), h_batch_norm, T.int64(1), T.int64(1))) + T_multiply_1 = Ts.sblock_alloc_buffer((n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)) + T_reshape_3 = Ts.sblock_alloc_buffer((T.int64(1), h_batch_norm, T.int64(1), T.int64(1))) + T_multiply_2 = Ts.sblock_alloc_buffer((c_batch_norm,)) + T_multiply_3 = Ts.sblock_alloc_buffer((h_batch_norm,)) + T_multiply_4 = Ts.sblock_alloc_buffer((c_batch_norm,)) + T_multiply_5 = Ts.sblock_alloc_buffer((h_batch_norm,)) + for ax0 in range(h_batch_norm): + for k0 in range(n_batch_norm): + for k2 in range(w_batch_norm): + for k3 in range(c_batch_norm): with Ts.sblock("x_red"): - v_ax0 = Ts.axis.spatial(h, ax0) - v_k0 = Ts.axis.reduce(n, k0) - v_k2 = Ts.axis.reduce(w, k2) - v_k3 = Ts.axis.reduce(c, k3) + v_ax0 = Ts.axis.spatial(h_batch_norm, ax0) + v_k0 = Ts.axis.reduce(n_batch_norm, k0) + v_k2 = Ts.axis.reduce(w_batch_norm, k2) + v_k3 = Ts.axis.reduce(c_batch_norm, k3) Ts.reads(x[v_k0, v_ax0, v_k2, v_k3]) Ts.writes(x_red[v_ax0]) with Ts.init(): x_red[v_ax0] = T.float32(0.0) x_red[v_ax0] = x_red[v_ax0] + x[v_k0, v_ax0, v_k2, v_k3] - for ax0 in range(h): + for ax0 in range(h_batch_norm): with Ts.sblock("T_divide"): - v_ax0 = Ts.axis.spatial(h, ax0) + v_ax0 = Ts.axis.spatial(h_batch_norm, ax0) Ts.reads(x_red[v_ax0]) Ts.writes(T_divide[v_ax0]) - T_divide[v_ax0] = x_red[v_ax0] / T.Cast("float32", n * w * c) + T_divide[v_ax0] = x_red[v_ax0] / T.Cast("float32", n_batch_norm * w_batch_norm * c_batch_norm) for ax0 in range(T.int64(1)): - for ax1 in range(h): + for ax1 in range(h_batch_norm): for ax2 in range(T.int64(1)): for ax3 in range(T.int64(1)): with Ts.sblock("T_reshape"): v_ax0 = Ts.axis.spatial(T.int64(1), ax0) - v_ax1 = Ts.axis.spatial(h, ax1) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) v_ax2 = Ts.axis.spatial(T.int64(1), ax2) v_ax3 = Ts.axis.spatial(T.int64(1), ax3) - Ts.reads(T_divide[(v_ax0 * h + v_ax1 + v_ax2 + v_ax3) % h]) + Ts.reads(T_divide[(v_ax0 * h_batch_norm + v_ax1 + v_ax2 + v_ax3) % h_batch_norm]) Ts.writes(T_reshape[v_ax0, v_ax1, v_ax2, v_ax3]) - T_reshape[v_ax0, v_ax1, v_ax2, v_ax3] = T_divide[(v_ax0 * h + v_ax1 + v_ax2 + v_ax3) % h] - for ax0 in range(n): - for ax1 in range(h): - for ax2 in range(w): - for ax3 in range(c): + T_reshape[v_ax0, v_ax1, v_ax2, v_ax3] = T_divide[(v_ax0 * h_batch_norm + v_ax1 + v_ax2 + v_ax3) % h_batch_norm] + for ax0 in range(n_batch_norm): + for ax1 in range(h_batch_norm): + for ax2 in range(w_batch_norm): + for ax3 in range(c_batch_norm): with Ts.sblock("T_subtract"): - v_ax0 = Ts.axis.spatial(n, ax0) - v_ax1 = Ts.axis.spatial(h, ax1) - v_ax2 = Ts.axis.spatial(w, ax2) - v_ax3 = Ts.axis.spatial(c, ax3) + v_ax0 = Ts.axis.spatial(n_batch_norm, ax0) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) + v_ax2 = Ts.axis.spatial(w_batch_norm, ax2) + v_ax3 = Ts.axis.spatial(c_batch_norm, ax3) Ts.reads(x[v_ax0, v_ax1, v_ax2, v_ax3], T_reshape[T.int64(0), v_ax1, T.int64(0), T.int64(0)]) Ts.writes(T_subtract[v_ax0, v_ax1, v_ax2, v_ax3]) T_subtract[v_ax0, v_ax1, v_ax2, v_ax3] = x[v_ax0, v_ax1, v_ax2, v_ax3] - T_reshape[T.int64(0), v_ax1, T.int64(0), T.int64(0)] - for ax0 in range(n): - for ax1 in range(h): - for ax2 in range(w): - for ax3 in range(c): + for ax0 in range(n_batch_norm): + for ax1 in range(h_batch_norm): + for ax2 in range(w_batch_norm): + for ax3 in range(c_batch_norm): with Ts.sblock("T_subtract_1"): - v_ax0 = Ts.axis.spatial(n, ax0) - v_ax1 = Ts.axis.spatial(h, ax1) - v_ax2 = Ts.axis.spatial(w, ax2) - v_ax3 = Ts.axis.spatial(c, ax3) + v_ax0 = Ts.axis.spatial(n_batch_norm, ax0) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) + v_ax2 = Ts.axis.spatial(w_batch_norm, ax2) + v_ax3 = Ts.axis.spatial(c_batch_norm, ax3) Ts.reads(x[v_ax0, v_ax1, v_ax2, v_ax3], T_reshape[T.int64(0), v_ax1, T.int64(0), T.int64(0)]) Ts.writes(T_subtract_1[v_ax0, v_ax1, v_ax2, v_ax3]) T_subtract_1[v_ax0, v_ax1, v_ax2, v_ax3] = x[v_ax0, v_ax1, v_ax2, v_ax3] - T_reshape[T.int64(0), v_ax1, T.int64(0), T.int64(0)] - for ax0 in range(n): - for ax1 in range(h): - for ax2 in range(w): - for ax3 in range(c): + for ax0 in range(n_batch_norm): + for ax1 in range(h_batch_norm): + for ax2 in range(w_batch_norm): + for ax3 in range(c_batch_norm): with Ts.sblock("T_subtract_2"): - v_ax0 = Ts.axis.spatial(n, ax0) - v_ax1 = Ts.axis.spatial(h, ax1) - v_ax2 = Ts.axis.spatial(w, ax2) - v_ax3 = Ts.axis.spatial(c, ax3) + v_ax0 = Ts.axis.spatial(n_batch_norm, ax0) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) + v_ax2 = Ts.axis.spatial(w_batch_norm, ax2) + v_ax3 = Ts.axis.spatial(c_batch_norm, ax3) Ts.reads(x[v_ax0, v_ax1, v_ax2, v_ax3], T_reshape[T.int64(0), v_ax1, T.int64(0), T.int64(0)]) Ts.writes(T_subtract_2[v_ax0, v_ax1, v_ax2, v_ax3]) T_subtract_2[v_ax0, v_ax1, v_ax2, v_ax3] = x[v_ax0, v_ax1, v_ax2, v_ax3] - T_reshape[T.int64(0), v_ax1, T.int64(0), T.int64(0)] - for ax0 in range(n): - for ax1 in range(h): - for ax2 in range(w): - for ax3 in range(c): + for ax0 in range(n_batch_norm): + for ax1 in range(h_batch_norm): + for ax2 in range(w_batch_norm): + for ax3 in range(c_batch_norm): with Ts.sblock("T_multiply"): - v_ax0 = Ts.axis.spatial(n, ax0) - v_ax1 = Ts.axis.spatial(h, ax1) - v_ax2 = Ts.axis.spatial(w, ax2) - v_ax3 = Ts.axis.spatial(c, ax3) + v_ax0 = Ts.axis.spatial(n_batch_norm, ax0) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) + v_ax2 = Ts.axis.spatial(w_batch_norm, ax2) + v_ax3 = Ts.axis.spatial(c_batch_norm, ax3) Ts.reads(T_subtract_1[v_ax0, v_ax1, v_ax2, v_ax3], T_subtract_2[v_ax0, v_ax1, v_ax2, v_ax3]) Ts.writes(T_multiply[v_ax0, v_ax1, v_ax2, v_ax3]) T_multiply[v_ax0, v_ax1, v_ax2, v_ax3] = T_subtract_1[v_ax0, v_ax1, v_ax2, v_ax3] * T_subtract_2[v_ax0, v_ax1, v_ax2, v_ax3] - for ax0 in range(h): - for k0 in range(n): - for k2 in range(w): - for k3 in range(c): + for ax0 in range(h_batch_norm): + for k0 in range(n_batch_norm): + for k2 in range(w_batch_norm): + for k3 in range(c_batch_norm): with Ts.sblock("T_multiply_red"): - v_ax0 = Ts.axis.spatial(h, ax0) - v_k0 = Ts.axis.reduce(n, k0) - v_k2 = Ts.axis.reduce(w, k2) - v_k3 = Ts.axis.reduce(c, k3) + v_ax0 = Ts.axis.spatial(h_batch_norm, ax0) + v_k0 = Ts.axis.reduce(n_batch_norm, k0) + v_k2 = Ts.axis.reduce(w_batch_norm, k2) + v_k3 = Ts.axis.reduce(c_batch_norm, k3) Ts.reads(T_multiply[v_k0, v_ax0, v_k2, v_k3]) Ts.writes(T_multiply_red[v_ax0]) with Ts.init(): T_multiply_red[v_ax0] = T.float32(0.0) T_multiply_red[v_ax0] = T_multiply_red[v_ax0] + T_multiply[v_k0, v_ax0, v_k2, v_k3] - for ax0 in range(h): + for ax0 in range(h_batch_norm): with Ts.sblock("T_divide_1"): - v_ax0 = Ts.axis.spatial(h, ax0) + v_ax0 = Ts.axis.spatial(h_batch_norm, ax0) Ts.reads(T_multiply_red[v_ax0]) Ts.writes(T_divide_1[v_ax0]) - T_divide_1[v_ax0] = T_multiply_red[v_ax0] / T.Cast("float32", n * w * c) + T_divide_1[v_ax0] = T_multiply_red[v_ax0] / T.Cast("float32", n_batch_norm * w_batch_norm * c_batch_norm) for ax0 in range(T.int64(1)): - for ax1 in range(h): + for ax1 in range(h_batch_norm): for ax2 in range(T.int64(1)): for ax3 in range(T.int64(1)): with Ts.sblock("T_reshape_1"): v_ax0 = Ts.axis.spatial(T.int64(1), ax0) - v_ax1 = Ts.axis.spatial(h, ax1) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) v_ax2 = Ts.axis.spatial(T.int64(1), ax2) v_ax3 = Ts.axis.spatial(T.int64(1), ax3) - Ts.reads(T_divide_1[(v_ax0 * h + v_ax1 + v_ax2 + v_ax3) % h]) + Ts.reads(T_divide_1[(v_ax0 * h_batch_norm + v_ax1 + v_ax2 + v_ax3) % h_batch_norm]) Ts.writes(T_reshape_1[v_ax0, v_ax1, v_ax2, v_ax3]) - T_reshape_1[v_ax0, v_ax1, v_ax2, v_ax3] = T_divide_1[(v_ax0 * h + v_ax1 + v_ax2 + v_ax3) % h] + T_reshape_1[v_ax0, v_ax1, v_ax2, v_ax3] = T_divide_1[(v_ax0 * h_batch_norm + v_ax1 + v_ax2 + v_ax3) % h_batch_norm] for ax0 in range(T.int64(1)): - for ax1 in range(h): + for ax1 in range(h_batch_norm): for ax2 in range(T.int64(1)): for ax3 in range(T.int64(1)): with Ts.sblock("T_add"): v_ax0 = Ts.axis.spatial(T.int64(1), ax0) - v_ax1 = Ts.axis.spatial(h, ax1) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) v_ax2 = Ts.axis.spatial(T.int64(1), ax2) v_ax3 = Ts.axis.spatial(T.int64(1), ax3) Ts.reads(T_reshape_1[v_ax0, v_ax1, v_ax2, v_ax3]) Ts.writes(T_add_3[v_ax0, v_ax1, v_ax2, v_ax3]) T_add_3[v_ax0, v_ax1, v_ax2, v_ax3] = T_reshape_1[v_ax0, v_ax1, v_ax2, v_ax3] + T.float32(1.0000000000000001e-05) for i0 in range(T.int64(1)): - for i1 in range(h): + for i1 in range(h_batch_norm): for i2 in range(T.int64(1)): for i3 in range(T.int64(1)): with Ts.sblock("compute"): v_i0 = Ts.axis.spatial(T.int64(1), i0) - v_i1 = Ts.axis.spatial(h, i1) + v_i1 = Ts.axis.spatial(h_batch_norm, i1) v_i2 = Ts.axis.spatial(T.int64(1), i2) v_i3 = Ts.axis.spatial(T.int64(1), i3) Ts.reads(T_add_3[v_i0, v_i1, v_i2, v_i3]) Ts.writes(compute[v_i0, v_i1, v_i2, v_i3]) compute[v_i0, v_i1, v_i2, v_i3] = T.sqrt(T_add_3[v_i0, v_i1, v_i2, v_i3]) - for ax0 in range(n): - for ax1 in range(h): - for ax2 in range(w): - for ax3 in range(c): + for ax0 in range(n_batch_norm): + for ax1 in range(h_batch_norm): + for ax2 in range(w_batch_norm): + for ax3 in range(c_batch_norm): with Ts.sblock("T_divide_2"): - v_ax0 = Ts.axis.spatial(n, ax0) - v_ax1 = Ts.axis.spatial(h, ax1) - v_ax2 = Ts.axis.spatial(w, ax2) - v_ax3 = Ts.axis.spatial(c, ax3) + v_ax0 = Ts.axis.spatial(n_batch_norm, ax0) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) + v_ax2 = Ts.axis.spatial(w_batch_norm, ax2) + v_ax3 = Ts.axis.spatial(c_batch_norm, ax3) Ts.reads(T_subtract[v_ax0, v_ax1, v_ax2, v_ax3], compute[T.int64(0), v_ax1, T.int64(0), T.int64(0)]) Ts.writes(T_divide_2[v_ax0, v_ax1, v_ax2, v_ax3]) T_divide_2[v_ax0, v_ax1, v_ax2, v_ax3] = T_subtract[v_ax0, v_ax1, v_ax2, v_ax3] / compute[T.int64(0), v_ax1, T.int64(0), T.int64(0)] for ax0 in range(T.int64(1)): - for ax1 in range(h): + for ax1 in range(h_batch_norm): for ax2 in range(T.int64(1)): for ax3 in range(T.int64(1)): with Ts.sblock("T_reshape_2"): v_ax0 = Ts.axis.spatial(T.int64(1), ax0) - v_ax1 = Ts.axis.spatial(h, ax1) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) v_ax2 = Ts.axis.spatial(T.int64(1), ax2) v_ax3 = Ts.axis.spatial(T.int64(1), ax3) - Ts.reads(gamma[(v_ax0 * h + v_ax1 + v_ax2 + v_ax3) % c]) + Ts.reads(gamma[(v_ax0 * h_batch_norm + v_ax1 + v_ax2 + v_ax3) % c_batch_norm]) Ts.writes(T_reshape_2[v_ax0, v_ax1, v_ax2, v_ax3]) - T_reshape_2[v_ax0, v_ax1, v_ax2, v_ax3] = gamma[(v_ax0 * h + v_ax1 + v_ax2 + v_ax3) % c] - for ax0 in range(n): - for ax1 in range(h): - for ax2 in range(w): - for ax3 in range(c): + T_reshape_2[v_ax0, v_ax1, v_ax2, v_ax3] = gamma[(v_ax0 * h_batch_norm + v_ax1 + v_ax2 + v_ax3) % c_batch_norm] + for ax0 in range(n_batch_norm): + for ax1 in range(h_batch_norm): + for ax2 in range(w_batch_norm): + for ax3 in range(c_batch_norm): with Ts.sblock("T_multiply_1"): - v_ax0 = Ts.axis.spatial(n, ax0) - v_ax1 = Ts.axis.spatial(h, ax1) - v_ax2 = Ts.axis.spatial(w, ax2) - v_ax3 = Ts.axis.spatial(c, ax3) + v_ax0 = Ts.axis.spatial(n_batch_norm, ax0) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) + v_ax2 = Ts.axis.spatial(w_batch_norm, ax2) + v_ax3 = Ts.axis.spatial(c_batch_norm, ax3) Ts.reads(T_divide_2[v_ax0, v_ax1, v_ax2, v_ax3], T_reshape_2[T.int64(0), v_ax1, T.int64(0), T.int64(0)]) Ts.writes(T_multiply_1[v_ax0, v_ax1, v_ax2, v_ax3]) T_multiply_1[v_ax0, v_ax1, v_ax2, v_ax3] = T_divide_2[v_ax0, v_ax1, v_ax2, v_ax3] * T_reshape_2[T.int64(0), v_ax1, T.int64(0), T.int64(0)] for ax0 in range(T.int64(1)): - for ax1 in range(h): + for ax1 in range(h_batch_norm): for ax2 in range(T.int64(1)): for ax3 in range(T.int64(1)): with Ts.sblock("T_reshape_3"): v_ax0 = Ts.axis.spatial(T.int64(1), ax0) - v_ax1 = Ts.axis.spatial(h, ax1) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) v_ax2 = Ts.axis.spatial(T.int64(1), ax2) v_ax3 = Ts.axis.spatial(T.int64(1), ax3) - Ts.reads(beta[(v_ax0 * h + v_ax1 + v_ax2 + v_ax3) % c]) + Ts.reads(beta[(v_ax0 * h_batch_norm + v_ax1 + v_ax2 + v_ax3) % c_batch_norm]) Ts.writes(T_reshape_3[v_ax0, v_ax1, v_ax2, v_ax3]) - T_reshape_3[v_ax0, v_ax1, v_ax2, v_ax3] = beta[(v_ax0 * h + v_ax1 + v_ax2 + v_ax3) % c] - for ax0 in range(n): - for ax1 in range(h): - for ax2 in range(w): - for ax3 in range(c): + T_reshape_3[v_ax0, v_ax1, v_ax2, v_ax3] = beta[(v_ax0 * h_batch_norm + v_ax1 + v_ax2 + v_ax3) % c_batch_norm] + for ax0 in range(n_batch_norm): + for ax1 in range(h_batch_norm): + for ax2 in range(w_batch_norm): + for ax3 in range(c_batch_norm): with Ts.sblock("T_add_1"): - v_ax0 = Ts.axis.spatial(n, ax0) - v_ax1 = Ts.axis.spatial(h, ax1) - v_ax2 = Ts.axis.spatial(w, ax2) - v_ax3 = Ts.axis.spatial(c, ax3) + v_ax0 = Ts.axis.spatial(n_batch_norm, ax0) + v_ax1 = Ts.axis.spatial(h_batch_norm, ax1) + v_ax2 = Ts.axis.spatial(w_batch_norm, ax2) + v_ax3 = Ts.axis.spatial(c_batch_norm, ax3) Ts.reads(T_multiply_1[v_ax0, v_ax1, v_ax2, v_ax3], T_reshape_3[T.int64(0), v_ax1, T.int64(0), T.int64(0)]) Ts.writes(T_add[v_ax0, v_ax1, v_ax2, v_ax3]) T_add[v_ax0, v_ax1, v_ax2, v_ax3] = T_multiply_1[v_ax0, v_ax1, v_ax2, v_ax3] + T_reshape_3[T.int64(0), v_ax1, T.int64(0), T.int64(0)] - for ax0 in range(c): + for ax0 in range(c_batch_norm): with Ts.sblock("T_multiply_2"): - v_ax0 = Ts.axis.spatial(c, ax0) + v_ax0 = Ts.axis.spatial(c_batch_norm, ax0) Ts.reads(moving_mean[v_ax0]) Ts.writes(T_multiply_2[v_ax0]) T_multiply_2[v_ax0] = T.float32(0.90000000000000002) * moving_mean[v_ax0] - for ax0 in range(h): + for ax0 in range(h_batch_norm): with Ts.sblock("T_multiply_3"): - v_ax0 = Ts.axis.spatial(h, ax0) + v_ax0 = Ts.axis.spatial(h_batch_norm, ax0) Ts.reads(T_divide[v_ax0]) Ts.writes(T_multiply_3[v_ax0]) T_multiply_3[v_ax0] = T.float32(0.10000000000000001) * T_divide[v_ax0] - for ax0 in range(T.max(c, h)): + for ax0 in range(T.max(c_batch_norm, h_batch_norm)): with Ts.sblock("T_add_2"): - v_ax0 = Ts.axis.spatial(T.max(c, h), ax0) + v_ax0 = Ts.axis.spatial(T.max(c_batch_norm, h_batch_norm), ax0) Ts.reads(T_multiply_2[v_ax0], T_multiply_3[v_ax0]) Ts.writes(T_add_1[v_ax0]) T_add_1[v_ax0] = T_multiply_2[v_ax0] + T_multiply_3[v_ax0] - for ax0 in range(c): + for ax0 in range(c_batch_norm): with Ts.sblock("T_multiply_4"): - v_ax0 = Ts.axis.spatial(c, ax0) + v_ax0 = Ts.axis.spatial(c_batch_norm, ax0) Ts.reads(moving_var[v_ax0]) Ts.writes(T_multiply_4[v_ax0]) T_multiply_4[v_ax0] = T.float32(0.90000000000000002) * moving_var[v_ax0] - for ax0 in range(h): + for ax0 in range(h_batch_norm): with Ts.sblock("T_multiply_5"): - v_ax0 = Ts.axis.spatial(h, ax0) + v_ax0 = Ts.axis.spatial(h_batch_norm, ax0) Ts.reads(T_divide_1[v_ax0]) Ts.writes(T_multiply_5[v_ax0]) T_multiply_5[v_ax0] = T.float32(0.10000000000000001) * T_divide_1[v_ax0] - for ax0 in range(T.max(c, h)): + for ax0 in range(T.max(c_batch_norm, h_batch_norm)): with Ts.sblock("T_add_3"): - v_ax0 = Ts.axis.spatial(T.max(c, h), ax0) + v_ax0 = Ts.axis.spatial(T.max(c_batch_norm, h_batch_norm), ax0) Ts.reads(T_multiply_4[v_ax0], T_multiply_5[v_ax0]) Ts.writes(T_add_2[v_ax0]) T_add_2[v_ax0] = T_multiply_4[v_ax0] + T_multiply_5[v_ax0] @R.function - def main(x: R.Tensor(("n", "h", "w", "c"), dtype="float32"), gamma: R.Tensor(("c",), dtype="float32"), beta: R.Tensor(("c",), dtype="float32"), moving_mean: R.Tensor(("c",), dtype="float32"), moving_var: R.Tensor(("c",), dtype="float32")) -> R.Tuple(R.Tensor(("n", "h", "w", "c"), dtype="float32"), R.Tensor(("T.max(c, h)",), dtype="float32"), R.Tensor(("T.max(c, h)",), dtype="float32")): - n = T.int64() - h = T.int64() - w = T.int64() - c = T.int64() + def main(x: R.Tensor((n_main, h_main, w_main, c_main), dtype="float32"), gamma: R.Tensor((c_main,), dtype="float32"), beta: R.Tensor((c_main,), dtype="float32"), moving_mean: R.Tensor((c_main,), dtype="float32"), moving_var: R.Tensor((c_main,), dtype="float32")) -> R.Tuple(R.Tensor((n_main, h_main, w_main, c_main), dtype="float32"), R.Tensor((T.max(c_main, h_main),), dtype="float32"), R.Tensor((T.max(c_main, h_main),), dtype="float32")): cls = Expected - gv = R.call_tir(cls.batch_norm, (x, gamma, beta, moving_mean, moving_var), out_ty=[R.Tensor((n, h, w, c), dtype="float32"), R.Tensor((T.max(c, h),), dtype="float32"), R.Tensor((T.max(c, h),), dtype="float32")]) + gv = R.call_tir(cls.batch_norm, (x, gamma, beta, moving_mean, moving_var), out_ty=[R.Tensor((n_main, h_main, w_main, c_main), dtype="float32"), R.Tensor((T.max(c_main, h_main),), dtype="float32"), R.Tensor((T.max(c_main, h_main),), dtype="float32")]) return gv mod = LegalizeOps()(BatchNorm) @@ -3002,39 +3063,43 @@ def main(x: R.Tensor((2, 3, 4, 5), dtype="float16"), gamma: R.Tensor((4, 5), dty def test_layer_norm_symbolic(): # fmt: off + n = T.dynamic("n") + s = T.dynamic("s") + f = T.dynamic("f") + @tvm.script.ir_module class LayerNorm: @R.function - def main(x: R.Tensor(("n", "s", "f"), "float32"), gamma: R.Tensor(("s", "f"), "float32"), beta: R.Tensor(("s", "f"), "float32")) -> R.Tensor(("n", "s", "f"), "float32"): - n = T.int64() - s = T.int64() - f = T.int64() + def main(x: R.Tensor((n, s, f), "float32"), gamma: R.Tensor((s, f), "float32"), beta: R.Tensor((s, f), "float32")) -> R.Tensor((n, s, f), "float32"): gv: R.Tensor((n, s, f), "float32") = R.nn.layer_norm(x, gamma, beta, axes=[1, 2]) return gv + n_main = T.dynamic("n") + s_main = T.dynamic("s") + f_main = T.dynamic("f") + n_layer_norm = T.dynamic("n") + s_layer_norm = T.dynamic("s") + f_layer_norm = T.dynamic("f") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("n", "s", "f"), "float32"), gamma: R.Tensor(("s", "f"), "float32"), beta: R.Tensor(("s", "f"), "float32")) -> R.Tensor(("n", "s", "f"), "float32"): - n = T.int64() - s = T.int64() - f = T.int64() - gv = R.call_tir(Expected.layer_norm, (x, gamma, beta), R.Tensor((n, s, f), dtype="float32")) + def main(x: R.Tensor((n_main, s_main, f_main), "float32"), gamma: R.Tensor((s_main, f_main), "float32"), beta: R.Tensor((s_main, f_main), "float32")) -> R.Tensor((n_main, s_main, f_main), "float32"): + gv = R.call_tir(Expected.layer_norm, (x, gamma, beta), R.Tensor((n_main, s_main, f_main), dtype="float32")) return gv @Ts.prim_func(private=True) def layer_norm(var_x: T.handle, var_gamma: T.handle, var_beta: T.handle, var_T_layer_norm: T.handle): T.func_attr({"tirx.noalias": True}) - n, s, f = T.int64(), T.int64(), T.int64() - x = T.match_buffer(var_x, (n, s, f)) - gamma = T.match_buffer(var_gamma, (s, f)) - beta = T.match_buffer(var_beta, (s, f)) - T_layer_norm = T.match_buffer(var_T_layer_norm, (n, s, f)) + x = T.match_buffer(var_x, (n_layer_norm, s_layer_norm, f_layer_norm)) + gamma = T.match_buffer(var_gamma, (s_layer_norm, f_layer_norm)) + beta = T.match_buffer(var_beta, (s_layer_norm, f_layer_norm)) + T_layer_norm = T.match_buffer(var_T_layer_norm, (n_layer_norm, s_layer_norm, f_layer_norm)) # with Ts.sblock("root"): - x_sum = Ts.sblock_alloc_buffer((n,)) - x_mean = Ts.sblock_alloc_buffer((n,)) - x_var_sum = Ts.sblock_alloc_buffer((n,)) - for ax0, k1, k2 in T.grid(n, s, f): + x_sum = Ts.sblock_alloc_buffer((n_layer_norm,)) + x_mean = Ts.sblock_alloc_buffer((n_layer_norm,)) + x_var_sum = Ts.sblock_alloc_buffer((n_layer_norm,)) + for ax0, k1, k2 in T.grid(n_layer_norm, s_layer_norm, f_layer_norm): with Ts.sblock("x_sum"): v_ax0, v_k1, v_k2 = Ts.axis.remap("SRR", [ax0, k1, k2]) Ts.reads(x[v_ax0, v_k1, v_k2]) @@ -3042,13 +3107,13 @@ def layer_norm(var_x: T.handle, var_gamma: T.handle, var_beta: T.handle, var_T_l with Ts.init(): x_sum[v_ax0] = T.float32(0.0) x_sum[v_ax0] = x_sum[v_ax0] + x[v_ax0, v_k1, v_k2] - for ax0 in range(n): + for ax0 in range(n_layer_norm): with Ts.sblock("x_mean"): - v_ax0 = Ts.axis.spatial(n, ax0) + v_ax0 = Ts.axis.spatial(n_layer_norm, ax0) Ts.reads(x_sum[v_ax0]) Ts.writes(x_mean[v_ax0]) - x_mean[v_ax0] = x_sum[v_ax0] / (T.Cast("float32", s) * T.Cast("float32", f)) - for ax0, k1, k2 in T.grid(n, s, f): + x_mean[v_ax0] = x_sum[v_ax0] / (T.Cast("float32", s_layer_norm) * T.Cast("float32", f_layer_norm)) + for ax0, k1, k2 in T.grid(n_layer_norm, s_layer_norm, f_layer_norm): with Ts.sblock("x_var_sum"): v_ax0, v_k1, v_k2 = Ts.axis.remap("SRR", [ax0, k1, k2]) Ts.reads(x[v_ax0, v_k1, v_k2], x_mean[v_ax0]) @@ -3056,12 +3121,12 @@ def layer_norm(var_x: T.handle, var_gamma: T.handle, var_beta: T.handle, var_T_l with Ts.init(): x_var_sum[v_ax0] = T.float32(0.0) x_var_sum[v_ax0] = x_var_sum[v_ax0] + (x[v_ax0, v_k1, v_k2] - x_mean[v_ax0]) * (x[v_ax0, v_k1, v_k2] - x_mean[v_ax0]) - for ax0, ax1, ax2 in T.grid(n, s, f): + for ax0, ax1, ax2 in T.grid(n_layer_norm, s_layer_norm, f_layer_norm): with Ts.sblock("T_layer_norm"): v_ax0, v_ax1, v_ax2 = Ts.axis.remap("SSS", [ax0, ax1, ax2]) Ts.reads(x[v_ax0, v_ax1, v_ax2], x_mean[v_ax0], x_var_sum[v_ax0], gamma[v_ax1, v_ax2], beta[v_ax1, v_ax2]) Ts.writes(T_layer_norm[v_ax0, v_ax1, v_ax2]) - T_layer_norm[v_ax0, v_ax1, v_ax2] = (x[v_ax0, v_ax1, v_ax2] - x_mean[v_ax0]) * T.rsqrt(x_var_sum[v_ax0] / (T.Cast("float32", s) * T.Cast("float32", f)) + T.float32(1.0000000000000001e-05)) * gamma[v_ax1, v_ax2] + beta[v_ax1, v_ax2] + T_layer_norm[v_ax0, v_ax1, v_ax2] = (x[v_ax0, v_ax1, v_ax2] - x_mean[v_ax0]) * T.rsqrt(x_var_sum[v_ax0] / (T.Cast("float32", s_layer_norm) * T.Cast("float32", f_layer_norm)) + T.float32(1.0000000000000001e-05)) * gamma[v_ax1, v_ax2] + beta[v_ax1, v_ax2] # fmt: on mod = LegalizeOps()(LayerNorm) tvm.ir.assert_structural_equal(mod, Expected) @@ -3222,43 +3287,49 @@ def group_norm(rxplaceholder: T.Buffer((T.int64(2), T.int64(4), T.int64(4), T.in def test_group_norm_symbolic(): # fmt: off + n = T.dynamic("n") + c = T.dynamic("c") + h = T.dynamic("h") + w = T.dynamic("w") + @tvm.script.ir_module class GroupNorm: @R.function - def main(s: R.Shape(["c"]), x: R.Tensor(("n", "4 * c", "h", "w"), "float32"), gamma: R.Tensor(("4 * c",), "float32"), beta: R.Tensor(("4 * c",), "float32")) -> R.Tensor(("n", "4 * c", "h", "w"), "float32"): - n = T.int64() - c = T.int64() - h = T.int64() - w = T.int64() + def main(s: R.Shape([c]), x: R.Tensor((n, 4 * c, h, w), "float32"), gamma: R.Tensor((4 * c,), "float32"), beta: R.Tensor((4 * c,), "float32")) -> R.Tensor((n, 4 * c, h, w), "float32"): gv: R.Tensor((n, 4 * c, h, w), "float32") = R.nn.group_norm(x, gamma, beta, num_groups=4, channel_axis=1, axes=[2, 3]) return gv + n_group_norm = T.dynamic("n") + h_group_norm = T.dynamic("h") + w_group_norm = T.dynamic("w") + n_main = T.dynamic("n") + c = T.dynamic("c") + h_main = T.dynamic("h") + w_main = T.dynamic("w") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) def group_norm(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_rxplaceholder_2: T.handle, c: T.int64, var_T_reshape: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() - h = T.int64() - w = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (n, T.int64(4) * c, h, w)) + rxplaceholder = T.match_buffer(var_rxplaceholder, (n_group_norm, T.int64(4) * c, h_group_norm, w_group_norm)) rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, (T.int64(4) * c,)) rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, (T.int64(4) * c,)) - T_reshape = T.match_buffer(var_T_reshape, (n, T.int64(4) * c, h, w)) + T_reshape = T.match_buffer(var_T_reshape, (n_group_norm, T.int64(4) * c, h_group_norm, w_group_norm)) # with Ts.sblock("root"): - T_reshape_1 = Ts.sblock_alloc_buffer((n, T.int64(4), T.int64(4) * c // T.int64(4), h, w)) - rxplaceholder_red_temp_v0 = Ts.sblock_alloc_buffer((n, T.int64(4))) - rxplaceholder_red_temp_v1 = Ts.sblock_alloc_buffer((n, T.int64(4))) + T_reshape_1 = Ts.sblock_alloc_buffer((n_group_norm, T.int64(4), T.int64(4) * c // T.int64(4), h_group_norm, w_group_norm)) + rxplaceholder_red_temp_v0 = Ts.sblock_alloc_buffer((n_group_norm, T.int64(4))) + rxplaceholder_red_temp_v1 = Ts.sblock_alloc_buffer((n_group_norm, T.int64(4))) T_reshape_2 = Ts.sblock_alloc_buffer((T.int64(4), T.int64(4) * c // T.int64(4))) T_reshape_3 = Ts.sblock_alloc_buffer((T.int64(4), T.int64(4) * c // T.int64(4))) - T_group_norm = Ts.sblock_alloc_buffer((n, T.int64(4), T.int64(4) * c // T.int64(4), h, w)) - for ax0, ax1, ax2, ax3, ax4 in T.grid(n, T.int64(4), c, h, w): + T_group_norm = Ts.sblock_alloc_buffer((n_group_norm, T.int64(4), T.int64(4) * c // T.int64(4), h_group_norm, w_group_norm)) + for ax0, ax1, ax2, ax3, ax4 in T.grid(n_group_norm, T.int64(4), c, h_group_norm, w_group_norm): with Ts.sblock("T_reshape"): v_ax0, v_ax1, v_ax2, v_ax3, v_ax4 = Ts.axis.remap("SSSSS", [ax0, ax1, ax2, ax3, ax4]) - Ts.reads(rxplaceholder[((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h + v_ax3) * w + v_ax4) // w // h // (c * T.int64(4)) % n, ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h + v_ax3) * w + v_ax4) // w // h % (c * T.int64(4)), ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h + v_ax3) * w + v_ax4) // w % h, ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h + v_ax3) * w + v_ax4) % w]) + Ts.reads(rxplaceholder[((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h_group_norm + v_ax3) * w_group_norm + v_ax4) // w_group_norm // h_group_norm // (c * T.int64(4)) % n_group_norm, ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h_group_norm + v_ax3) * w_group_norm + v_ax4) // w_group_norm // h_group_norm % (c * T.int64(4)), ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h_group_norm + v_ax3) * w_group_norm + v_ax4) // w_group_norm % h_group_norm, ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h_group_norm + v_ax3) * w_group_norm + v_ax4) % w_group_norm]) Ts.writes(T_reshape_1[v_ax0, v_ax1, v_ax2, v_ax3, v_ax4]) - T_reshape_1[v_ax0, v_ax1, v_ax2, v_ax3, v_ax4] = rxplaceholder[((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h + v_ax3) * w + v_ax4) // w // h // (c * T.int64(4)) % n, ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h + v_ax3) * w + v_ax4) // w // h % (c * T.int64(4)), ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h + v_ax3) * w + v_ax4) // w % h, ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h + v_ax3) * w + v_ax4) % w] - for ax0, ax1, k2, k3, k4 in T.grid(n, T.int64(4), c, h, w): + T_reshape_1[v_ax0, v_ax1, v_ax2, v_ax3, v_ax4] = rxplaceholder[((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h_group_norm + v_ax3) * w_group_norm + v_ax4) // w_group_norm // h_group_norm // (c * T.int64(4)) % n_group_norm, ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h_group_norm + v_ax3) * w_group_norm + v_ax4) // w_group_norm // h_group_norm % (c * T.int64(4)), ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h_group_norm + v_ax3) * w_group_norm + v_ax4) // w_group_norm % h_group_norm, ((((v_ax0 * T.int64(4) + v_ax1) * c + v_ax2) * h_group_norm + v_ax3) * w_group_norm + v_ax4) % w_group_norm] + for ax0, ax1, k2, k3, k4 in T.grid(n_group_norm, T.int64(4), c, h_group_norm, w_group_norm): with Ts.sblock("rxplaceholder_red_temp"): v_ax0, v_ax1, v_k2, v_k3, v_k4 = Ts.axis.remap("SSRRR", [ax0, ax1, k2, k3, k4]) Ts.reads(T_reshape_1[v_ax0, v_ax1, v_k2, v_k3, v_k4]) @@ -3282,26 +3353,22 @@ def group_norm(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_r Ts.reads(rxplaceholder_2[(v_ax0 * c + v_ax1) % (c * T.int64(4))]) Ts.writes(T_reshape_3[v_ax0, v_ax1]) T_reshape_3[v_ax0, v_ax1] = rxplaceholder_2[(v_ax0 * c + v_ax1) % (c * T.int64(4))] - for ax0, ax1, ax2, ax3, ax4 in T.grid(n, T.int64(4), c, h, w): + for ax0, ax1, ax2, ax3, ax4 in T.grid(n_group_norm, T.int64(4), c, h_group_norm, w_group_norm): with Ts.sblock("T_group_norm"): v_ax0, v_ax1, v_ax2, v_ax3, v_ax4 = Ts.axis.remap("SSSSS", [ax0, ax1, ax2, ax3, ax4]) Ts.reads(T_reshape_1[v_ax0, v_ax1, v_ax2, v_ax3, v_ax4], rxplaceholder_red_temp_v0[v_ax0, v_ax1], rxplaceholder_red_temp_v1[v_ax0, v_ax1], T_reshape_2[v_ax1, v_ax2], T_reshape_3[v_ax1, v_ax2]) Ts.writes(T_group_norm[v_ax0, v_ax1, v_ax2, v_ax3, v_ax4]) - T_group_norm[v_ax0, v_ax1, v_ax2, v_ax3, v_ax4] = (T_reshape_1[v_ax0, v_ax1, v_ax2, v_ax3, v_ax4] - rxplaceholder_red_temp_v0[v_ax0, v_ax1] / (T.Cast("float32", c) * T.Cast("float32", h) * T.Cast("float32", w))) * T.rsqrt(rxplaceholder_red_temp_v1[v_ax0, v_ax1] / (T.Cast("float32", c) * T.Cast("float32", h) * T.Cast("float32", w)) - rxplaceholder_red_temp_v0[v_ax0, v_ax1] / (T.Cast("float32", c) * T.Cast("float32", h) * T.Cast("float32", w)) * (rxplaceholder_red_temp_v0[v_ax0, v_ax1] / (T.Cast("float32", c) * T.Cast("float32", h) * T.Cast("float32", w))) + T.float32(1.0000000000000001e-05)) * T_reshape_2[v_ax1, v_ax2] + T_reshape_3[v_ax1, v_ax2] - for ax0, ax1, ax2, ax3 in T.grid(n, c * T.int64(4), h, w): + T_group_norm[v_ax0, v_ax1, v_ax2, v_ax3, v_ax4] = (T_reshape_1[v_ax0, v_ax1, v_ax2, v_ax3, v_ax4] - rxplaceholder_red_temp_v0[v_ax0, v_ax1] / (T.Cast("float32", c) * T.Cast("float32", h_group_norm) * T.Cast("float32", w_group_norm))) * T.rsqrt(rxplaceholder_red_temp_v1[v_ax0, v_ax1] / (T.Cast("float32", c) * T.Cast("float32", h_group_norm) * T.Cast("float32", w_group_norm)) - rxplaceholder_red_temp_v0[v_ax0, v_ax1] / (T.Cast("float32", c) * T.Cast("float32", h_group_norm) * T.Cast("float32", w_group_norm)) * (rxplaceholder_red_temp_v0[v_ax0, v_ax1] / (T.Cast("float32", c) * T.Cast("float32", h_group_norm) * T.Cast("float32", w_group_norm))) + T.float32(1.0000000000000001e-05)) * T_reshape_2[v_ax1, v_ax2] + T_reshape_3[v_ax1, v_ax2] + for ax0, ax1, ax2, ax3 in T.grid(n_group_norm, c * T.int64(4), h_group_norm, w_group_norm): with Ts.sblock("T_reshape_3"): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) - Ts.reads(T_group_norm[(((v_ax0 * c * T.int64(4) + v_ax1) * h + v_ax2) * w + v_ax3) // w // h // c // T.int64(4) % n, (((v_ax0 * c * T.int64(4) + v_ax1) * h + v_ax2) * w + v_ax3) // w // h // c % T.int64(4), (((v_ax0 * c * T.int64(4) + v_ax1) * h + v_ax2) * w + v_ax3) // w // h % c, (((v_ax0 * c * T.int64(4) + v_ax1) * h + v_ax2) * w + v_ax3) // w % h, (((v_ax0 * c * T.int64(4) + v_ax1) * h + v_ax2) * w + v_ax3) % w]) + Ts.reads(T_group_norm[(((v_ax0 * c * T.int64(4) + v_ax1) * h_group_norm + v_ax2) * w_group_norm + v_ax3) // w_group_norm // h_group_norm // c // T.int64(4) % n_group_norm, (((v_ax0 * c * T.int64(4) + v_ax1) * h_group_norm + v_ax2) * w_group_norm + v_ax3) // w_group_norm // h_group_norm // c % T.int64(4), (((v_ax0 * c * T.int64(4) + v_ax1) * h_group_norm + v_ax2) * w_group_norm + v_ax3) // w_group_norm // h_group_norm % c, (((v_ax0 * c * T.int64(4) + v_ax1) * h_group_norm + v_ax2) * w_group_norm + v_ax3) // w_group_norm % h_group_norm, (((v_ax0 * c * T.int64(4) + v_ax1) * h_group_norm + v_ax2) * w_group_norm + v_ax3) % w_group_norm]) Ts.writes(T_reshape[v_ax0, v_ax1, v_ax2, v_ax3]) - T_reshape[v_ax0, v_ax1, v_ax2, v_ax3] = T_group_norm[(((v_ax0 * c * T.int64(4) + v_ax1) * h + v_ax2) * w + v_ax3) // w // h // c // T.int64(4) % n, (((v_ax0 * c * T.int64(4) + v_ax1) * h + v_ax2) * w + v_ax3) // w // h // c % T.int64(4), (((v_ax0 * c * T.int64(4) + v_ax1) * h + v_ax2) * w + v_ax3) // w // h % c, (((v_ax0 * c * T.int64(4) + v_ax1) * h + v_ax2) * w + v_ax3) // w % h, (((v_ax0 * c * T.int64(4) + v_ax1) * h + v_ax2) * w + v_ax3) % w] + T_reshape[v_ax0, v_ax1, v_ax2, v_ax3] = T_group_norm[(((v_ax0 * c * T.int64(4) + v_ax1) * h_group_norm + v_ax2) * w_group_norm + v_ax3) // w_group_norm // h_group_norm // c // T.int64(4) % n_group_norm, (((v_ax0 * c * T.int64(4) + v_ax1) * h_group_norm + v_ax2) * w_group_norm + v_ax3) // w_group_norm // h_group_norm // c % T.int64(4), (((v_ax0 * c * T.int64(4) + v_ax1) * h_group_norm + v_ax2) * w_group_norm + v_ax3) // w_group_norm // h_group_norm % c, (((v_ax0 * c * T.int64(4) + v_ax1) * h_group_norm + v_ax2) * w_group_norm + v_ax3) // w_group_norm % h_group_norm, (((v_ax0 * c * T.int64(4) + v_ax1) * h_group_norm + v_ax2) * w_group_norm + v_ax3) % w_group_norm] @R.function - def main(s: R.Shape(["c"]), x: R.Tensor(("n", "4 * c", "h", "w"), dtype="float32"), gamma: R.Tensor(("4 * c",), dtype="float32"), beta: R.Tensor(("4 * c",), dtype="float32")) -> R.Tensor(("n", "4 * c", "h", "w"), dtype="float32"): - n = T.int64() - c = T.int64() - h = T.int64() - w = T.int64() - gv = R.call_tir(Expected.group_norm, (x, gamma, beta, c), out_ty=R.Tensor((n, 4 * c, h, w), dtype="float32")) + def main(s: R.Shape([c]), x: R.Tensor((n_main, 4 * c, h_main, w_main), dtype="float32"), gamma: R.Tensor((4 * c,), dtype="float32"), beta: R.Tensor((4 * c,), dtype="float32")) -> R.Tensor((n_main, 4 * c, h_main, w_main), dtype="float32"): + gv = R.call_tir(Expected.group_norm, (x, gamma, beta, c), out_ty=R.Tensor((n_main, 4 * c, h_main, w_main), dtype="float32")) return gv # fmt: on mod = LegalizeOps()(GroupNorm) @@ -3462,45 +3529,52 @@ def main(x: R.Tensor((2, 3, 4, 5), dtype="float16"), weight: R.Tensor((4, 5), dt def test_rms_norm_symbolic(): # fmt: off + n = T.dynamic("n") + s = T.dynamic("s") + f = T.dynamic("f") + @tvm.script.ir_module class RMSNorm: @R.function - def main(x: R.Tensor(("n", "s", "f"), "float32"), weight: R.Tensor(("s", "f"), "float32")) -> R.Tensor(("n", "s", "f"), "float32"): - n = T.int64() - s = T.int64() - f = T.int64() + def main(x: R.Tensor((n, s, f), "float32"), weight: R.Tensor((s, f), "float32")) -> R.Tensor((n, s, f), "float32"): gv: R.Tensor((n, s, f), "float32") = R.nn.rms_norm(x, weight, axes=[1, 2]) return gv + n_rms_norm = T.dynamic("n") + s_rms_norm = T.dynamic("s") + f_rms_norm = T.dynamic("f") + n_main = T.dynamic("n") + s_main = T.dynamic("s") + f_main = T.dynamic("f") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) def rms_norm(var_A: T.handle, var_B: T.handle, var_T_cast: T.handle): T.func_attr({"tirx.noalias": True}) - n, s, f = T.int64(), T.int64(), T.int64() - A = T.match_buffer(var_A, (n, s, f)) - B = T.match_buffer(var_B, (s, f)) - T_cast = T.match_buffer(var_T_cast, (n, s, f)) + A = T.match_buffer(var_A, (n_rms_norm, s_rms_norm, f_rms_norm)) + B = T.match_buffer(var_B, (s_rms_norm, f_rms_norm)) + T_cast = T.match_buffer(var_T_cast, (n_rms_norm, s_rms_norm, f_rms_norm)) # with Ts.sblock("root"): - T_cast_1 = Ts.sblock_alloc_buffer((n, s, f)) - T_multiply = Ts.sblock_alloc_buffer((n, s, f)) - T_multiply_red = Ts.sblock_alloc_buffer((n,)) - rsqrt = Ts.sblock_alloc_buffer((n,)) - T_cast_2 = Ts.sblock_alloc_buffer((s, f)) - T_rms_norm = Ts.sblock_alloc_buffer((n, s, f)) - for ax0, ax1, ax2 in T.grid(n, s, f): + T_cast_1 = Ts.sblock_alloc_buffer((n_rms_norm, s_rms_norm, f_rms_norm)) + T_multiply = Ts.sblock_alloc_buffer((n_rms_norm, s_rms_norm, f_rms_norm)) + T_multiply_red = Ts.sblock_alloc_buffer((n_rms_norm,)) + rsqrt = Ts.sblock_alloc_buffer((n_rms_norm,)) + T_cast_2 = Ts.sblock_alloc_buffer((s_rms_norm, f_rms_norm)) + T_rms_norm = Ts.sblock_alloc_buffer((n_rms_norm, s_rms_norm, f_rms_norm)) + for ax0, ax1, ax2 in T.grid(n_rms_norm, s_rms_norm, f_rms_norm): with Ts.sblock("T_cast"): v_ax0, v_ax1, v_ax2 = Ts.axis.remap("SSS", [ax0, ax1, ax2]) Ts.reads(A[v_ax0, v_ax1, v_ax2]) Ts.writes(T_cast_1[v_ax0, v_ax1, v_ax2]) T_cast_1[v_ax0, v_ax1, v_ax2] = A[v_ax0, v_ax1, v_ax2] - for ax0, ax1, ax2 in T.grid(n, s, f): + for ax0, ax1, ax2 in T.grid(n_rms_norm, s_rms_norm, f_rms_norm): with Ts.sblock("T_multiply"): v_ax0, v_ax1, v_ax2 = Ts.axis.remap("SSS", [ax0, ax1, ax2]) Ts.reads(T_cast_1[v_ax0, v_ax1, v_ax2]) Ts.writes(T_multiply[v_ax0, v_ax1, v_ax2]) T_multiply[v_ax0, v_ax1, v_ax2] = T_cast_1[v_ax0, v_ax1, v_ax2] * T_cast_1[v_ax0, v_ax1, v_ax2] - for ax0, k1, k2 in T.grid(n, s, f): + for ax0, k1, k2 in T.grid(n_rms_norm, s_rms_norm, f_rms_norm): with Ts.sblock("T_multiply_red"): v_ax0, v_k1, v_k2 = Ts.axis.remap("SRR", [ax0, k1, k2]) Ts.reads(T_multiply[v_ax0, v_k1, v_k2]) @@ -3508,25 +3582,25 @@ def rms_norm(var_A: T.handle, var_B: T.handle, var_T_cast: T.handle): with Ts.init(): T_multiply_red[v_ax0] = T.float32(0) T_multiply_red[v_ax0] = T_multiply_red[v_ax0] + T_multiply[v_ax0, v_k1, v_k2] - for ax0 in range(n): + for ax0 in range(n_rms_norm): with Ts.sblock("rsqrt"): - v_ax0 = Ts.axis.spatial(n, ax0) + v_ax0 = Ts.axis.spatial(n_rms_norm, ax0) Ts.reads(T_multiply_red[v_ax0]) Ts.writes(rsqrt[v_ax0]) - rsqrt[v_ax0] = T.rsqrt(T_multiply_red[v_ax0] / (T.Cast("float32", s) * T.Cast("float32", f)) + T.float32(1.0000000000000001e-05)) - for ax0, ax1 in T.grid(s, f): + rsqrt[v_ax0] = T.rsqrt(T_multiply_red[v_ax0] / (T.Cast("float32", s_rms_norm) * T.Cast("float32", f_rms_norm)) + T.float32(1.0000000000000001e-05)) + for ax0, ax1 in T.grid(s_rms_norm, f_rms_norm): with Ts.sblock("T_cast_1"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads(B[v_ax0, v_ax1]) Ts.writes(T_cast_2[v_ax0, v_ax1]) T_cast_2[v_ax0, v_ax1] = B[v_ax0, v_ax1] - for ax0, ax1, ax2 in T.grid(n, s, f): + for ax0, ax1, ax2 in T.grid(n_rms_norm, s_rms_norm, f_rms_norm): with Ts.sblock("T_rms_norm"): v_ax0, v_ax1, v_ax2 = Ts.axis.remap("SSS", [ax0, ax1, ax2]) Ts.reads(rsqrt[v_ax0], T_cast_1[v_ax0, v_ax1, v_ax2], T_cast_2[v_ax1, v_ax2]) Ts.writes(T_rms_norm[v_ax0, v_ax1, v_ax2]) T_rms_norm[v_ax0, v_ax1, v_ax2] = rsqrt[v_ax0] * T_cast_1[v_ax0, v_ax1, v_ax2] * T_cast_2[v_ax1, v_ax2] - for ax0, ax1, ax2 in T.grid(n, s, f): + for ax0, ax1, ax2 in T.grid(n_rms_norm, s_rms_norm, f_rms_norm): with Ts.sblock("T_cast_2"): v_ax0, v_ax1, v_ax2 = Ts.axis.remap("SSS", [ax0, ax1, ax2]) Ts.reads(T_rms_norm[v_ax0, v_ax1, v_ax2]) @@ -3534,12 +3608,9 @@ def rms_norm(var_A: T.handle, var_B: T.handle, var_T_cast: T.handle): T_cast[v_ax0, v_ax1, v_ax2] = T_rms_norm[v_ax0, v_ax1, v_ax2] @R.function - def main(x: R.Tensor(("n", "s", "f"), dtype="float32"), weight: R.Tensor(("s", "f"), dtype="float32")) -> R.Tensor(("n", "s", "f"), dtype="float32"): - n = T.int64() - s = T.int64() - f = T.int64() + def main(x: R.Tensor((n_main, s_main, f_main), dtype="float32"), weight: R.Tensor((s_main, f_main), dtype="float32")) -> R.Tensor((n_main, s_main, f_main), dtype="float32"): cls = Expected - gv = R.call_tir(cls.rms_norm, (x, weight), out_ty=R.Tensor((n, s, f), dtype="float32")) + gv = R.call_tir(cls.rms_norm, (x, weight), out_ty=R.Tensor((n_main, s_main, f_main), dtype="float32")) return gv # fmt: on mod = LegalizeOps()(RMSNorm) @@ -3681,9 +3752,9 @@ def attention_bias(q: T.Buffer((T.int64(4), T.int64(16), T.int64(32), T.int64(8) Ts.reads(T_transpose_2[((v_ax2 // T.int64(8) + v_ax1) // T.int64(8) + v_ax0) % T.int64(128) // T.int64(32), ((v_ax2 // T.int64(8) + v_ax1) // T.int64(8) + v_ax0) % T.int64(32), (v_ax2 // T.int64(8) + v_ax1) % T.int64(8), v_ax2 % T.int64(8)]) Ts.writes(T_reshape_1[v_ax0, v_ax1, v_ax2]) T_reshape_1[v_ax0, v_ax1, v_ax2] = T_transpose_2[((v_ax2 // T.int64(8) + v_ax1) // T.int64(8) + v_ax0) % T.int64(128) // T.int64(32), ((v_ax2 // T.int64(8) + v_ax1) // T.int64(8) + v_ax0) % T.int64(32), (v_ax2 // T.int64(8) + v_ax1) % T.int64(8), v_ax2 % T.int64(8)] - for b, i, j, k_1 in T.grid(T.int64(128), T.int64(16), T.int64(8), T.int64(8)): + for b_index, i, j, k_1 in T.grid(T.int64(128), T.int64(16), T.int64(8), T.int64(8)): with Ts.sblock("T_batch_matmul_NT"): - v_b, v_i, v_j, v_k = Ts.axis.remap("SSSR", [b, i, j, k_1]) + v_b, v_i, v_j, v_k = Ts.axis.remap("SSSR", [b_index, i, j, k_1]) Ts.reads(T_reshape[v_b, v_i, v_k], T_reshape_1[v_b, v_j, v_k]) Ts.writes(T_batch_matmul_NT[v_b, v_i, v_j]) Ts.sblock_attr({"layout_free_placeholders": [T_reshape_1]}) @@ -3772,9 +3843,9 @@ def attention_bias(q: T.Buffer((T.int64(4), T.int64(16), T.int64(32), T.int64(8) Ts.reads(T_transpose_3[((v_ax2 // T.int64(16) + v_ax1) // T.int64(8) + v_ax0) % T.int64(128) // T.int64(32), ((v_ax2 // T.int64(16) + v_ax1) // T.int64(8) + v_ax0) % T.int64(32), (v_ax2 // T.int64(16) + v_ax1) % T.int64(8), v_ax2 % T.int64(16)]) Ts.writes(T_reshape_4[v_ax0, v_ax1, v_ax2]) T_reshape_4[v_ax0, v_ax1, v_ax2] = T_transpose_3[((v_ax2 // T.int64(16) + v_ax1) // T.int64(8) + v_ax0) % T.int64(128) // T.int64(32), ((v_ax2 // T.int64(16) + v_ax1) // T.int64(8) + v_ax0) % T.int64(32), (v_ax2 // T.int64(16) + v_ax1) % T.int64(8), v_ax2 % T.int64(16)] - for b, i, j, k_1 in T.grid(T.int64(128), T.int64(16), T.int64(16), T.int64(8)): + for b_index, i, j, k_1 in T.grid(T.int64(128), T.int64(16), T.int64(16), T.int64(8)): with Ts.sblock("T_batch_matmul_NN"): - v_b, v_i, v_j, v_k = Ts.axis.remap("SSSR", [b, i, j, k_1]) + v_b, v_i, v_j, v_k = Ts.axis.remap("SSSR", [b_index, i, j, k_1]) Ts.reads(T_divide[v_b, v_i, v_k], T_reshape_4[v_b, v_k, v_j]) Ts.writes(T_batch_matmul_NN[v_b, v_i, v_j]) Ts.sblock_attr({"layout_free_placeholders": [T_reshape_4]}) @@ -3812,14 +3883,17 @@ def test_dynamic_attention(): legalization. """ + seq_len = T.dynamic("seq_len") + seq_len_kv = T.dynamic("seq_len_kv") + @tvm.script.ir_module class Attention: @R.function def main( - q: R.Tensor((4, "seq_len", 32, 8), "float32"), - k: R.Tensor((4, "seq_len_kv", 32, 8), "float32"), - v: R.Tensor((4, "seq_len_kv", 32, 16), "float32"), - bias: R.Tensor((4, 32, "seq_len", "seq_len_kv"), "float32"), + q: R.Tensor((4, seq_len, 32, 8), "float32"), + k: R.Tensor((4, seq_len_kv, 32, 8), "float32"), + v: R.Tensor((4, seq_len_kv, 32, 16), "float32"), + bias: R.Tensor((4, 32, seq_len, seq_len_kv), "float32"), ): gv = R.nn.attention( q, k, v, bias, scale=T.FloatImm("float32", 0.1), causal_mask="BottomRight" @@ -3835,27 +3909,31 @@ def test_dynamic_batch_attention(): fix https://github.com/apache/tvm/issues/19696 """ + batch_size = T.dynamic("batch_size") + @tvm.script.ir_module class Attention: @R.function def main( - q: R.Tensor(("batch_size", 16, 32, 8), "float32"), - k: R.Tensor(("batch_size", 8, 32, 8), "float32"), - v: R.Tensor(("batch_size", 8, 32, 16), "float32"), + q: R.Tensor((batch_size, 16, 32, 8), "float32"), + k: R.Tensor((batch_size, 8, 32, 8), "float32"), + v: R.Tensor((batch_size, 8, 32, 16), "float32"), ): gv = R.nn.attention(q, k, v) return gv LegalizeOps()(Attention) + batch_size = T.dynamic("batch_size") + @tvm.script.ir_module class AttentionBias: @R.function def main( - q: R.Tensor(("batch_size", 16, 32, 8), "float32"), - k: R.Tensor(("batch_size", 8, 32, 8), "float32"), - v: R.Tensor(("batch_size", 8, 32, 16), "float32"), - bias: R.Tensor(("batch_size", 32, 16, 8), "float32"), + q: R.Tensor((batch_size, 16, 32, 8), "float32"), + k: R.Tensor((batch_size, 8, 32, 8), "float32"), + v: R.Tensor((batch_size, 8, 32, 16), "float32"), + bias: R.Tensor((batch_size, 32, 16, 8), "float32"), ): gv = R.nn.attention( q, k, v, bias, scale=T.FloatImm("float32", 0.1), causal_mask="BottomRight" @@ -4017,27 +4095,30 @@ def nll_loss_without_weight(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.i def test_nll_no_batch(): # fmt: off + C = T.dynamic("C") + @tvm.script.ir_module class NLLLoss: @R.function - def main(predictions: R.Tensor(("C",), "float32"), targets: R.Tensor((), "int64"), weights: R.Tensor(("C",), "float32")) -> R.Tensor((), "float32"): + def main(predictions: R.Tensor((C,), "float32"), targets: R.Tensor((), "int64"), weights: R.Tensor((C,), "float32")) -> R.Tensor((), "float32"): gv = R.nn.nll_loss(predictions, targets, weights, reduction="mean", ignore_index=1) return gv + C_main = T.dynamic("C") + C_nll_loss = T.dynamic("C") + @tvm.script.ir_module class Expected: @R.function - def main(predictions: R.Tensor(("C",), dtype="float32"), targets: R.Tensor((), dtype="int64"), weights: R.Tensor(("C",), dtype="float32")) -> R.Tensor((), dtype="float32"): - C = T.int64() + def main(predictions: R.Tensor((C_main,), dtype="float32"), targets: R.Tensor((), dtype="int64"), weights: R.Tensor((C_main,), dtype="float32")) -> R.Tensor((), dtype="float32"): gv = R.call_tir(Expected.nll_loss, (predictions, targets, weights), out_ty=R.Tensor((), dtype="float32")) return gv @Ts.prim_func(private=True) def nll_loss(var_rxplaceholder: T.handle, rxplaceholder: T.Buffer((), "int64"), var_rxplaceholder_1: T.handle, T_divide: T.Buffer((), "float32")): T.func_attr({"tirx.noalias": True}) - C = T.int64() - rxplaceholder_1 = T.match_buffer(var_rxplaceholder, (C,)) - rxplaceholder_2 = T.match_buffer(var_rxplaceholder_1, (C,)) + rxplaceholder_1 = T.match_buffer(var_rxplaceholder, (C_nll_loss,)) + rxplaceholder_2 = T.match_buffer(var_rxplaceholder_1, (C_nll_loss,)) # with Ts.sblock("root"): nll_loss = Ts.sblock_alloc_buffer(()) nll_loss_1 = Ts.sblock_alloc_buffer(()) @@ -4064,17 +4145,31 @@ def nll_loss(var_rxplaceholder: T.handle, rxplaceholder: T.Buffer((), "int64"), def test_nll_loss_symbolic(): # fmt: off + N = T.dynamic("N") + C = T.dynamic("C") + d1 = T.dynamic("d1") + d2 = T.dynamic("d2") + @tvm.script.ir_module class NLLLoss: @R.function - def main(predictions: R.Tensor(("N", "C", "d1", "d2"), "float32"), targets: R.Tensor(("N", "d1", "d2"), "int64"), weights: R.Tensor(("C",), "float32")) -> R.Tensor((), "float32"): + def main(predictions: R.Tensor((N, C, d1, d2), "float32"), targets: R.Tensor((N, d1, d2), "int64"), weights: R.Tensor((C,), "float32")) -> R.Tensor((), "float32"): gv: R.Tensor((), "float32") = R.nn.nll_loss(predictions, targets, weights, reduction="mean", ignore_index=-1) return gv + N_main = T.dynamic("N") + C_main = T.dynamic("C") + d1_main = T.dynamic("d1") + d2_main = T.dynamic("d2") + C_nll_loss = T.dynamic("C") + N_nll_loss = T.dynamic("N") + d1_nll_loss = T.dynamic("d1") + d2_nll_loss = T.dynamic("d2") + @tvm.script.ir_module class Expected: @R.function - def main(predictions: R.Tensor(("N", "C", "d1", "d2"), dtype="float32"), targets: R.Tensor(("N", "d1", "d2"), dtype="int64"), weights: R.Tensor(("C",), dtype="float32")) -> R.Tensor((), dtype="float32"): + def main(predictions: R.Tensor((N_main, C_main, d1_main, d2_main), dtype="float32"), targets: R.Tensor((N_main, d1_main, d2_main), dtype="int64"), weights: R.Tensor((C_main,), dtype="float32")) -> R.Tensor((), dtype="float32"): # block 0 gv = R.call_tir(Expected.nll_loss, (predictions, targets, weights), R.Tensor((), dtype="float32")) return gv @@ -4083,26 +4178,22 @@ def main(predictions: R.Tensor(("N", "C", "d1", "d2"), dtype="float32"), targets def nll_loss(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_rxplaceholder_2: T.handle, T_divide: T.Buffer((), "float32"),): # function attr dict T.func_attr({"tirx.noalias": True}) - C = T.int64() - N = T.int64() - d1 = T.int64() - d2 = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [N, C, d1, d2], dtype="float32") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [N, d1, d2], dtype="int64") - rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, [C], dtype="float32") + rxplaceholder = T.match_buffer(var_rxplaceholder, [N_nll_loss, C_nll_loss, d1_nll_loss, d2_nll_loss], dtype="float32") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [N_nll_loss, d1_nll_loss, d2_nll_loss], dtype="int64") + rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, [C_nll_loss], dtype="float32") # body # with Ts.sblock("root") - nll_loss = Ts.sblock_alloc_buffer([N, d1, d2], dtype="float32") + nll_loss = Ts.sblock_alloc_buffer([N_nll_loss, d1_nll_loss, d2_nll_loss], dtype="float32") nll_loss_red = Ts.sblock_alloc_buffer([], dtype="float32") - nll_loss_1 = Ts.sblock_alloc_buffer([N, d1, d2], dtype="float32") + nll_loss_1 = Ts.sblock_alloc_buffer([N_nll_loss, d1_nll_loss, d2_nll_loss], dtype="float32") nll_loss_red_1 = Ts.sblock_alloc_buffer([], dtype="float32") - for ax0, ax1, ax2 in T.grid(N, d1, d2): + for ax0, ax1, ax2 in T.grid(N_nll_loss, d1_nll_loss, d2_nll_loss): with Ts.sblock("nll_loss"): v_ax0, v_ax1, v_ax2 = Ts.axis.remap("SSS", [ax0, ax1, ax2]) Ts.reads(rxplaceholder_1[v_ax0, v_ax1, v_ax2], rxplaceholder[v_ax0, rxplaceholder_1[v_ax0, v_ax1, v_ax2], v_ax1, v_ax2],rxplaceholder_2[rxplaceholder_1[v_ax0, v_ax1, v_ax2]],) Ts.writes(nll_loss[v_ax0, v_ax1, v_ax2]) nll_loss[v_ax0, v_ax1, v_ax2] = T.Select(rxplaceholder_1[v_ax0, v_ax1, v_ax2] != T.int64(-1), (T.float32(0) - rxplaceholder[v_ax0, rxplaceholder_1[v_ax0, v_ax1, v_ax2], v_ax1, v_ax2]) * rxplaceholder_2[rxplaceholder_1[v_ax0, v_ax1, v_ax2]], T.float32(0),) - for k0, k1, k2 in T.grid(N, d1, d2): + for k0, k1, k2 in T.grid(N_nll_loss, d1_nll_loss, d2_nll_loss): with Ts.sblock("nll_loss_red"): v_k0, v_k1, v_k2 = Ts.axis.remap("RRR", [k0, k1, k2]) Ts.reads(nll_loss[v_k0, v_k1, v_k2]) @@ -4110,13 +4201,13 @@ def nll_loss(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_rxp with Ts.init(): nll_loss_red[()] = T.float32(0) nll_loss_red[()] = nll_loss_red[()] + nll_loss[v_k0, v_k1, v_k2] - for ax0, ax1, ax2 in T.grid(N, d1, d2): + for ax0, ax1, ax2 in T.grid(N_nll_loss, d1_nll_loss, d2_nll_loss): with Ts.sblock("nll_loss_1"): v_ax0, v_ax1, v_ax2 = Ts.axis.remap("SSS", [ax0, ax1, ax2]) Ts.reads(rxplaceholder_1[v_ax0, v_ax1, v_ax2], rxplaceholder_2[rxplaceholder_1[v_ax0, v_ax1, v_ax2]],) Ts.writes(nll_loss_1[v_ax0, v_ax1, v_ax2]) nll_loss_1[v_ax0, v_ax1, v_ax2] = T.Select(rxplaceholder_1[v_ax0, v_ax1, v_ax2] != T.int64(-1), rxplaceholder_2[rxplaceholder_1[v_ax0, v_ax1, v_ax2]], T.float32(0),) - for k0, k1, k2 in T.grid(N, d1, d2): + for k0, k1, k2 in T.grid(N_nll_loss, d1_nll_loss, d2_nll_loss): with Ts.sblock("nll_loss_red_1"): v_k0, v_k1, v_k2 = Ts.axis.remap("RRR", [k0, k1, k2]) Ts.reads(nll_loss_1[v_k0, v_k1, v_k2]) diff --git a/tests/python/relax/test_transform_legalize_ops_qdq.py b/tests/python/relax/test_transform_legalize_ops_qdq.py index 6195eb2ce2da..098ccca98d90 100644 --- a/tests/python/relax/test_transform_legalize_ops_qdq.py +++ b/tests/python/relax/test_transform_legalize_ops_qdq.py @@ -132,29 +132,33 @@ def main( def test_quantize_fp32_to_int8_symbolic(): + n = T.dynamic("n") + @tvm.script.ir_module class Quantize: @R.function def main( - data: R.Tensor((4, "n"), "float32"), - scale: R.Tensor(("n",), "float32"), - zp: R.Tensor(("n",), "int8"), - ) -> R.Tensor((4, "n"), "int8"): + data: R.Tensor((4, n), "float32"), + scale: R.Tensor((n,), "float32"), + zp: R.Tensor((n,), "int8"), + ) -> R.Tensor((4, n), "int8"): out = R.quantize(data, scale, zp, axis=-1, out_dtype="int8") return out + n_quantize = T.dynamic("n") + n_main = T.dynamic("n") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) def quantize(var_A: T.handle, var_B: T.handle, var_C: T.handle, var_quantized: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() - A = T.match_buffer(var_A, (T.int64(4), n)) - B = T.match_buffer(var_B, (n,)) - C = T.match_buffer(var_C, (n,), "int8") - quantized = T.match_buffer(var_quantized, (T.int64(4), n), "int8") + A = T.match_buffer(var_A, (T.int64(4), n_quantize)) + B = T.match_buffer(var_B, (n_quantize,)) + C = T.match_buffer(var_C, (n_quantize,), "int8") + quantized = T.match_buffer(var_quantized, (T.int64(4), n_quantize), "int8") # with Ts.sblock("root"): - for i0, i1 in T.grid(T.int64(4), n): + for i0, i1 in T.grid(T.int64(4), n_quantize): with Ts.sblock("quantized"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(A[v_i0, v_i1], B[v_i1], C[v_i1]) @@ -172,12 +176,13 @@ def quantize(var_A: T.handle, var_B: T.handle, var_C: T.handle, var_quantized: T @R.function def main( - data: R.Tensor((4, "n"), dtype="float32"), - scale: R.Tensor(("n",), dtype="float32"), - zp: R.Tensor(("n",), dtype="int8"), - ) -> R.Tensor((4, "n"), dtype="int8"): - n = T.int64() - out = R.call_tir(Expected.quantize, (data, scale, zp), out_ty=R.Tensor((4, n), "int8")) + data: R.Tensor((4, n_main), dtype="float32"), + scale: R.Tensor((n_main,), dtype="float32"), + zp: R.Tensor((n_main,), dtype="int8"), + ) -> R.Tensor((4, n_main), dtype="int8"): + out = R.call_tir( + Expected.quantize, (data, scale, zp), out_ty=R.Tensor((4, n_main), "int8") + ) return out mod = LegalizeOps()(Quantize) @@ -414,17 +419,22 @@ def main(data: R.Tensor((2, 4), dtype="int8")) -> R.Tensor((2, 4), dtype="float3 def test_dequantize_int8_to_fp32_symbolic(): + n = T.dynamic("n") + @tvm.script.ir_module class Dequantize: @R.function def main( - data: R.Tensor((2, "n"), "int8"), - scale: R.Tensor(("n",), "float32"), - zp: R.Tensor(("n",), "int8"), - ) -> R.Tensor((2, "n"), "float32"): + data: R.Tensor((2, n), "int8"), + scale: R.Tensor((n,), "float32"), + zp: R.Tensor((n,), "int8"), + ) -> R.Tensor((2, n), "float32"): out = R.dequantize(data, scale, zp, axis=-1, out_dtype="float32") return out + n_dequantize = T.dynamic("n") + n_main = T.dynamic("n") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) @@ -432,13 +442,12 @@ def dequantize( var_A: T.handle, var_B: T.handle, var_C: T.handle, var_dequantized: T.handle ): T.func_attr({"tirx.noalias": True}) - n = T.int64() - A = T.match_buffer(var_A, (T.int64(2), n), "int8") - B = T.match_buffer(var_B, (n,)) - C = T.match_buffer(var_C, (n,), "int8") - dequantized = T.match_buffer(var_dequantized, (T.int64(2), n)) + A = T.match_buffer(var_A, (T.int64(2), n_dequantize), "int8") + B = T.match_buffer(var_B, (n_dequantize,)) + C = T.match_buffer(var_C, (n_dequantize,), "int8") + dequantized = T.match_buffer(var_dequantized, (T.int64(2), n_dequantize)) # with Ts.sblock("root"): - for i0, i1 in T.grid(T.int64(2), n): + for i0, i1 in T.grid(T.int64(2), n_dequantize): with Ts.sblock("dequantized"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(A[v_i0, v_i1], C[v_i1], B[v_i1]) @@ -450,13 +459,14 @@ def dequantize( @R.function def main( - data: R.Tensor((2, "n"), dtype="int8"), - scale: R.Tensor(("n",), dtype="float32"), - zp: R.Tensor(("n",), dtype="int8"), - ) -> R.Tensor((2, "n"), dtype="float32"): - n = T.int64() + data: R.Tensor((2, n_main), dtype="int8"), + scale: R.Tensor((n_main,), dtype="float32"), + zp: R.Tensor((n_main,), dtype="int8"), + ) -> R.Tensor((2, n_main), dtype="float32"): out = R.call_tir( - Expected.dequantize, (data, scale, zp), out_ty=R.Tensor((2, n), dtype="float32") + Expected.dequantize, + (data, scale, zp), + out_ty=R.Tensor((2, n_main), dtype="float32"), ) return out diff --git a/tests/python/relax/test_transform_legalize_ops_search_statistical.py b/tests/python/relax/test_transform_legalize_ops_search_statistical.py index 7492fc3ce914..bcce5d7a9f6a 100644 --- a/tests/python/relax/test_transform_legalize_ops_search_statistical.py +++ b/tests/python/relax/test_transform_legalize_ops_search_statistical.py @@ -61,37 +61,39 @@ def where(rxplaceholder: T.Buffer((T.int64(3), T.int64(2), T.int64(1)), "bool"), def test_where_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + @tvm.script.ir_module class Where: @R.function - def main(condition: R.Tensor(("a", "b", 1), "bool"), x: R.Tensor(("b", "c"), "float32"), y: R.Tensor(("b", 1), "float32")) -> R.Tensor(("a", "b", "c"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() + def main(condition: R.Tensor((a, b, 1), "bool"), x: R.Tensor((b, c), "float32"), y: R.Tensor((b, 1), "float32")) -> R.Tensor((a, b, c), "float32"): gv: R.Tensor((a, b, c), "float32") = R.where(condition, x, y) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + a_where = T.dynamic("a") + b_where = T.dynamic("b") + c_where = T.dynamic("c") + @tvm.script.ir_module class Expected: @R.function - def main(condition: R.Tensor(("a", "b", 1), "bool"), x: R.Tensor(("b", "c"), "float32"), y: R.Tensor(("b", 1), "float32")) -> R.Tensor(("a", "b", "c"), "float32"): - a = T.int64() - b = T.int64() - c = T.int64() - gv = R.call_tir(Expected.where, (condition, x, y), R.Tensor((a, b, c), dtype="float32")) + def main(condition: R.Tensor((a_main, b_main, 1), "bool"), x: R.Tensor((b_main, c_main), "float32"), y: R.Tensor((b_main, 1), "float32")) -> R.Tensor((a_main, b_main, c_main), "float32"): + gv = R.call_tir(Expected.where, (condition, x, y), R.Tensor((a_main, b_main, c_main), dtype="float32")) return gv @Ts.prim_func(private=True) def where(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle, var_rxplaceholder_2: T.handle, var_T_where: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, T.int64(1)], dtype="bool") - rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [b, c], dtype="float32") - rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, [b, T.int64(1)], dtype="float32") - T_where = T.match_buffer(var_T_where, [a, b, c], dtype="float32") - for i0, i1, i2 in T.grid(a, b, c): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_where, b_where, T.int64(1)], dtype="bool") + rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [b_where, c_where], dtype="float32") + rxplaceholder_2 = T.match_buffer(var_rxplaceholder_2, [b_where, T.int64(1)], dtype="float32") + T_where = T.match_buffer(var_T_where, [a_where, b_where, c_where], dtype="float32") + for i0, i1, i2 in T.grid(a_where, b_where, c_where): with Ts.sblock("T_where"): ax0, ax1, ax2 = Ts.axis.remap("SSS", [i0, i1, i2]) Ts.reads(rxplaceholder[ax0, ax1, T.int64(0)], rxplaceholder_1[ax1, ax2], rxplaceholder_2[ax1, T.int64(0)]) @@ -150,39 +152,43 @@ def argmax(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64( def test_argmax_symbolic(): # fmt: off + a = T.dynamic("a") + c = T.dynamic("c") + d = T.dynamic("d") + b = T.dynamic("b") + @tvm.script.ir_module class Argmax: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor(("a", 1, "c", "d"), "int64"): - a = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((a, 1, c, d), "int64"): gv: R.Tensor((a, 1, c, d), "int64") = R.argmax(x, axis=1, keepdims=True) return gv + a_main = T.dynamic("a") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + b_main = T.dynamic("b") + a_argmax = T.dynamic("a") + b_argmax = T.dynamic("b") + c_argmax = T.dynamic("c") + d_argmax = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), dtype="float32")) -> R.Tensor(("a", 1, "c", "d"), dtype="int64"): - a = T.int64() - c = T.int64() - d = T.int64() - gv = R.call_tir(Expected.argmax, (x,), out_ty=R.Tensor((a, 1, c, d), dtype="int64")) + def main(x: R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) -> R.Tensor((a_main, 1, c_main, d_main), dtype="int64"): + gv = R.call_tir(Expected.argmax, (x,), out_ty=R.Tensor((a_main, 1, c_main, d_main), dtype="int64")) return gv @Ts.prim_func(private=True) def argmax(var_rxplaceholder: T.handle, var_rxplaceholder_red: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (a, b, c, d)) - rxplaceholder_red = T.match_buffer(var_rxplaceholder_red, (a, T.int64(1), c, d), "int64") + rxplaceholder = T.match_buffer(var_rxplaceholder, (a_argmax, b_argmax, c_argmax, d_argmax)) + rxplaceholder_red = T.match_buffer(var_rxplaceholder_red, (a_argmax, T.int64(1), c_argmax, d_argmax), "int64") # with Ts.sblock("root"): - rxplaceholder_red_temp_v0 = Ts.sblock_alloc_buffer((a, T.int64(1), c, d), "int64") - rxplaceholder_red_temp_v1 = Ts.sblock_alloc_buffer((a, T.int64(1), c, d)) - for ax0, ax1, ax2, ax3, k1 in T.grid(a, T.int64(1), c, d, b): + rxplaceholder_red_temp_v0 = Ts.sblock_alloc_buffer((a_argmax, T.int64(1), c_argmax, d_argmax), "int64") + rxplaceholder_red_temp_v1 = Ts.sblock_alloc_buffer((a_argmax, T.int64(1), c_argmax, d_argmax)) + for ax0, ax1, ax2, ax3, k1 in T.grid(a_argmax, T.int64(1), c_argmax, d_argmax, b_argmax): with Ts.sblock("rxplaceholder_red_temp"): v_ax0, v_ax1, v_ax2, v_ax3, v_k1 = Ts.axis.remap("SSSSR", [ax0, ax1, ax2, ax3, k1]) Ts.reads(rxplaceholder[v_ax0, v_k1, v_ax2, v_ax3]) @@ -194,7 +200,7 @@ def argmax(var_rxplaceholder: T.handle, var_rxplaceholder_red: T.handle): v_rxplaceholder_red_temp_v1: T.let[T.float32] = T.Select(rxplaceholder_red_temp_v1[v_ax0, v_ax1, v_ax2, v_ax3] > rxplaceholder[v_ax0, v_k1, v_ax2, v_ax3], rxplaceholder_red_temp_v1[v_ax0, v_ax1, v_ax2, v_ax3], rxplaceholder[v_ax0, v_k1, v_ax2, v_ax3]) rxplaceholder_red_temp_v0[v_ax0, v_ax1, v_ax2, v_ax3] = v_rxplaceholder_red_temp_v0 rxplaceholder_red_temp_v1[v_ax0, v_ax1, v_ax2, v_ax3] = v_rxplaceholder_red_temp_v1 - for ax0, ax1, ax2, ax3 in T.grid(a, T.int64(1), c, d): + for ax0, ax1, ax2, ax3 in T.grid(a_argmax, T.int64(1), c_argmax, d_argmax): with Ts.sblock("rxplaceholder_red"): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) Ts.reads(rxplaceholder_red_temp_v0[v_ax0, v_ax1, v_ax2, v_ax3]) @@ -252,26 +258,36 @@ def main(x: R.Tensor((2, 3, 4, 5), dtype="float32")) -> R.Tensor((), dtype="int6 def test_argmin_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Argmin: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor((1, 1, 1, 1), "int64"): + def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((1, 1, 1, 1), "int64"): gv: R.Tensor((1, 1, 1, 1), "int64") = R.argmin(x, keepdims=True) return gv + a_argmin = T.dynamic("a") + b_argmin = T.dynamic("b") + c_argmin = T.dynamic("c") + d_argmin = T.dynamic("d") + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) def argmin(var_rxplaceholder: T.handle, rxplaceholder_red: T.Buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "int64")): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (a, b, c, d)) + rxplaceholder = T.match_buffer(var_rxplaceholder, (a_argmin, b_argmin, c_argmin, d_argmin)) rxplaceholder_red_temp_v0 = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "int64") rxplaceholder_red_temp_v1 = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1))) - for ax0, ax1, ax2, ax3, k0, k1, k2, k3 in T.grid(T.int64(1), T.int64(1), T.int64(1), T.int64(1), a, b, c, d): + for ax0, ax1, ax2, ax3, k0, k1, k2, k3 in T.grid(T.int64(1), T.int64(1), T.int64(1), T.int64(1), a_argmin, b_argmin, c_argmin, d_argmin): with Ts.sblock("rxplaceholder_red_temp"): v_ax0, v_ax1, v_ax2, v_ax3, v_k0, v_k1, v_k2, v_k3 = Ts.axis.remap("SSSSRRRR", [ax0, ax1, ax2, ax3, k0, k1, k2, k3]) Ts.reads(rxplaceholder[v_k0, v_k1, v_k2, v_k3]) @@ -279,7 +295,7 @@ def argmin(var_rxplaceholder: T.handle, rxplaceholder_red: T.Buffer((T.int64(1), with Ts.init(): rxplaceholder_red_temp_v0[v_ax0, v_ax1, v_ax2, v_ax3] = T.int64(-1) rxplaceholder_red_temp_v1[v_ax0, v_ax1, v_ax2, v_ax3] = T.max_value("float32") - v_rxplaceholder_red_temp_v0: T.let[T.int64] = T.Select(rxplaceholder_red_temp_v1[v_ax0, v_ax1, v_ax2, v_ax3] < rxplaceholder[v_k0, v_k1, v_k2, v_k3] or (rxplaceholder_red_temp_v1[v_ax0, v_ax1, v_ax2, v_ax3] == rxplaceholder[v_k0, v_k1, v_k2, v_k3] and rxplaceholder_red_temp_v0[v_ax0, v_ax1, v_ax2, v_ax3] < ((v_k0 * b + v_k1) * c + v_k2) * d + v_k3), rxplaceholder_red_temp_v0[v_ax0, v_ax1, v_ax2, v_ax3], ((v_k0 * b + v_k1) * c + v_k2) * d + v_k3) + v_rxplaceholder_red_temp_v0: T.let[T.int64] = T.Select(rxplaceholder_red_temp_v1[v_ax0, v_ax1, v_ax2, v_ax3] < rxplaceholder[v_k0, v_k1, v_k2, v_k3] or (rxplaceholder_red_temp_v1[v_ax0, v_ax1, v_ax2, v_ax3] == rxplaceholder[v_k0, v_k1, v_k2, v_k3] and rxplaceholder_red_temp_v0[v_ax0, v_ax1, v_ax2, v_ax3] < ((v_k0 * b_argmin + v_k1) * c_argmin + v_k2) * d_argmin + v_k3), rxplaceholder_red_temp_v0[v_ax0, v_ax1, v_ax2, v_ax3], ((v_k0 * b_argmin + v_k1) * c_argmin + v_k2) * d_argmin + v_k3) v_rxplaceholder_red_temp_v1: T.let[T.float32] = T.Select(rxplaceholder_red_temp_v1[v_ax0, v_ax1, v_ax2, v_ax3] < rxplaceholder[v_k0, v_k1, v_k2, v_k3], rxplaceholder_red_temp_v1[v_ax0, v_ax1, v_ax2, v_ax3], rxplaceholder[v_k0, v_k1, v_k2, v_k3]) rxplaceholder_red_temp_v0[v_ax0, v_ax1, v_ax2, v_ax3] = v_rxplaceholder_red_temp_v0 rxplaceholder_red_temp_v1[v_ax0, v_ax1, v_ax2, v_ax3] = v_rxplaceholder_red_temp_v1 @@ -291,7 +307,7 @@ def argmin(var_rxplaceholder: T.handle, rxplaceholder_red: T.Buffer((T.int64(1), rxplaceholder_red[v_ax0, v_ax1, v_ax2, v_ax3] = rxplaceholder_red_temp_v0[v_ax0, v_ax1, v_ax2, v_ax3] @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), dtype="float32")) -> R.Tensor((1, 1, 1, 1), dtype="int64"): + def main(x: R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) -> R.Tensor((1, 1, 1, 1), dtype="int64"): gv = R.call_tir(Expected.argmin, (x,), out_ty=R.Tensor((1, 1, 1, 1), dtype="int64")) return gv # fmt: on @@ -338,34 +354,40 @@ def max(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)) def test_max_symbolic(): # fmt: off + a = T.dynamic("a") + d = T.dynamic("d") + b = T.dynamic("b") + c = T.dynamic("c") + @tvm.script.ir_module class Max: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor(("a", "d"), "float32"): - a = T.int64() - d = T.int64() + def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((a, d), "float32"): gv: R.Tensor((a, d), "float32") = R.max(x, axis=[1, 2]) return gv + a_main = T.dynamic("a") + d_main = T.dynamic("d") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + a_max = T.dynamic("a") + b_max = T.dynamic("b") + c_max = T.dynamic("c") + d_max = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor(("a", "d"), "float32"): - a = T.int64() - d = T.int64() - gv = R.call_tir(Expected.max, (x,), R.Tensor((a, d), dtype="float32")) + def main(x: R.Tensor((a_main, b_main, c_main, d_main), "float32")) -> R.Tensor((a_main, d_main), "float32"): + gv = R.call_tir(Expected.max, (x,), R.Tensor((a_main, d_main), dtype="float32")) return gv @Ts.prim_func(private=True) def max(var_rxplaceholder: T.handle, var_rxplaceholder_red: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, c, d], dtype="float32") - rxplaceholder_red = T.match_buffer(var_rxplaceholder_red, [a, d], dtype="float32") - for i0, i1, i2, i3 in T.grid(a, d, b, c): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_max, b_max, c_max, d_max], dtype="float32") + rxplaceholder_red = T.match_buffer(var_rxplaceholder_red, [a_max, d_max], dtype="float32") + for i0, i1, i2, i3 in T.grid(a_max, d_max, b_max, c_max): with Ts.sblock("rxplaceholder_red"): ax0, ax1, k1, k2 = Ts.axis.remap("SSRR", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[ax0, k1, k2, ax1]) @@ -414,34 +436,40 @@ def min(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)) def test_min_symbolic(): # fmt: off + a = T.dynamic("a") + d = T.dynamic("d") + b = T.dynamic("b") + c = T.dynamic("c") + @tvm.script.ir_module class Min: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor(("a", 1, 1, "d"), "float32"): - a = T.int64() - d = T.int64() + def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((a, 1, 1, d), "float32"): gv: R.Tensor((a, 1, 1, d), "float32") = R.min(x, axis=[1, 2], keepdims=True) return gv + a_main = T.dynamic("a") + d_main = T.dynamic("d") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + a_min = T.dynamic("a") + b_min = T.dynamic("b") + c_min = T.dynamic("c") + d_min = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor(("a", 1, 1, "d"), "float32"): - a = T.int64() - d = T.int64() - gv = R.call_tir(Expected.min, (x,), R.Tensor((a, 1, 1, d), dtype="float32")) + def main(x: R.Tensor((a_main, b_main, c_main, d_main), "float32")) -> R.Tensor((a_main, 1, 1, d_main), "float32"): + gv = R.call_tir(Expected.min, (x,), R.Tensor((a_main, 1, 1, d_main), dtype="float32")) return gv @Ts.prim_func(private=True) def min(var_rxplaceholder: T.handle, var_rxplaceholder_red: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, c, d], dtype="float32") - rxplaceholder_red = T.match_buffer(var_rxplaceholder_red, [a, T.int64(1), T.int64(1), d], dtype="float32") - for i0, i1, i2, i3, i4, i5 in T.grid(a, T.int64(1), T.int64(1), d, b, c): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_min, b_min, c_min, d_min], dtype="float32") + rxplaceholder_red = T.match_buffer(var_rxplaceholder_red, [a_min, T.int64(1), T.int64(1), d_min], dtype="float32") + for i0, i1, i2, i3, i4, i5 in T.grid(a_min, T.int64(1), T.int64(1), d_min, b_min, c_min): with Ts.sblock("rxplaceholder_red"): ax0, ax1, ax2, ax3, k1, k2 = Ts.axis.remap("SSSSRR", [i0, i1, i2, i3, i4, i5]) Ts.reads(rxplaceholder[ax0, k1, k2, ax3]) @@ -490,29 +518,39 @@ def sum(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)) def test_sum_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Sum: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor((), "float32"): + def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((), "float32"): gv: R.Tensor((), "float32") = R.sum(x) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_sum = T.dynamic("a") + b_sum = T.dynamic("b") + c_sum = T.dynamic("c") + d_sum = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor((), "float32"): + def main(x: R.Tensor((a_main, b_main, c_main, d_main), "float32")) -> R.Tensor((), "float32"): gv = R.call_tir(Expected.sum, (x,), R.Tensor((), dtype="float32")) return gv @Ts.prim_func(private=True) def sum(var_rxplaceholder: T.handle, rxplaceholder_red: T.Buffer((), "float32")): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, c, d], dtype="float32") - for i0, i1, i2, i3 in T.grid(a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_sum, b_sum, c_sum, d_sum], dtype="float32") + for i0, i1, i2, i3 in T.grid(a_sum, b_sum, c_sum, d_sum): with Ts.sblock("rxplaceholder_red"): k0, k1, k2, k3 = Ts.axis.remap("RRRR", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[k0, k1, k2, k3]) @@ -594,29 +632,39 @@ def prod(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5) def test_prod_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Prod: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor((1, 1, 1, 1), "float32"): + def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((1, 1, 1, 1), "float32"): gv: R.Tensor((1, 1, 1, 1), "float32") = R.prod(x, keepdims=True) return gv + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + a_prod = T.dynamic("a") + b_prod = T.dynamic("b") + c_prod = T.dynamic("c") + d_prod = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor((1, 1, 1, 1), "float32"): + def main(x: R.Tensor((a_main, b_main, c_main, d_main), "float32")) -> R.Tensor((1, 1, 1, 1), "float32"): gv = R.call_tir(Expected.prod, (x,), R.Tensor((1, 1, 1, 1), dtype="float32")) return gv @Ts.prim_func(private=True) def prod(var_rxplaceholder: T.handle, rxplaceholder_red: T.Buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "float32")): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, c, d], dtype="float32") - for i0, i1, i2, i3, i4, i5, i6, i7 in T.grid(T.int64(1), T.int64(1), T.int64(1), T.int64(1), a, b, c, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_prod, b_prod, c_prod, d_prod], dtype="float32") + for i0, i1, i2, i3, i4, i5, i6, i7 in T.grid(T.int64(1), T.int64(1), T.int64(1), T.int64(1), a_prod, b_prod, c_prod, d_prod): with Ts.sblock("rxplaceholder_red"): ax0, ax1, ax2, ax3, k0, k1, k2, k3 = Ts.axis.remap("SSSSRRRR", [i0, i1, i2, i3, i4, i5, i6, i7]) Ts.reads(rxplaceholder[k0, k1, k2, k3]) @@ -736,35 +784,41 @@ def mean(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5) def test_mean_symbolic(): # fmt: off + b = T.dynamic("b") + c = T.dynamic("c") + a = T.dynamic("a") + d = T.dynamic("d") + @tvm.script.ir_module class Mean: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor(("b", "c"), "float32"): - b = T.int64() - c = T.int64() + def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((b, c), "float32"): gv: R.Tensor((b, c), "float32") = R.mean(x, [0, 3]) return gv + b_main = T.dynamic("b") + c_main = T.dynamic("c") + a_main = T.dynamic("a") + d_main = T.dynamic("d") + a_mean = T.dynamic("a") + b_mean = T.dynamic("b") + c_mean = T.dynamic("c") + d_mean = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), dtype="float32")) -> R.Tensor(("b", "c"), dtype="float32"): - b = T.int64() - c = T.int64() - gv = R.call_tir(Expected.mean, (x,), R.Tensor((b, c), dtype="float32")) + def main(x: R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) -> R.Tensor((b_main, c_main), dtype="float32"): + gv = R.call_tir(Expected.mean, (x,), R.Tensor((b_main, c_main), dtype="float32")) return gv @Ts.prim_func(private=True) def mean(var_rxplaceholder: T.handle, var_T_divide: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, c, d], dtype="float32") - T_divide = T.match_buffer(var_T_divide, [b, c], dtype="float32") - rxplaceholder_red = Ts.sblock_alloc_buffer([b, c], dtype="float32") - for i0, i1, i2, i3 in T.grid(b, c, a, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_mean, b_mean, c_mean, d_mean], dtype="float32") + T_divide = T.match_buffer(var_T_divide, [b_mean, c_mean], dtype="float32") + rxplaceholder_red = Ts.sblock_alloc_buffer([b_mean, c_mean], dtype="float32") + for i0, i1, i2, i3 in T.grid(b_mean, c_mean, a_mean, d_mean): with Ts.sblock("rxplaceholder_red"): ax0, ax1, k0, k3 = Ts.axis.remap("SSRR", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[k0, ax0, ax1, k3]) @@ -772,12 +826,12 @@ def mean(var_rxplaceholder: T.handle, var_T_divide: T.handle): with Ts.init(): rxplaceholder_red[ax0, ax1] = T.float32(0) rxplaceholder_red[ax0, ax1] = rxplaceholder_red[ax0, ax1] + rxplaceholder[k0, ax0, ax1, k3] - for i0, i1 in T.grid(b, c): + for i0, i1 in T.grid(b_mean, c_mean): with Ts.sblock("T_divide"): ax0, ax1 = Ts.axis.remap("SS", [i0, i1]) Ts.reads(rxplaceholder_red[ax0, ax1]) Ts.writes(T_divide[ax0, ax1]) - T_divide[ax0, ax1] = rxplaceholder_red[ax0, ax1] / T.Cast("float32", a * d) + T_divide[ax0, ax1] = rxplaceholder_red[ax0, ax1] / T.Cast("float32", a_mean * d_mean) # fmt: on mod = LegalizeOps()(Mean) @@ -941,28 +995,41 @@ def main(x: R.Tensor((2, 3, 4, 5), dtype="float32")) -> R.Tensor((), dtype="floa def test_std_symbolic(): # fmt: off + a = T.dynamic("a") + b = T.dynamic("b") + c = T.dynamic("c") + d = T.dynamic("d") + @tvm.script.ir_module class Std: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor((), "float32"): + def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((), "float32"): gv: R.Tensor((), "float32") = R.std(x) return gv + a_std = T.dynamic("a") + b_std = T.dynamic("b") + c_std = T.dynamic("c") + d_std = T.dynamic("d") + a_main = T.dynamic("a") + b_main = T.dynamic("b") + c_main = T.dynamic("c") + d_main = T.dynamic("d") + @I.ir_module class Expected: @Ts.prim_func(private=True) def std(var_rxplaceholder: T.handle, compute: T.Buffer((), "float32")): T.func_attr({"tirx.noalias": True}) - a, b, c, d = T.int64(), T.int64(), T.int64(), T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, (a, b, c, d)) + rxplaceholder = T.match_buffer(var_rxplaceholder, (a_std, b_std, c_std, d_std)) # with Ts.sblock("root"): rxplaceholder_red = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1))) T_divide = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1))) - T_subtract = Ts.sblock_alloc_buffer((a, b, c, d)) - T_multiply = Ts.sblock_alloc_buffer((a, b, c, d)) + T_subtract = Ts.sblock_alloc_buffer((a_std, b_std, c_std, d_std)) + T_multiply = Ts.sblock_alloc_buffer((a_std, b_std, c_std, d_std)) T_multiply_red = Ts.sblock_alloc_buffer(()) T_divide_1 = Ts.sblock_alloc_buffer(()) - for ax0, ax1, ax2, ax3, k0, k1, k2, k3 in T.grid(T.int64(1), T.int64(1), T.int64(1), T.int64(1), a, b, c, d): + for ax0, ax1, ax2, ax3, k0, k1, k2, k3 in T.grid(T.int64(1), T.int64(1), T.int64(1), T.int64(1), a_std, b_std, c_std, d_std): with Ts.sblock("rxplaceholder_red"): v_ax0, v_ax1, v_ax2, v_ax3, v_k0, v_k1, v_k2, v_k3 = Ts.axis.remap("SSSSRRRR", [ax0, ax1, ax2, ax3, k0, k1, k2, k3]) Ts.reads(rxplaceholder[v_k0, v_k1, v_k2, v_k3]) @@ -975,20 +1042,20 @@ def std(var_rxplaceholder: T.handle, compute: T.Buffer((), "float32")): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) Ts.reads(rxplaceholder_red[v_ax0, v_ax1, v_ax2, v_ax3]) Ts.writes(T_divide[v_ax0, v_ax1, v_ax2, v_ax3]) - T_divide[v_ax0, v_ax1, v_ax2, v_ax3] = rxplaceholder_red[v_ax0, v_ax1, v_ax2, v_ax3] / T.Cast("float32", a * b * c * d) - for ax0, ax1, ax2, ax3 in T.grid(a, b, c, d): + T_divide[v_ax0, v_ax1, v_ax2, v_ax3] = rxplaceholder_red[v_ax0, v_ax1, v_ax2, v_ax3] / T.Cast("float32", a_std * b_std * c_std * d_std) + for ax0, ax1, ax2, ax3 in T.grid(a_std, b_std, c_std, d_std): with Ts.sblock("T_subtract"): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) Ts.reads(rxplaceholder[v_ax0, v_ax1, v_ax2, v_ax3], T_divide[T.int64(0), T.int64(0), T.int64(0), T.int64(0)]) Ts.writes(T_subtract[v_ax0, v_ax1, v_ax2, v_ax3]) T_subtract[v_ax0, v_ax1, v_ax2, v_ax3] = rxplaceholder[v_ax0, v_ax1, v_ax2, v_ax3] - T_divide[T.int64(0), T.int64(0), T.int64(0), T.int64(0)] - for ax0, ax1, ax2, ax3 in T.grid(a, b, c, d): + for ax0, ax1, ax2, ax3 in T.grid(a_std, b_std, c_std, d_std): with Ts.sblock("T_multiply"): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) Ts.reads(T_subtract[v_ax0, v_ax1, v_ax2, v_ax3]) Ts.writes(T_multiply[v_ax0, v_ax1, v_ax2, v_ax3]) T_multiply[v_ax0, v_ax1, v_ax2, v_ax3] = T_subtract[v_ax0, v_ax1, v_ax2, v_ax3] * T_subtract[v_ax0, v_ax1, v_ax2, v_ax3] - for k0, k1, k2, k3 in T.grid(a, b, c, d): + for k0, k1, k2, k3 in T.grid(a_std, b_std, c_std, d_std): with Ts.sblock("T_multiply_red"): v_k0, v_k1, v_k2, v_k3 = Ts.axis.remap("RRRR", [k0, k1, k2, k3]) Ts.reads(T_multiply[v_k0, v_k1, v_k2, v_k3]) @@ -1000,7 +1067,7 @@ def std(var_rxplaceholder: T.handle, compute: T.Buffer((), "float32")): vi = Ts.axis.spatial(T.int64(1), T.int64(0)) Ts.reads(T_multiply_red[()]) Ts.writes(T_divide_1[()]) - T_divide_1[()] = T_multiply_red[()] / T.Cast("float32", a * b * c * d) + T_divide_1[()] = T_multiply_red[()] / T.Cast("float32", a_std * b_std * c_std * d_std) with Ts.sblock("compute"): vi = Ts.axis.spatial(T.int64(1), T.int64(0)) Ts.reads(T_divide_1[()]) @@ -1008,11 +1075,7 @@ def std(var_rxplaceholder: T.handle, compute: T.Buffer((), "float32")): compute[()] = T.sqrt(T_divide_1[()]) @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), dtype="float32")) -> R.Tensor((), dtype="float32"): - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() + def main(x: R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) -> R.Tensor((), dtype="float32"): cls = Expected gv = R.call_tir(cls.std, (x,), out_ty=R.Tensor((), dtype="float32")) return gv @@ -1094,39 +1157,45 @@ def variance(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int6 def test_variance_symbolic(): # fmt: off + b = T.dynamic("b") + c = T.dynamic("c") + a = T.dynamic("a") + d = T.dynamic("d") + @tvm.script.ir_module class Variance: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor((1, "b", "c", 1), "float32"): - b = T.int64() - c = T.int64() + def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((1, b, c, 1), "float32"): gv: R.Tensor((1, b, c, 1), "float32") = R.variance(x, [0, 3], keepdims=True) return gv + b_main = T.dynamic("b") + c_main = T.dynamic("c") + a_main = T.dynamic("a") + d_main = T.dynamic("d") + a_variance = T.dynamic("a") + b_variance = T.dynamic("b") + c_variance = T.dynamic("c") + d_variance = T.dynamic("d") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("a", "b", "c", "d"), "float32")) -> R.Tensor((1, "b", "c", 1), "float32"): - b = T.int64() - c = T.int64() - gv = R.call_tir(Expected.variance, (x,), R.Tensor((1, b, c, 1), dtype="float32")) + def main(x: R.Tensor((a_main, b_main, c_main, d_main), "float32")) -> R.Tensor((1, b_main, c_main, 1), "float32"): + gv = R.call_tir(Expected.variance, (x,), R.Tensor((1, b_main, c_main, 1), dtype="float32")) return gv @Ts.prim_func(private=True) def variance(var_rxplaceholder: T.handle, var_T_divide: T.handle): T.func_attr({"tirx.noalias": True}) - a = T.int64() - b = T.int64() - c = T.int64() - d = T.int64() - rxplaceholder = T.match_buffer(var_rxplaceholder, [a, b, c, d], dtype="float32") - T_divide = T.match_buffer(var_T_divide, [T.int64(1), b, c, T.int64(1)], dtype="float32") - rxplaceholder_red = Ts.sblock_alloc_buffer([T.int64(1), b, c, T.int64(1)], dtype="float32") - T_divide_1 = Ts.sblock_alloc_buffer([T.int64(1), b, c, T.int64(1)], dtype="float32") - T_subtract = Ts.sblock_alloc_buffer([a, b, c, d], dtype="float32") - T_multiply = Ts.sblock_alloc_buffer([a, b, c, d], dtype="float32") - T_multiply_red = Ts.sblock_alloc_buffer([T.int64(1), b, c, T.int64(1)], dtype="float32") - for i0, i1, i2, i3, i4, i5 in T.grid(T.int64(1), b, c, T.int64(1), a, d): + rxplaceholder = T.match_buffer(var_rxplaceholder, [a_variance, b_variance, c_variance, d_variance], dtype="float32") + T_divide = T.match_buffer(var_T_divide, [T.int64(1), b_variance, c_variance, T.int64(1)], dtype="float32") + rxplaceholder_red = Ts.sblock_alloc_buffer([T.int64(1), b_variance, c_variance, T.int64(1)], dtype="float32") + T_divide_1 = Ts.sblock_alloc_buffer([T.int64(1), b_variance, c_variance, T.int64(1)], dtype="float32") + T_subtract = Ts.sblock_alloc_buffer([a_variance, b_variance, c_variance, d_variance], dtype="float32") + T_multiply = Ts.sblock_alloc_buffer([a_variance, b_variance, c_variance, d_variance], dtype="float32") + T_multiply_red = Ts.sblock_alloc_buffer([T.int64(1), b_variance, c_variance, T.int64(1)], dtype="float32") + for i0, i1, i2, i3, i4, i5 in T.grid(T.int64(1), b_variance, c_variance, T.int64(1), a_variance, d_variance): with Ts.sblock("rxplaceholder_red"): ax0, ax1, ax2, ax3, k0, k3 = Ts.axis.remap("SSSSRR", [i0, i1, i2, i3, i4, i5]) Ts.reads(rxplaceholder[k0, ax1, ax2, k3]) @@ -1134,25 +1203,25 @@ def variance(var_rxplaceholder: T.handle, var_T_divide: T.handle): with Ts.init(): rxplaceholder_red[ax0, ax1, ax2, ax3] = T.float32(0) rxplaceholder_red[ax0, ax1, ax2, ax3] = rxplaceholder_red[ax0, ax1, ax2, ax3] + rxplaceholder[k0, ax1, ax2, k3] - for i0, i1, i2, i3 in T.grid(T.int64(1), b, c, T.int64(1)): + for i0, i1, i2, i3 in T.grid(T.int64(1), b_variance, c_variance, T.int64(1)): with Ts.sblock("T_divide"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder_red[ax0, ax1, ax2, ax3]) Ts.writes(T_divide_1[ax0, ax1, ax2, ax3]) - T_divide_1[ax0, ax1, ax2, ax3] = rxplaceholder_red[ax0, ax1, ax2, ax3] / T.Cast("float32", a * d) - for i0, i1, i2, i3 in T.grid(a, b, c, d): + T_divide_1[ax0, ax1, ax2, ax3] = rxplaceholder_red[ax0, ax1, ax2, ax3] / T.Cast("float32", a_variance * d_variance) + for i0, i1, i2, i3 in T.grid(a_variance, b_variance, c_variance, d_variance): with Ts.sblock("T_subtract"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(rxplaceholder[ax0, ax1, ax2, ax3], T_divide_1[T.int64(0), ax1, ax2, T.int64(0)]) Ts.writes(T_subtract[ax0, ax1, ax2, ax3]) T_subtract[ax0, ax1, ax2, ax3] = rxplaceholder[ax0, ax1, ax2, ax3] - T_divide_1[T.int64(0), ax1, ax2, T.int64(0)] - for i0, i1, i2, i3 in T.grid(a, b, c, d): + for i0, i1, i2, i3 in T.grid(a_variance, b_variance, c_variance, d_variance): with Ts.sblock("T_multiply"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(T_subtract[ax0, ax1, ax2, ax3]) Ts.writes(T_multiply[ax0, ax1, ax2, ax3]) T_multiply[ax0, ax1, ax2, ax3] = T_subtract[ax0, ax1, ax2, ax3] * T_subtract[ax0, ax1, ax2, ax3] - for i0, i1, i2, i3, i4, i5 in T.grid(T.int64(1), b, c, T.int64(1), a, d): + for i0, i1, i2, i3, i4, i5 in T.grid(T.int64(1), b_variance, c_variance, T.int64(1), a_variance, d_variance): with Ts.sblock("T_multiply_red"): ax0, ax1, ax2, ax3, k0, k3 = Ts.axis.remap("SSSSRR", [i0, i1, i2, i3, i4, i5]) Ts.reads(T_multiply[k0, ax1, ax2, k3]) @@ -1160,12 +1229,12 @@ def variance(var_rxplaceholder: T.handle, var_T_divide: T.handle): with Ts.init(): T_multiply_red[ax0, ax1, ax2, ax3] = T.float32(0) T_multiply_red[ax0, ax1, ax2, ax3] = T_multiply_red[ax0, ax1, ax2, ax3] + T_multiply[k0, ax1, ax2, k3] - for i0, i1, i2, i3 in T.grid(T.int64(1), b, c, T.int64(1)): + for i0, i1, i2, i3 in T.grid(T.int64(1), b_variance, c_variance, T.int64(1)): with Ts.sblock("T_divide_1"): ax0, ax1, ax2, ax3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) Ts.reads(T_multiply_red[ax0, ax1, ax2, ax3]) Ts.writes(T_divide[ax0, ax1, ax2, ax3]) - T_divide[ax0, ax1, ax2, ax3] = T_multiply_red[ax0, ax1, ax2, ax3] / T.Cast("float32", a * d) + T_divide[ax0, ax1, ax2, ax3] = T_multiply_red[ax0, ax1, ax2, ax3] / T.Cast("float32", a_variance * d_variance) # fmt: on mod = LegalizeOps()(Variance) diff --git a/tests/python/relax/test_transform_legalize_ops_unary.py b/tests/python/relax/test_transform_legalize_ops_unary.py index 4f8ee67d99ee..0ee8828beeff 100644 --- a/tests/python/relax/test_transform_legalize_ops_unary.py +++ b/tests/python/relax/test_transform_legalize_ops_unary.py @@ -52,18 +52,24 @@ def main(x: R.Tensor((2, 3), dtype)): def _test_symbolic_shape(name: str, relax_op: Callable, te_func: Callable, dtype: str): + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Before: @R.function - def main(x: R.Tensor(("m", "n"), dtype)): + def main(x: R.Tensor((m, n), dtype)): nonlocal dtype gv = relax_op(x) return gv + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), dtype)): + def main(x: R.Tensor((m, n), dtype)): nonlocal dtype gv = R.emit_te(te_func, x, primfunc_name_hint=f"tir_{name}") return gv diff --git a/tests/python/relax/test_transform_lift_transform_params.py b/tests/python/relax/test_transform_lift_transform_params.py index 71632963b925..9a5d66fc3252 100644 --- a/tests/python/relax/test_transform_lift_transform_params.py +++ b/tests/python/relax/test_transform_lift_transform_params.py @@ -1398,16 +1398,19 @@ def func1_transform_params( def test_symbolic_var_1(): + n = T.dynamic("n") + @tvm.script.ir_module class Before: @R.function - def main(shape: R.Shape(["n"])): + def main(shape: R.Shape([n])): R.func_attr({"num_input": 1}) - n = T.int64() with R.dataflow(): zeros = R.zeros((n, n), "float32") return shape + n = T.dynamic("n") + @I.ir_module class Expected: @R.function @@ -1418,9 +1421,8 @@ def main_transform_params(params: R.Tuple) -> R.Tuple: return R.tuple() @R.function - def main(shape: R.Shape(["n"])) -> R.Shape(["n"]): + def main(shape: R.Shape([n])) -> R.Shape([n]): R.func_attr({"num_input": 1}) - n = T.int64() with R.dataflow(): zeros: R.Tensor((n, n), dtype="float32") = R.zeros(R.shape([n, n]), dtype="float32") R.output() @@ -1432,14 +1434,16 @@ def main(shape: R.Shape(["n"])) -> R.Shape(["n"]): def test_symbolic_var_2(): + n_zeros = T.dynamic("n") + n_main = T.dynamic("n") + @I.ir_module class Before: @Ts.prim_func def zeros(var_T_full: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() - T_full = T.match_buffer(var_T_full, (n, n)) - for ax0, ax1 in T.grid(n, n): + T_full = T.match_buffer(var_T_full, (n_zeros, n_zeros)) + for ax0, ax1 in T.grid(n_zeros, n_zeros): with Ts.sblock("T_full"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads() @@ -1447,24 +1451,27 @@ def zeros(var_T_full: T.handle): T_full[v_ax0, v_ax1] = T.float32(0) @R.function - def main(shape: R.Shape(["n"])) -> R.Shape(["n"]): + def main(shape: R.Shape([n_main])) -> R.Shape([n_main]): R.func_attr({"num_input": 1}) - n = T.int64() cls = Before with R.dataflow(): - zeros = R.call_tir(cls.zeros, R.tuple(), out_ty=R.Tensor((n, n), dtype="float32")) + zeros = R.call_tir( + cls.zeros, R.tuple(), out_ty=R.Tensor((n_main, n_main), dtype="float32") + ) R.output() return shape + n_zeros = T.dynamic("n") + n_main = T.dynamic("n") + @I.ir_module class Expected: @Ts.prim_func def zeros(var_T_full: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() - T_full = T.match_buffer(var_T_full, (n, n)) + T_full = T.match_buffer(var_T_full, (n_zeros, n_zeros)) # with Ts.sblock("root"): - for ax0, ax1 in T.grid(n, n): + for ax0, ax1 in T.grid(n_zeros, n_zeros): with Ts.sblock("T_full"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) Ts.reads() @@ -1477,12 +1484,13 @@ def main_transform_params(params: R.Tuple) -> R.Tuple: return R.tuple() @R.function - def main(shape: R.Shape(["n"])) -> R.Shape(["n"]): + def main(shape: R.Shape([n_main])) -> R.Shape([n_main]): R.func_attr({"num_input": 1}) - n = T.int64() cls = Expected with R.dataflow(): - zeros = R.call_tir(cls.zeros, R.tuple(), out_ty=R.Tensor((n, n), dtype="float32")) + zeros = R.call_tir( + cls.zeros, R.tuple(), out_ty=R.Tensor((n_main, n_main), dtype="float32") + ) R.output() return shape @@ -1492,16 +1500,17 @@ def main(shape: R.Shape(["n"])) -> R.Shape(["n"]): def test_symbolic_var_from_shape(): + slice_index = T.dynamic("slice_index") + @I.ir_module class Before: @R.function def main( A: R.Tensor([16, 16], "int32"), B: R.Tensor([16, 16], "int32"), - shape: R.Shape(["slice_index"]), + shape: R.Shape([slice_index]), ) -> R.Tensor([16], "int32"): R.func_attr({"num_input": 1}) - slice_index = T.int64() cls = Before with R.dataflow(): B_slice = R.call_tir( @@ -1530,21 +1539,23 @@ def slice( vj = Ts.axis.remap("S", [j]) Output_Slice[vj] = Input_2d[slice_index, vj] + slice_index_main = T.dynamic("slice_index") + slice_index_main_transform_params = T.dynamic("slice_index") + @I.ir_module class Expected: @R.function def main( A: R.Tensor([16, 16], "int32"), - shape: R.Shape(["slice_index"]), + shape: R.Shape([slice_index_main]), B_slice: R.Tensor([16], "int32"), ) -> R.Tensor([16], "int32"): R.func_attr({"num_input": 1}) - slice_index = T.int64() cls = Expected with R.dataflow(): A_slice = R.call_tir( cls.slice, - [A, slice_index], + [A, slice_index_main], out_ty=R.Tensor([16], dtype="int32"), ) A_scale = R.multiply(A_slice, B_slice) @@ -1553,20 +1564,21 @@ def main( @R.function def main_transform_params( - params: R.Tuple(R.Tensor([16, 16], "int32"), R.Shape(["slice_index"])), + params: R.Tuple( + R.Tensor([16, 16], "int32"), R.Shape([slice_index_main_transform_params]) + ), ): R.func_attr({"num_input": 0}) - slice_index = T.int64() cls = Expected with R.dataflow(): B = params[0] # extra_symbolic_vars = params[1] B_slice = R.call_tir( cls.slice, - [B, slice_index], + [B, slice_index_main_transform_params], out_ty=R.Tensor([16], dtype="int32"), ) - output = (R.ShapeExpr([slice_index]), B_slice) + output = (R.ShapeExpr([slice_index_main_transform_params]), B_slice) R.output(output) return output @@ -1588,16 +1600,17 @@ def slice( def test_symbolic_var_in_param_shape(): + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Before: @R.function def main( - x: R.Tensor((1, 16, 224, "n"), "float32"), - w1: R.Tensor((16, "m", 3, 3), "float32"), - w2: R.Tensor((16, "m", 3, 3), "float32"), - ) -> R.Tensor((1, 16, 224, "n"), "float32"): - m = T.int64() - n = T.int64() + x: R.Tensor((1, 16, 224, n), "float32"), + w1: R.Tensor((16, m, 3, 3), "float32"), + w2: R.Tensor((16, m, 3, 3), "float32"), + ) -> R.Tensor((1, 16, 224, n), "float32"): R.func_attr({"num_input": 1}) with R.dataflow(): zeros = R.zeros((n, n), "float32") @@ -1609,38 +1622,42 @@ def main( R.output(conv2) return conv2 + m_main_transform_params = T.dynamic("m") + n = T.dynamic("n") + m_main = T.dynamic("m") + @I.ir_module class Expected: @R.function def main_transform_params( params: R.Tuple( - R.Tensor((16, "m", 3, 3), dtype="float32"), - R.Tensor((16, "m", 3, 3), dtype="float32"), + R.Tensor((16, m_main_transform_params, 3, 3), dtype="float32"), + R.Tensor((16, m_main_transform_params, 3, 3), dtype="float32"), ), ) -> R.Tuple( - R.Tensor((16, "m", 3, 3), dtype="float32"), R.Tensor((16, "m", 3, 3), dtype="float32") + R.Tensor((16, m_main_transform_params, 3, 3), dtype="float32"), + R.Tensor((16, m_main_transform_params, 3, 3), dtype="float32"), ): R.func_attr({"num_input": 0}) - m = T.int64() with R.dataflow(): - lv1: R.Tensor((16, m, 3, 3), dtype="float32") = params[0] - lv2: R.Tensor((16, m, 3, 3), dtype="float32") = R.add(lv1, R.const(1, "float32")) - lv: R.Tensor((16, m, 3, 3), dtype="float32") = params[1] + lv1: R.Tensor((16, m_main_transform_params, 3, 3), dtype="float32") = params[0] + lv2: R.Tensor((16, m_main_transform_params, 3, 3), dtype="float32") = R.add( + lv1, R.const(1, "float32") + ) + lv: R.Tensor((16, m_main_transform_params, 3, 3), dtype="float32") = params[1] gv: R.Tuple( - R.Tensor((16, m, 3, 3), dtype="float32"), - R.Tensor((16, m, 3, 3), dtype="float32"), + R.Tensor((16, m_main_transform_params, 3, 3), dtype="float32"), + R.Tensor((16, m_main_transform_params, 3, 3), dtype="float32"), ) = (lv, lv2) R.output(gv) return gv @R.function def main( - x: R.Tensor((1, 16, 224, "n"), dtype="float32"), - transformed_param_0: R.Tensor((16, "m", 3, 3), dtype="float32"), - transformed_param_1: R.Tensor((16, "m", 3, 3), dtype="float32"), - ) -> R.Tensor((1, 16, 224, "n"), dtype="float32"): - n = T.int64() - m = T.int64() + x: R.Tensor((1, 16, 224, n), dtype="float32"), + transformed_param_0: R.Tensor((16, m_main, 3, 3), dtype="float32"), + transformed_param_1: R.Tensor((16, m_main, 3, 3), dtype="float32"), + ) -> R.Tensor((1, 16, 224, n), dtype="float32"): R.func_attr({"num_input": 1}) with R.dataflow(): zeros: R.Tensor((n, n), dtype="float32") = R.zeros(R.shape([n, n]), dtype="float32") @@ -1687,15 +1704,16 @@ def test_symbolic_var_defined_in_params_but_used_in_weights(): not variable definitions. """ + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Before: @R.function def main( - x: R.Tensor(["m", "n"], "float32"), - weight: R.Tensor(["m * n"], "float32"), - ) -> R.Tensor(["m", "n"], "float32"): - m = T.int64() - n = T.int64() + x: R.Tensor([m, n], "float32"), + weight: R.Tensor([m * n], "float32"), + ) -> R.Tensor([m, n], "float32"): R.func_attr({"num_input": 1}) with R.dataflow(): weight = R.add(weight, R.const(1, "float32")) @@ -1704,14 +1722,17 @@ def main( R.output(output) return output + k = T.dynamic("k") + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class Expected: @R.function - def main_transform_params(params: R.Tuple(R.Tensor(("k",), dtype="float32"))) -> R.Tuple( + def main_transform_params(params: R.Tuple(R.Tensor((k,), dtype="float32"))) -> R.Tuple( R.Tensor(dtype="float32", ndim=1) ): R.func_attr({"num_input": 0}) - k = T.int64() with R.dataflow(): lv: R.Tensor((k,), dtype="float32") = params[0] gv: R.Tuple(R.Tensor((k,), dtype="float32")) = (lv,) @@ -1720,11 +1741,9 @@ def main_transform_params(params: R.Tuple(R.Tensor(("k",), dtype="float32"))) -> @R.function def main( - x: R.Tensor(("m", "n"), dtype="float32"), + x: R.Tensor((m, n), dtype="float32"), transformed_param_0: R.Tensor(dtype="float32", ndim=1), - ) -> R.Tensor(("m", "n"), dtype="float32"): - m = T.int64() - n = T.int64() + ) -> R.Tensor((m, n), dtype="float32"): R.func_attr({"num_input": 1}) with R.dataflow(): lv: R.Tensor(dtype="float32", ndim=1) = transformed_param_0 @@ -1796,14 +1815,17 @@ def main_transform_params(params: R.Tuple([R.Tensor([16], "int32")])): def test_lift_transform_is_idempotent(shared_transform): """Multiple applicates of LiftTransformParams are allowed""" + batch_size = T.dynamic("batch_size") + lora_rank = T.dynamic("lora_rank") + @I.ir_module class Module: @R.function def main( - state: R.Tensor(["batch_size", 4096], "float16"), + state: R.Tensor([batch_size, 4096], "float16"), base_weights: R.Tensor([4096, 4096], "float16"), - lora_A: R.Tensor([4096, "lora_rank"], "float16"), - lora_B: R.Tensor(["lora_rank", 4096], "float16"), + lora_A: R.Tensor([4096, lora_rank], "float16"), + lora_B: R.Tensor([lora_rank, 4096], "float16"), ): R.func_attr({"num_input": 1}) folded_weights = base_weights + R.matmul(lora_A, lora_B) @@ -1825,14 +1847,18 @@ def test_lift_transform_when_one_already_exists(): """If the module already contains `transform_params`, the functions are composed together""" + batch_size = T.dynamic("batch_size") + lora_rank_main = T.dynamic("lora_rank") + lora_rank_main_transform_params = T.dynamic("lora_rank") + @I.ir_module class Module: @R.function def main( - state: R.Tensor(["batch_size", 4096], "float16"), + state: R.Tensor([batch_size, 4096], "float16"), base_weights: R.Tensor([4096, 4096], "float16"), - lora_A: R.Tensor([4096, "lora_rank"], "float16"), - lora_B: R.Tensor(["lora_rank", 4096], "float16"), + lora_A: R.Tensor([4096, lora_rank_main], "float16"), + lora_B: R.Tensor([lora_rank_main, 4096], "float16"), ): R.func_attr({"num_input": 1}) folded_weights = base_weights + R.matmul(lora_A, lora_B) @@ -1843,8 +1869,8 @@ def main( def main_transform_params( model_params: R.Tuple( R.Tensor([4096, 4096], "float16"), - R.Tensor([4096, "lora_rank"], "float16"), - R.Tensor(["lora_rank", 4096], "float16"), + R.Tensor([4096, lora_rank_main_transform_params], "float16"), + R.Tensor([lora_rank_main_transform_params, 4096], "float16"), ), ): R.func_attr({"num_input": 0}) @@ -1869,7 +1895,7 @@ class Before: def main( x: R.Tensor(dtype="float32", ndim=1), extent: T.int64, - weight: R.Tensor(["extent"], "float32"), + weight: 'R.Tensor([extent], "float32")', ): R.func_attr({"num_input": 1}) transformed = R.multiply(weight, weight) diff --git a/tests/python/relax/test_transform_lower_gpu_ipc_alloc_storage.py b/tests/python/relax/test_transform_lower_gpu_ipc_alloc_storage.py index f9be099f6f7b..45c99848b6c0 100644 --- a/tests/python/relax/test_transform_lower_gpu_ipc_alloc_storage.py +++ b/tests/python/relax/test_transform_lower_gpu_ipc_alloc_storage.py @@ -24,12 +24,13 @@ def test_alloc_storage(): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Module: @R.function(pure=False) - def main(shape: R.Shape(["m", "n"])): # type: ignore - m = T.int64() - n = T.int64() + def main(shape: R.Shape([m, n])): # type: ignore storage: R.Any = R.memory.alloc_storage( R.shape([m, n]), R.prim_value(0), R.str("ipc_memory"), R.dtype("float16") ) @@ -38,12 +39,13 @@ def main(shape: R.Shape(["m", "n"])): # type: ignore ) return alloc + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Expected: @R.function(pure=False) - def main(shape: R.Shape(["m", "n"])): # type: ignore - m = T.int64() - n = T.int64() + def main(shape: R.Shape([m, n])): # type: ignore storage: R.Any = R.call_packed( "runtime.disco.cuda_ipc.alloc_storage", R.shape([m, n]), @@ -60,23 +62,25 @@ def main(shape: R.Shape(["m", "n"])): # type: ignore def test_builtin_alloc_tensor(): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Module: @R.function(pure=False) - def main(shape: R.Shape(["m", "n"])): # type: ignore - m = T.int64() - n = T.int64() + def main(shape: R.Shape([m, n])): # type: ignore tensor: R.Any = R.builtin.alloc_tensor( R.shape([m, n]), R.dtype("float16"), R.prim_value(0), R.str("ipc_memory") ) return tensor + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Expected: @R.function(pure=False) - def main(shape: R.Shape(["m", "n"])): # type: ignore - m = T.int64() - n = T.int64() + def main(shape: R.Shape([m, n])): # type: ignore gv: R.Any = R.call_packed( "runtime.disco.cuda_ipc.alloc_storage", R.shape([m, n]), diff --git a/tests/python/relax/test_transform_normalize.py b/tests/python/relax/test_transform_normalize.py index d136c80df38e..26cdcb78c9d5 100644 --- a/tests/python/relax/test_transform_normalize.py +++ b/tests/python/relax/test_transform_normalize.py @@ -20,15 +20,15 @@ import tvm import tvm.script import tvm.testing -from tvm import relax, tirx +from tvm import relax from tvm.ir.base import assert_structural_equal from tvm.script import relax as R from tvm.script import tirx as T def test_normalize_function(): - m = tirx.Var("m", "int64") - n = tirx.Var("n", "int64") + m = T.dynamic("m", "int64") + n = T.dynamic("n", "int64") x = relax.Var("x", R.Tensor([m, n], "float16")) # Note: the parser automatically normalize the IR written in TVMScript, @@ -45,7 +45,7 @@ def test_normalize_function(): after_mod = relax.transform.Normalize()(before_mod) @R.function(private=True) - def expected(x: R.Tensor(("m", "n"), "float16")) -> R.Tensor(dtype="float16", ndim=2): + def expected(x: R.Tensor((m, n), "float16")) -> R.Tensor(dtype="float16", ndim=2): gv = R.add(x, x) gv1 = R.add(x, x) return R.multiply(gv, gv1) @@ -118,11 +118,13 @@ def f(x: R.Tensor(dtype="float32")): after_mod = relax.transform.Normalize()(before_mod) assert_structural_equal(before_mod, after_mod, map_free_vars=True) + m = T.dynamic("m") + n = T.dynamic("n") + @tvm.script.ir_module class ANFMod2: @R.function - def foo(x: R.Tensor(("m", "n"), "float32")): - m, n = T.int64(), T.int64() + def foo(x: R.Tensor((m, n), "float32")): with R.dataflow(): lv0 = R.call_dps_packed("test.op.identity", (x,), R.Tensor((m, n), dtype="float32")) gv0 = R.call_dps_packed( diff --git a/tests/python/relax/test_transform_remove_unused_parameters.py b/tests/python/relax/test_transform_remove_unused_parameters.py index 989866a3f39d..3334354264b9 100644 --- a/tests/python/relax/test_transform_remove_unused_parameters.py +++ b/tests/python/relax/test_transform_remove_unused_parameters.py @@ -68,34 +68,40 @@ def test_replace_symbolic_variables(): parameter, which actually *defines* the variable. """ + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_func = T.dynamic("m") + n_func = T.dynamic("n") + @I.ir_module class Before: @R.function - def main(A: R.Tensor(["m", "n"], "float32")) -> R.Tensor(["m", "n"], "float32"): + def main(A: R.Tensor([m_main, n_main], "float32")) -> R.Tensor([m_main, n_main], "float32"): return Before.func(A) @R.function(private=True) - def func(A: R.Tensor(["m", "n"], "float32")) -> R.Tensor(["m", "n"], "float32"): - m = T.int64() - n = T.int64() - return R.zeros(R.shape([m, n]), dtype="float32") + def func(A: R.Tensor([m_func, n_func], "float32")) -> R.Tensor([m_func, n_func], "float32"): + return R.zeros(R.shape([m_func, n_func]), dtype="float32") + + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_func = T.dynamic("m") + n_func = T.dynamic("n") @I.ir_module class Expected: @R.function - def main(A: R.Tensor(["m", "n"], "float32")) -> R.Tensor(["m", "n"], "float32"): - m = T.int64() - n = T.int64() - out: R.Tensor([m, n], "float32") = Expected.func(R.shape([n]), R.shape([m])) + def main(A: R.Tensor([m_main, n_main], "float32")) -> R.Tensor([m_main, n_main], "float32"): + out: R.Tensor([m_main, n_main], "float32") = Expected.func( + R.shape([n_main]), R.shape([m_main]) + ) return out @R.function(private=True) - def func(param_n: R.Shape(["n"]), param_m: R.Shape(["m"])) -> R.Tensor( - ["m", "n"], "float32" + def func(param_n: R.Shape([n_func]), param_m: R.Shape([m_func])) -> R.Tensor( + [m_func, n_func], "float32" ): - m = T.int64() - n = T.int64() - return R.zeros(R.shape([m, n]), dtype="float32") + return R.zeros(R.shape([m_func, n_func]), dtype="float32") After = tvm.relax.transform.RemoveUnusedParameters()(Before) tvm.ir.assert_structural_equal(After, Expected) @@ -109,17 +115,20 @@ def test_no_extra_symbolic_variables(): distinct parameter. """ + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_func = T.dynamic("m") + n_func = T.dynamic("n") + @I.ir_module class Before: @R.function - def main(A: R.Tensor(["m", "n"], "float32")) -> R.Tensor(["m", "n"], "float32"): + def main(A: R.Tensor([m_main, n_main], "float32")) -> R.Tensor([m_main, n_main], "float32"): return Before.func(A) @R.function(private=True) - def func(A: R.Tensor(["m", "n"], "float32")) -> R.Tensor(["m", "n"], "float32"): - m = T.int64() - n = T.int64() - zeros = R.zeros(R.shape([m, n]), dtype="float32") + def func(A: R.Tensor([m_func, n_func], "float32")) -> R.Tensor([m_func, n_func], "float32"): + zeros = R.zeros(R.shape([m_func, n_func]), dtype="float32") out = R.add(A, zeros) return out @@ -136,37 +145,41 @@ def test_remove_extra_prim_parameters(): dtype-only scalar parameters are unused by the private function. """ + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_func = T.dynamic("m") + n_func = T.dynamic("n") + @I.ir_module class Before: @R.function - def main(A: R.Tensor(["m", "n"], "float32")) -> R.Tensor(["m", "n"], "float32"): - m = T.int64() - n = T.int64() - return Before.func(A, R.prim_value(m), R.prim_value(n)) + def main(A: R.Tensor([m_main, n_main], "float32")) -> R.Tensor([m_main, n_main], "float32"): + return Before.func(A, R.prim_value(m_main), R.prim_value(n_main)) @R.function(private=True) def func( - A: R.Tensor(["m", "n"], "float32"), + A: R.Tensor([m_func, n_func], "float32"), _m: T.int64, _n: T.int64, - ) -> R.Tensor(["m", "n"], "float32"): - m = T.int64() - n = T.int64() - zeros = R.zeros(R.shape([m, n]), dtype="float32") + ) -> R.Tensor([m_func, n_func], "float32"): + zeros = R.zeros(R.shape([m_func, n_func]), dtype="float32") out = R.add(A, zeros) return out + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_func = T.dynamic("m") + n_func = T.dynamic("n") + @I.ir_module class Expected: @R.function - def main(A: R.Tensor(["m", "n"], "float32")) -> R.Tensor(["m", "n"], "float32"): + def main(A: R.Tensor([m_main, n_main], "float32")) -> R.Tensor([m_main, n_main], "float32"): return Expected.func(A) @R.function(private=True) - def func(A: R.Tensor(["m", "n"], "float32")) -> R.Tensor(["m", "n"], "float32"): - m = T.int64() - n = T.int64() - zeros = R.zeros(R.shape([m, n]), dtype="float32") + def func(A: R.Tensor([m_func, n_func], "float32")) -> R.Tensor([m_func, n_func], "float32"): + zeros = R.zeros(R.shape([m_func, n_func]), dtype="float32") out = R.add(A, zeros) return out @@ -182,36 +195,40 @@ def test_remove_extra_shape_variables(): different parameter, then the `R.Shape` parameter can be removed. """ + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_func = T.dynamic("m") + n_func = T.dynamic("n") + @I.ir_module class Before: @R.function - def main(A: R.Tensor(["m", "n"], "float32")) -> R.Tensor(["m", "n"], "float32"): - m = T.int64() - n = T.int64() - return Before.func(A, R.shape([m, n])) + def main(A: R.Tensor([m_main, n_main], "float32")) -> R.Tensor([m_main, n_main], "float32"): + return Before.func(A, R.shape([m_main, n_main])) @R.function(private=True) def func( - A: R.Tensor(["m", "n"], "float32"), - _: R.Shape(["m", "n"]), - ) -> R.Tensor(["m", "n"], "float32"): - m = T.int64() - n = T.int64() - zeros = R.zeros(R.shape([m, n]), dtype="float32") + A: R.Tensor([m_func, n_func], "float32"), + _: R.Shape([m_func, n_func]), + ) -> R.Tensor([m_func, n_func], "float32"): + zeros = R.zeros(R.shape([m_func, n_func]), dtype="float32") out = R.add(A, zeros) return out + m_main = T.dynamic("m") + n_main = T.dynamic("n") + m_func = T.dynamic("m") + n_func = T.dynamic("n") + @I.ir_module class Expected: @R.function - def main(A: R.Tensor(["m", "n"], "float32")) -> R.Tensor(["m", "n"], "float32"): + def main(A: R.Tensor([m_main, n_main], "float32")) -> R.Tensor([m_main, n_main], "float32"): return Expected.func(A) @R.function(private=True) - def func(A: R.Tensor(["m", "n"], "float32")) -> R.Tensor(["m", "n"], "float32"): - m = T.int64() - n = T.int64() - zeros = R.zeros(R.shape([m, n]), dtype="float32") + def func(A: R.Tensor([m_func, n_func], "float32")) -> R.Tensor([m_func, n_func], "float32"): + zeros = R.zeros(R.shape([m_func, n_func]), dtype="float32") out = R.add(A, zeros) return out diff --git a/tests/python/relax/test_transform_reorder_take_after_matmul.py b/tests/python/relax/test_transform_reorder_take_after_matmul.py index c07a86911670..cc9c529dea1f 100644 --- a/tests/python/relax/test_transform_reorder_take_after_matmul.py +++ b/tests/python/relax/test_transform_reorder_take_after_matmul.py @@ -14,7 +14,6 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -# ruff: noqa: F841 import inspect @@ -39,16 +38,19 @@ def test_compare(self): tvm.ir.assert_structural_equal(self.Expected, after) +weight_table_size_before = T.dynamic("weight_table_size") +weight_table_size_expected = T.dynamic("weight_table_size") + + class TestSimple(Base): @I.ir_module class Before: @R.function def main( x: R.Tensor([1, 16], "float32"), - weight_table: R.Tensor([16, "weight_table_size"], "float32"), + weight_table: R.Tensor([16, weight_table_size_before], "float32"), routing_table: R.Tensor([32], "int64"), ) -> R.Tensor([1, 32], "float32"): - weight_table_size = T.int64() with R.dataflow(): weight: R.Tensor([16, 32], "float32") = R.take(weight_table, routing_table, axis=1) out: R.Tensor([1, 32], "float32") = R.matmul(x, weight) @@ -60,31 +62,36 @@ class Expected: @R.function def main( x: R.Tensor([1, 16], "float32"), - weight_table: R.Tensor([16, "weight_table_size"], "float32"), + weight_table: R.Tensor([16, weight_table_size_expected], "float32"), routing_table: R.Tensor([32], "int64"), ) -> R.Tensor([1, 32], "float32"): - weight_table_size = T.int64() with R.dataflow(): - out_table: R.Tensor([1, weight_table_size], "float32") = R.matmul(x, weight_table) + out_table: R.Tensor([1, weight_table_size_expected], "float32") = R.matmul( + x, weight_table + ) out: R.Tensor([1, 32], "float32") = R.take(out_table, routing_table, axis=1) R.output(out) return out +batch_size_before = T.dynamic("batch_size") +weight_table_size_before = T.dynamic("weight_table_size") +batch_size_expected = T.dynamic("batch_size") +weight_table_size_expected = T.dynamic("weight_table_size") + + class TestBatchedActivations(Base): @I.ir_module class Before: @R.function def main( - x: R.Tensor(["batch_size", 1, 16], "float32"), - weight_table: R.Tensor([16, "weight_table_size"], "float32"), + x: R.Tensor([batch_size_before, 1, 16], "float32"), + weight_table: R.Tensor([16, weight_table_size_before], "float32"), routing_table: R.Tensor([32], "int64"), - ) -> R.Tensor(["batch_size", 1, 32], "float32"): - batch_size = T.int64() - weight_table_size = T.int64() + ) -> R.Tensor([batch_size_before, 1, 32], "float32"): with R.dataflow(): weight: R.Tensor([16, 32], "float32") = R.take(weight_table, routing_table, axis=1) - out: R.Tensor([batch_size, 1, 32], "float32") = R.matmul(x, weight) + out: R.Tensor([batch_size_before, 1, 32], "float32") = R.matmul(x, weight) R.output(out) return out @@ -92,34 +99,36 @@ def main( class Expected: @R.function def main( - x: R.Tensor(["batch_size", 1, 16], "float32"), - weight_table: R.Tensor([16, "weight_table_size"], "float32"), + x: R.Tensor([batch_size_expected, 1, 16], "float32"), + weight_table: R.Tensor([16, weight_table_size_expected], "float32"), routing_table: R.Tensor([32], "int64"), - ) -> R.Tensor(["batch_size", 1, 32], "float32"): - batch_size = T.int64() - weight_table_size = T.int64() + ) -> R.Tensor([batch_size_expected, 1, 32], "float32"): with R.dataflow(): - out_table: R.Tensor([batch_size, 1, weight_table_size], "float32") = R.matmul( - x, weight_table - ) - out: R.Tensor([batch_size, 1, 32], "float32") = R.take( + out_table: R.Tensor( + [batch_size_expected, 1, weight_table_size_expected], "float32" + ) = R.matmul(x, weight_table) + out: R.Tensor([batch_size_expected, 1, 32], "float32") = R.take( out_table, routing_table, axis=2 ) R.output(out) return out +batch_size_before = T.dynamic("batch_size") +routing_table_size_before = T.dynamic("routing_table_size") +batch_size_expected = T.dynamic("batch_size") +routing_table_size_expected = T.dynamic("routing_table_size") + + class TestStaticBatchedActivationsAndWeights(Base): @I.ir_module class Before: @R.function def main( x: R.Tensor([128, 1, 16], "float32"), - weight_table: R.Tensor(["routing_table_size", 16, 32], "float32"), + weight_table: R.Tensor([routing_table_size_before, 16, 32], "float32"), routing_table: R.Tensor([128], "int64"), ) -> R.Tensor([128, 1, 32], "float32"): - batch_size = T.int64() - routing_table_size = T.int64() with R.dataflow(): weight = R.take(weight_table, routing_table, axis=0) out = R.matmul(x, weight) @@ -131,33 +140,37 @@ class Expected: @R.function def main( x: R.Tensor([128, 1, 16], "float32"), - weight_table: R.Tensor(["routing_table_size", 16, 32], "float32"), + weight_table: R.Tensor([routing_table_size_expected, 16, 32], "float32"), routing_table: R.Tensor([128], "int64"), ) -> R.Tensor([128, 1, 32], "float32"): - batch_size = T.int64() - routing_table_size = T.int64() with R.dataflow(): reordered_weight = R.permute_dims(weight_table, [1, 0, 2]) - fused_weight = R.reshape(reordered_weight, [16, routing_table_size * 32]) + fused_weight = R.reshape(reordered_weight, [16, routing_table_size_expected * 32]) fused_output = R.matmul(x, fused_weight) - reordered_output = R.reshape(fused_output, [128, 1, routing_table_size, 32]) + reordered_output = R.reshape( + fused_output, [128, 1, routing_table_size_expected, 32] + ) tabular_output = R.take(reordered_output, routing_table, axis=2) out = R.einsum([tabular_output], "ijik->ijk") R.output(out) return out +batch_size_before = T.dynamic("batch_size") +routing_table_size_before = T.dynamic("routing_table_size") +batch_size_expected = T.dynamic("batch_size") +routing_table_size_expected = T.dynamic("routing_table_size") + + class TestDynamicBatchedActivationsAndWeights(Base): @I.ir_module class Before: @R.function def main( - x: R.Tensor(["batch_size", 1, 16], "float32"), - weight_table: R.Tensor(["routing_table_size", 16, 32], "float32"), - routing_table: R.Tensor(["batch_size"], "int64"), - ) -> R.Tensor(["batch_size", 1, 32], "float32"): - batch_size = T.int64() - routing_table_size = T.int64() + x: R.Tensor([batch_size_before, 1, 16], "float32"), + weight_table: R.Tensor([routing_table_size_before, 16, 32], "float32"), + routing_table: R.Tensor([batch_size_before], "int64"), + ) -> R.Tensor([batch_size_before, 1, 32], "float32"): with R.dataflow(): weight = R.take(weight_table, routing_table, axis=0) out = R.matmul(x, weight) @@ -168,17 +181,17 @@ def main( class Expected: @R.function def main( - x: R.Tensor(["batch_size", 1, 16], "float32"), - weight_table: R.Tensor(["routing_table_size", 16, 32], "float32"), - routing_table: R.Tensor(["batch_size"], "int64"), - ) -> R.Tensor(["batch_size", 1, 32], "float32"): - batch_size = T.int64() - routing_table_size = T.int64() + x: R.Tensor([batch_size_expected, 1, 16], "float32"), + weight_table: R.Tensor([routing_table_size_expected, 16, 32], "float32"), + routing_table: R.Tensor([batch_size_expected], "int64"), + ) -> R.Tensor([batch_size_expected, 1, 32], "float32"): with R.dataflow(): reordered_weight = R.permute_dims(weight_table, [1, 0, 2]) - fused_weight = R.reshape(reordered_weight, [16, routing_table_size * 32]) + fused_weight = R.reshape(reordered_weight, [16, routing_table_size_expected * 32]) fused_output = R.matmul(x, fused_weight) - reordered_output = R.reshape(fused_output, [batch_size, 1, routing_table_size, 32]) + reordered_output = R.reshape( + fused_output, [batch_size_expected, 1, routing_table_size_expected, 32] + ) tabular_output = R.take(reordered_output, routing_table, axis=2) out = R.einsum([tabular_output], "ijik->ijk") R.output(out) diff --git a/tests/python/relax/test_transform_rewrite_cuda_graph.py b/tests/python/relax/test_transform_rewrite_cuda_graph.py index 078bb157a20d..fdcffb2fd86d 100644 --- a/tests/python/relax/test_transform_rewrite_cuda_graph.py +++ b/tests/python/relax/test_transform_rewrite_cuda_graph.py @@ -763,55 +763,59 @@ def main() -> R.Tuple: def test_dynamic_capture(): + m_add_one = T.dynamic("m") + m_main = T.dynamic("m") + @I.ir_module class Before: @Ts.prim_func def add_one(x_handle: T.handle, y_handle: T.handle): - m = T.int64() - x = T.match_buffer(x_handle, (m,), "float32") - y = T.match_buffer(y_handle, (m,), "float32") + x = T.match_buffer(x_handle, (m_add_one,), "float32") + y = T.match_buffer(y_handle, (m_add_one,), "float32") # Use T.serial with explicit int64 min so the inner sblock iter_var # dom is all-int64 (matches what Expected emits via Ts.axis.spatial(m, i)). - for i in T.serial(T.int64(0), m): + for i in T.serial(T.int64(0), m_add_one): with Ts.sblock("add"): vi = Ts.axis.remap("S", [i]) y[vi] = x[vi] + T.float32(1) @R.function - def main(x: R.Tensor(("m",), "float32")) -> R.Tensor(("m",), "float32"): + def main(x: R.Tensor((m_main,), "float32")) -> R.Tensor((m_main,), "float32"): R.func_attr( {"relax.rewrite_cuda_graph.capture_symbolic_vars": ["m"], "relax.force_pure": True} ) - m = T.int64() storage: R.Any = R.memory.alloc_storage( R.shape([16]), 0, "global", "float32" ) # assume m is upper-bounded - alloc1: R.Tensor((m,), "float32") = R.memory.alloc_tensor( - storage, 0, R.shape([m]), "float32" + alloc1: R.Tensor((m_main,), "float32") = R.memory.alloc_tensor( + storage, 0, R.shape([m_main]), "float32" ) _ = Before.add_one(x, alloc1) storage1: R.Any = R.memory.alloc_storage(R.shape([16]), 0, "global", "float32") - alloc2: R.Tensor((m,), "float32") = R.memory.alloc_tensor( - storage1, 0, R.shape([m]), "float32" + alloc2: R.Tensor((m_main,), "float32") = R.memory.alloc_tensor( + storage1, 0, R.shape([m_main]), "float32" ) _ = Before.add_one(alloc1, alloc2) - alloc3: R.Tensor((m,), "float32") = R.builtin.alloc_tensor( - R.shape([m]), "float32", 0, "global" + alloc3: R.Tensor((m_main,), "float32") = R.builtin.alloc_tensor( + R.shape([m_main]), "float32", 0, "global" ) _ = Before.add_one(alloc2, alloc3) return alloc3 + m_add_one = T.dynamic("m") + m_main_cuda_graph_capture = T.dynamic("m") + m_main = T.dynamic("m") + @I.ir_module class Expected: @Ts.prim_func def add_one(x_handle: T.handle, y_handle: T.handle): - m = T.int64() - x = T.match_buffer(x_handle, (m,)) - y = T.match_buffer(y_handle, (m,)) + x = T.match_buffer(x_handle, (m_add_one,)) + y = T.match_buffer(y_handle, (m_add_one,)) # with Ts.sblock("root"): - for i in T.serial(T.int64(0), m): + for i in T.serial(T.int64(0), m_add_one): with Ts.sblock("add"): - vi = Ts.axis.spatial(m, i) + vi = Ts.axis.spatial(m_add_one, i) Ts.reads(x[vi]) Ts.writes(y[vi]) y[vi] = x[vi] + T.float32(1) @@ -830,11 +834,10 @@ def cuda_graph_alloc() -> R.Tuple(R.Any, R.Any): @R.function(private=True) def main_cuda_graph_capture( - alloc1: R.Tensor(("m",), dtype="float32"), - alloc2: R.Tensor(("m",), dtype="float32"), - shape_expr: R.Shape(["m"]), + alloc1: R.Tensor((m_main_cuda_graph_capture,), dtype="float32"), + alloc2: R.Tensor((m_main_cuda_graph_capture,), dtype="float32"), + shape_expr: R.Shape([m_main_cuda_graph_capture]), ): - m = T.int64() R.func_attr({"relax.force_pure": True}) cls = Expected cls.add_one(alloc1, alloc2) @@ -842,8 +845,7 @@ def main_cuda_graph_capture( return R.tuple() @R.function - def main(x: R.Tensor(("m",), dtype="float32")) -> R.Tensor(("m",), dtype="float32"): - m = T.int64() + def main(x: R.Tensor((m_main,), dtype="float32")) -> R.Tensor((m_main,), dtype="float32"): R.func_attr( {"relax.force_pure": True, "relax.rewrite_cuda_graph.capture_symbolic_vars": ["m"]} ) @@ -854,26 +856,26 @@ def main(x: R.Tensor(("m",), dtype="float32")) -> R.Tensor(("m",), dtype="float3 ty_args=(R.Tuple(R.Any, R.Any),), ) storage: R.Any = gv[0] - alloc1: R.Tensor((m,), dtype="float32") = R.memory.alloc_tensor( - storage, R.prim_value(0), R.shape([m]), R.dtype("float32") + alloc1: R.Tensor((m_main,), dtype="float32") = R.memory.alloc_tensor( + storage, R.prim_value(0), R.shape([m_main]), R.dtype("float32") ) cls.add_one(x, alloc1) storage1: R.Any = gv[1] - alloc2: R.Tensor((m,), dtype="float32") = R.memory.alloc_tensor( - storage1, R.prim_value(0), R.shape([m]), R.dtype("float32") + alloc2: R.Tensor((m_main,), dtype="float32") = R.memory.alloc_tensor( + storage1, R.prim_value(0), R.shape([m_main]), R.dtype("float32") ) R.call_builtin_with_ctx( "vm.builtin.cuda_graph.run_or_capture", ( cls.main_cuda_graph_capture, - (alloc1, alloc2, R.shape([m])), + (alloc1, alloc2, R.shape([m_main])), R.prim_value(0), - R.shape([m]), + R.shape([m_main]), ), ty_args=(R.Tuple,), ) - alloc3: R.Tensor((m,), dtype="float32") = R.builtin.alloc_tensor( - R.shape([m]), R.dtype("float32"), R.prim_value(0), R.str("global") + alloc3: R.Tensor((m_main,), dtype="float32") = R.builtin.alloc_tensor( + R.shape([m_main]), R.dtype("float32"), R.prim_value(0), R.str("global") ) cls.add_one(alloc2, alloc3) return alloc3 @@ -1102,11 +1104,12 @@ def main(x: R.Tensor((8,), dtype="float32")) -> R.Tuple(R.Tensor((8,), dtype="fl def test_static_input_with_symbolic_shape(): + m = T.dynamic("m") + @I.ir_module class Before: @R.function - def main(x: R.Tensor((8,), "float16"), w: R.Tensor(("m",))): - m = T.int64() + def main(x: R.Tensor((8,), "float16"), w: R.Tensor((m,))): R.func_attr({"relax.force_pure": True, "num_input": 1}) storage1 = R.memory.alloc_storage(R.shape([8]), 0, "global", "float16") alloc1 = R.memory.alloc_tensor(storage1, 0, R.shape([8]), "float16") @@ -1120,6 +1123,9 @@ def main(x: R.Tensor((8,), "float16"), w: R.Tensor(("m",))): gv = (alloc3,) return gv + m_main_cuda_graph_capture = T.dynamic("m") + m_main = T.dynamic("m") + @I.ir_module class Expected: @R.function(private=True) @@ -1137,21 +1143,19 @@ def cuda_graph_alloc() -> R.Tuple(R.Any, R.Any): @R.function(private=True) def main_cuda_graph_capture( alloc1: R.Tensor((8,), dtype="float16"), - w: R.Tensor(("m",)), + w: R.Tensor((m_main_cuda_graph_capture,)), alloc2: R.Tensor((8,), dtype="float16"), - shape_expr: R.Shape(["m"]), + shape_expr: R.Shape([m_main_cuda_graph_capture]), ) -> R.Tuple: - m = T.int64() R.func_attr({"relax.force_pure": True}) R.call_packed("dummy", alloc1, w, alloc2, ty_args=(R.Tuple,)) R.tuple() return R.tuple() @R.function - def main(x: R.Tensor((8,), dtype="float16"), w: R.Tensor(("m",))) -> R.Tuple( + def main(x: R.Tensor((8,), dtype="float16"), w: R.Tensor((m_main,))) -> R.Tuple( R.Tensor((8,), dtype="float16") ): - m = T.int64() R.func_attr({"num_input": 1, "relax.force_pure": True}) cls = Expected gv: R.Tuple(R.Any, R.Any) = R.call_builtin_with_ctx( @@ -1172,9 +1176,9 @@ def main(x: R.Tensor((8,), dtype="float16"), w: R.Tensor(("m",))) -> R.Tuple( "vm.builtin.cuda_graph.run_or_capture", ( cls.main_cuda_graph_capture, - (alloc1, w, alloc2, R.shape([m])), + (alloc1, w, alloc2, R.shape([m_main])), R.prim_value(0), - R.shape([m]), + R.shape([m_main]), ), ty_args=(R.Tuple,), ) diff --git a/tests/python/relax/test_transform_rewrite_dataflow_reshape.py b/tests/python/relax/test_transform_rewrite_dataflow_reshape.py index 6070c0543ea3..1dcecd0733db 100644 --- a/tests/python/relax/test_transform_rewrite_dataflow_reshape.py +++ b/tests/python/relax/test_transform_rewrite_dataflow_reshape.py @@ -224,12 +224,13 @@ def main(x: R.Tensor((2, 4096, 320), dtype="float32")) -> R.Tensor((2, 1, 4096, def test_reshape_dynamic_shape(): + n = T.dynamic("n", "int32") + @tvm.script.ir_module class Module: @Ts.prim_func(private=True) def reshape(var_A: T.handle, var_T_reshape: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int32() A = T.match_buffer(var_A, (n, 16, 128), "float16") T_reshape = T.match_buffer(var_T_reshape, (1, n, 16, 128), "float16") # with Ts.sblock("root"): @@ -264,12 +265,13 @@ def main(x: R.Tensor((8, 16, 128), dtype="float16")) -> R.Tensor( R.output(z) return z + n = T.dynamic("n", "int32") + @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) def reshape(var_A: T.handle, var_T_reshape: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int32() A = T.match_buffer(var_A, (n, 16, 128), "float16") T_reshape = T.match_buffer(var_T_reshape, (1, n, 16, 128), "float16") # with Ts.sblock("root"): @@ -638,11 +640,12 @@ def add( # def test_rewrite_dynamic_reshape(): +# N = T.dynamic("N") +# # @I.ir_module # class Before: # @R.function -# def main(x: R.Tensor(["N"], dtype="float32")): -# N = T.int64() +# def main(x: R.Tensor([N], dtype="float32")): # with R.dataflow(): # y = R.reshape(x, [N // 4, 4]) # z = R.add(y, y) @@ -652,8 +655,7 @@ def add( # @I.ir_module # class Expected: # @R.function -# def main(x: R.Tensor(["N"], dtype="float32")): -# N = T.int64() +# def main(x: R.Tensor([N], dtype="float32")): # cls = Expected # with R.dataflow(): @@ -702,22 +704,24 @@ def add( def test_rewrite_dynamic_reshape(): + N = T.dynamic("N") + @I.ir_module class Before: @R.function - def main(x: R.Tensor(["N", 16], dtype="float32")): - N = T.int64() + def main(x: R.Tensor([N, 16], dtype="float32")): with R.dataflow(): y = R.reshape(x, [N * 4, T.int64(4)]) z = R.add(y, y) R.output(z) return z + N = T.dynamic("N") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(["N", 16], dtype="float32")): - N = T.int64() + def main(x: R.Tensor([N, 16], dtype="float32")): cls = Expected with R.dataflow(): diff --git a/tests/python/relax/test_transform_static_plan_block_memory.py b/tests/python/relax/test_transform_static_plan_block_memory.py index 146657901962..abc9ada82ba3 100644 --- a/tests/python/relax/test_transform_static_plan_block_memory.py +++ b/tests/python/relax/test_transform_static_plan_block_memory.py @@ -681,52 +681,59 @@ def main(x: R.Tensor((2, 3), dtype="float32")) -> R.Tensor((2, 3), dtype="float3 def test_symbolic_shape(): + m_exp = T.dynamic("m") + n_exp = T.dynamic("n") + m_main = T.dynamic("m") + n_main = T.dynamic("n") + @tvm.script.ir_module class Module: @Ts.prim_func def exp(var_A: T.handle, var_B: T.handle): - m = T.int64() - n = T.int64() - A = T.match_buffer(var_A, (m, n), "float32") - B = T.match_buffer(var_B, (m, n), "float32") + A = T.match_buffer(var_A, (m_exp, n_exp), "float32") + B = T.match_buffer(var_B, (m_exp, n_exp), "float32") T.evaluate(0) @R.function - def main(x: R.Tensor(("m", "n"), "float32")): + def main(x: R.Tensor((m_main, n_main), "float32")): R.func_attr({"relax.force_pure": True}) - m = T.int64() - n = T.int64() - alloc: R.Tensor((m, n), dtype="float32") = R.builtin.alloc_tensor( - R.shape([m, n]), dtype="float32", runtime_device_index=0 + alloc: R.Tensor((m_main, n_main), dtype="float32") = R.builtin.alloc_tensor( + R.shape([m_main, n_main]), dtype="float32", runtime_device_index=0 ) _ = Module.exp(x, alloc) - y: R.Tensor((m, n), dtype="float32") = alloc + y: R.Tensor((m_main, n_main), dtype="float32") = alloc return x + m_exp = T.dynamic("m") + n_exp = T.dynamic("n") + m_main = T.dynamic("m") + n_main = T.dynamic("n") + @tvm.script.ir_module class Expected: @Ts.prim_func def exp(var_A: T.handle, var_B: T.handle): - m = T.int64() - n = T.int64() - A = T.match_buffer(var_A, (m, n), "float32") - B = T.match_buffer(var_B, (m, n), "float32") + A = T.match_buffer(var_A, (m_exp, n_exp), "float32") + B = T.match_buffer(var_B, (m_exp, n_exp), "float32") T.evaluate(0) @R.function - def main(x: R.Tensor(("m", "n"), dtype="float32")) -> R.Tensor(("m", "n"), dtype="float32"): - m = T.int64() - n = T.int64() + def main(x: R.Tensor((m_main, n_main), dtype="float32")) -> R.Tensor( + (m_main, n_main), dtype="float32" + ): R.func_attr({"relax.force_pure": True}) cls = Expected storage: R.Any = R.memory.alloc_storage( - R.shape([4 * (m * n)]), R.prim_value(0), R.str("global"), R.dtype("float32") + R.shape([4 * (m_main * n_main)]), + R.prim_value(0), + R.str("global"), + R.dtype("float32"), ) - alloc: R.Tensor((m, n), dtype="float32") = R.memory.alloc_tensor( - storage, R.prim_value(0), R.shape([m, n]), R.dtype("float32") + alloc: R.Tensor((m_main, n_main), dtype="float32") = R.memory.alloc_tensor( + storage, R.prim_value(0), R.shape([m_main, n_main]), R.dtype("float32") ) _: R.Tuple = cls.exp(x, alloc) - y: R.Tensor((m, n), dtype="float32") = alloc + y: R.Tensor((m_main, n_main), dtype="float32") = alloc return x mod = relax.transform.StaticPlanBlockMemory()(Module) @@ -915,6 +922,8 @@ def func2( def test_tir_var_upper_bound(): # fmt: off + n = T.dynamic("n") + @tvm.script.ir_module class Module: @Ts.prim_func @@ -942,9 +951,8 @@ def pad(rxplaceholder: T.handle, PadInput: T.handle): T.evaluate(0) @R.function - def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dtype="float32"): + def main(x: R.Tensor((2, n), dtype="float32")) -> R.Tensor((2 * n + 2,), dtype="float32"): R.func_attr({"tir_var_upper_bound": {"n": 4}, "relax.force_pure": True}) - n = T.int64() cls = Module alloc: R.Tensor((2, n), dtype="float32") = R.builtin.alloc_tensor(R.shape([2, n]), dtype="float32", runtime_device_index=0) _: R.Tuple() = cls.exp(x, alloc) @@ -964,6 +972,8 @@ def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dty gv: R.Tensor((2 * n + 2,), dtype="float32") = alloc4 return gv + n = T.dynamic("n") + @I.ir_module class Expected: @Ts.prim_func @@ -991,8 +1001,7 @@ def reshape(rxplaceholder: T.handle, T_reshape: T.handle): T.evaluate(0) @R.function - def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dtype="float32"): - n = T.int64() + def main(x: R.Tensor((2, n), dtype="float32")) -> R.Tensor((2 * n + 2,), dtype="float32"): R.func_attr({"tir_var_upper_bound": {"n": 4}, "relax.force_pure": True}) cls = Expected storage: R.Any = R.memory.alloc_storage(R.shape([32]), R.prim_value(0), R.str("global"), R.dtype("float32")) @@ -1022,6 +1031,8 @@ def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dty def test_lower_bound_only(): # fmt: off + n = T.dynamic("n") + @tvm.script.ir_module class Module: @Ts.prim_func @@ -1049,9 +1060,8 @@ def pad(rxplaceholder: T.handle, PadInput: T.handle): T.evaluate(0) @R.function - def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dtype="float32"): + def main(x: R.Tensor((2, n), dtype="float32")) -> R.Tensor((2 * n + 2,), dtype="float32"): R.func_attr({"tir_var_lower_bound": {"n": 2}, "relax.force_pure": True}) - n = T.int64() cls = Module alloc: R.Tensor((2, n), dtype="float32") = R.builtin.alloc_tensor(R.shape([2, n]), dtype="float32", runtime_device_index=0) _: R.Tuple() = cls.exp(x, alloc) @@ -1071,6 +1081,8 @@ def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dty gv: R.Tensor((2 * n + 2,), dtype="float32") = alloc4 return gv + n = T.dynamic("n") + @I.ir_module class Expected: @Ts.prim_func @@ -1098,8 +1110,7 @@ def reshape(rxplaceholder: T.handle, T_reshape: T.handle): T.evaluate(0) @R.function - def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dtype="float32"): - n = T.int64() + def main(x: R.Tensor((2, n), dtype="float32")) -> R.Tensor((2 * n + 2,), dtype="float32"): R.func_attr({"tir_var_lower_bound": {"n": 2}, "relax.force_pure": True}) cls = Expected storage: R.Any = R.memory.alloc_storage(R.shape([8 * n]), R.prim_value(0), R.str("global"), R.dtype("float32")) @@ -1130,6 +1141,8 @@ def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dty def test_upper_and_lower_bounds(): # fmt: off + n = T.dynamic("n") + @tvm.script.ir_module class Module: @Ts.prim_func @@ -1157,9 +1170,8 @@ def pad(rxplaceholder: T.handle, PadInput: T.handle): T.evaluate(0) @R.function - def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dtype="float32"): + def main(x: R.Tensor((2, n), dtype="float32")) -> R.Tensor((2 * n + 2,), dtype="float32"): R.func_attr({"tir_var_upper_bound": {"n": 4}, "tir_var_lower_bound": {"n": 2}, "relax.force_pure": True}) - n = T.int64() cls = Module alloc: R.Tensor((2, n), dtype="float32") = R.builtin.alloc_tensor(R.shape([2, n]), dtype="float32", runtime_device_index=0) _: R.Tuple() = cls.exp(x, alloc) @@ -1179,6 +1191,8 @@ def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dty gv: R.Tensor((2 * n + 2,), dtype="float32") = alloc4 return gv + n = T.dynamic("n") + @I.ir_module class Expected: @Ts.prim_func @@ -1206,8 +1220,7 @@ def reshape(rxplaceholder: T.handle, T_reshape: T.handle): T.evaluate(0) @R.function - def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dtype="float32"): - n = T.int64() + def main(x: R.Tensor((2, n), dtype="float32")) -> R.Tensor((2 * n + 2,), dtype="float32"): R.func_attr({"tir_var_upper_bound": {"n": 4}, "tir_var_lower_bound": {"n": 2}, "relax.force_pure": True}) cls = Expected storage: R.Any = R.memory.alloc_storage(R.shape([32]), R.prim_value(0), R.str("global"), R.dtype("float32")) @@ -1236,10 +1249,12 @@ def main(x: R.Tensor((2, "n"), dtype="float32")) -> R.Tensor(("2 * n + 2",), dty def test_invalid_tir_var_upper_bound(): + n = T.dynamic("n") + @tvm.script.ir_module class Module: @R.function - def main(x: R.Tensor((2, "n"), dtype="float32")): + def main(x: R.Tensor((2, n), dtype="float32")): R.func_attr({"tir_var_upper_bound": {"n": [4]}, "relax.force_pure": True}) return x @@ -1248,10 +1263,12 @@ def main(x: R.Tensor((2, "n"), dtype="float32")): def test_invalid_tir_var_lower_bound(): + n = T.dynamic("n") + @tvm.script.ir_module class Module: @R.function - def main(x: R.Tensor((2, "n"), dtype="float32")): + def main(x: R.Tensor((2, n), dtype="float32")): R.func_attr({"tir_var_lower_bound": {"n": [4]}, "relax.force_pure": True}) return x @@ -1261,6 +1278,9 @@ def main(x: R.Tensor((2, "n"), dtype="float32")): def test_tir_var_decreasing_monotone(): # fmt: off + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Module: @Ts.prim_func @@ -1268,9 +1288,7 @@ def tir_exp(var_rxplaceholder: T.handle, var_compute: T.handle): T.evaluate(0) @R.function - def main(x: R.Tensor(("n", "m", "T.max(n - m, 1)"), dtype="float32")) -> R.Tensor(("n", "m", "T.max(n - m, 1)"), dtype="float32"): - n = T.int64() - m = T.int64() + def main(x: R.Tensor((n, m, T.max(n - m, 1)), dtype="float32")) -> R.Tensor((n, m, T.max(n - m, 1)), dtype="float32"): R.func_attr({"tir_var_upper_bound": {"m": 5, "n": 20}, "relax.force_pure": True}) cls = Module alloc: R.Tensor((n, m, T.max(n - m, 1)), dtype="float32") = R.builtin.alloc_tensor(R.shape([n, m, T.max(n - m, 1)]), R.dtype("float32"), R.prim_value(0)) @@ -1284,6 +1302,9 @@ def main(x: R.Tensor(("n", "m", "T.max(n - m, 1)"), dtype="float32")) -> R.Tenso r: R.Tensor((n, m, T.max(n - m, 1)), dtype="float32") = alloc2 return r + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Expected: @Ts.prim_func @@ -1291,9 +1312,7 @@ def tir_exp(var_rxplaceholder: T.handle, var_compute: T.handle): T.evaluate(0) @R.function - def main(x: R.Tensor(("n", "m", "T.max(n - m, 1)"), dtype="float32")) -> R.Tensor(("n", "m", "T.max(n - m, 1)"), dtype="float32"): - n = T.int64() - m = T.int64() + def main(x: R.Tensor((n, m, T.max(n - m, 1)), dtype="float32")) -> R.Tensor((n, m, T.max(n - m, 1)), dtype="float32"): R.func_attr({"tir_var_upper_bound": {"m": 5, "n": 20}, "relax.force_pure": True}) cls = Expected storage: R.Any = R.memory.alloc_storage(R.shape([8000]), R.prim_value(0), R.str("global"), R.dtype("float32")) @@ -1316,6 +1335,8 @@ def main(x: R.Tensor(("n", "m", "T.max(n - m, 1)"), dtype="float32")) -> R.Tenso def test_call_tir_dyn(): # fmt: off + n = T.dynamic("n") + @I.ir_module class Module: @Ts.prim_func @@ -1327,8 +1348,7 @@ def tir_exp(var_rxplaceholder: T.handle, var_compute: T.handle): T.evaluate(0) @R.function - def main(s: R.Shape(["n"])) -> R.Tensor(("n",), dtype="float32"): - n = T.int64() + def main(s: R.Shape([n])) -> R.Tensor((n,), dtype="float32"): R.func_attr({"tir_var_upper_bound": {"n": 20}, "relax.force_pure": True}) cls = Module alloc: R.Tensor((n,), dtype="float32") = R.builtin.alloc_tensor(R.shape([n]), R.dtype("float32"), R.prim_value(0)) @@ -1342,6 +1362,8 @@ def main(s: R.Shape(["n"])) -> R.Tensor(("n",), dtype="float32"): lv3: R.Tensor((n,), dtype="float32") = alloc2 return lv3 + n = T.dynamic("n") + @I.ir_module class Expected: @Ts.prim_func @@ -1353,8 +1375,7 @@ def tir_full(var_full: T.handle, n: T.int64): T.evaluate(0) @R.function - def main(s: R.Shape(["n"])) -> R.Tensor(("n",), dtype="float32"): - n = T.int64() + def main(s: R.Shape([n])) -> R.Tensor((n,), dtype="float32"): R.func_attr({"tir_var_upper_bound": {"n": 20}, "relax.force_pure": True}) cls = Expected storage: R.Any = R.memory.alloc_storage(R.shape([80]), R.prim_value(0), R.str("global"), R.dtype("float32")) @@ -1377,6 +1398,8 @@ def main(s: R.Shape(["n"])) -> R.Tensor(("n",), dtype="float32"): def test_call_tir_dyn_plan_dynamic_func_output(): # fmt: off + n = T.dynamic("n") + @I.ir_module class Module: @Ts.prim_func @@ -1388,8 +1411,7 @@ def tir_exp(var_rxplaceholder: T.handle, var_compute: T.handle): T.evaluate(0) @R.function - def main(s: R.Shape(["n"])) -> R.Tensor(("n",), dtype="float32"): - n = T.int64() + def main(s: R.Shape([n])) -> R.Tensor((n,), dtype="float32"): R.func_attr({"tir_var_upper_bound": {"n": 20}, "relax.force_pure": True, "relax.memory_plan_dynamic_func_output": True}) cls = Module alloc: R.Tensor((n,), dtype="float32") = R.builtin.alloc_tensor(R.shape([n]), R.dtype("float32"), R.prim_value(0)) @@ -1403,6 +1425,8 @@ def main(s: R.Shape(["n"])) -> R.Tensor(("n",), dtype="float32"): lv3: R.Tensor((n,), dtype="float32") = alloc2 return lv3 + n = T.dynamic("n") + @I.ir_module class Expected: @Ts.prim_func @@ -1414,8 +1438,7 @@ def tir_full(var_full: T.handle, n: T.int64): T.evaluate(0) @R.function - def main(s: R.Shape(["n"])) -> R.Tensor(("n",), dtype="float32"): - n = T.int64() + def main(s: R.Shape([n])) -> R.Tensor((n,), dtype="float32"): R.func_attr({"tir_var_upper_bound": {"n": 20}, "relax.force_pure": True}) cls = Expected storage: R.Any = R.memory.alloc_storage(R.shape([80]), R.prim_value(0), R.str("global"), R.dtype("float32")) @@ -1439,6 +1462,9 @@ def main(s: R.Shape(["n"])) -> R.Tensor(("n",), dtype="float32"): def test_call_tir_dyn_plan_partially_dynamic(): # fmt: off + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Module: @Ts.prim_func @@ -1450,9 +1476,7 @@ def tir_exp(var_rxplaceholder: T.handle, var_compute: T.handle): T.evaluate(0) @R.function - def main(s: R.Shape(["n", "m"])) -> R.Tensor(("n", "m"), dtype="float32"): - n = T.int64() - m = T.int64() + def main(s: R.Shape([n, m])) -> R.Tensor((n, m), dtype="float32"): R.func_attr({"tir_var_upper_bound": {"n": 20}, "relax.force_pure": True, "relax.memory_plan_dynamic_func_output": True}) cls = Module alloc: R.Tensor((n, m), dtype="float32") = R.builtin.alloc_tensor(R.shape([n, m]), R.dtype("float32"), R.prim_value(0)) @@ -1469,6 +1493,9 @@ def main(s: R.Shape(["n", "m"])) -> R.Tensor(("n", "m"), dtype="float32"): lv4: R.Tensor((n, m), dtype="float32") = alloc3 return lv4 + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Expected: @Ts.prim_func @@ -1480,9 +1507,7 @@ def tir_exp(var_rxplaceholder: T.handle, var_compute: T.handle): T.evaluate(0) @R.function - def main(s: R.Shape(["n", "m"])) -> R.Tensor(("n", "m"), dtype="float32"): - n = T.int64() - m = T.int64() + def main(s: R.Shape([n, m])) -> R.Tensor((n, m), dtype="float32"): R.func_attr({"relax.force_pure": True, "tir_var_upper_bound": {"n": 20}}) cls = Expected storage: R.Any = R.memory.alloc_storage(R.shape([80 * m]), R.prim_value(0), R.str("global"), R.dtype("float32")) @@ -1577,6 +1602,9 @@ def func2(x: R.Tensor((10,), dtype="float32")) -> R.Tensor((10,), dtype="float32 def test_add(): + batch_size = T.dynamic("batch_size") + vocab_size = T.dynamic("vocab_size") + @I.ir_module class Module: @Ts.prim_func(private=True) @@ -1584,11 +1612,9 @@ def cumsum(var_A: T.handle, var_A_1: T.handle, var_exclusive_scan_thrust: T.hand T.evaluate(0) @R.function - def main(probs: R.Tensor(("batch_size", "vocab_size"), dtype="float32")) -> R.Tensor( - ("batch_size", "vocab_size"), dtype="float32" + def main(probs: R.Tensor((batch_size, vocab_size), dtype="float32")) -> R.Tensor( + (batch_size, vocab_size), dtype="float32" ): - batch_size = T.int64() - vocab_size = T.int64() R.func_attr( { "relax.force_pure": True, @@ -1623,6 +1649,9 @@ def main(probs: R.Tensor(("batch_size", "vocab_size"), dtype="float32")) -> R.Te ) return lv1_1 + batch_size = T.dynamic("batch_size") + vocab_size = T.dynamic("vocab_size") + @I.ir_module class Expected: @Ts.prim_func(private=True) @@ -1630,11 +1659,9 @@ def cumsum(var_A: T.handle, var_A_1: T.handle, var_exclusive_scan_thrust: T.hand T.evaluate(0) @R.function - def main(probs: R.Tensor(("batch_size", "vocab_size"), dtype="float32")) -> R.Tensor( - ("batch_size", "vocab_size"), dtype="float32" + def main(probs: R.Tensor((batch_size, vocab_size), dtype="float32")) -> R.Tensor( + (batch_size, vocab_size), dtype="float32" ): - batch_size = T.int64() - vocab_size = T.int64() R.func_attr( { "relax.force_pure": True, diff --git a/tests/python/relax/test_utils.py b/tests/python/relax/test_utils.py index 33a68f3e183c..7f2a2d781b6c 100644 --- a/tests/python/relax/test_utils.py +++ b/tests/python/relax/test_utils.py @@ -42,8 +42,10 @@ def before(x: R.Tensor((3,), "float32"), y: R.Tensor((3,), "float32")): def test_copy_with_new_vars_copied_symbolic_vars(): + m = T.dynamic("m") + @R.function - def before(x: R.Tensor(("m",), "float32"), y: R.Tensor(("m",), "float32")): + def before(x: R.Tensor((m,), "float32"), y: R.Tensor((m,), "float32")): gv = R.add(x, y) return gv diff --git a/tests/python/relax/test_vm_build.py b/tests/python/relax/test_vm_build.py index 3804d72345d8..89c10c0180a2 100644 --- a/tests/python/relax/test_vm_build.py +++ b/tests/python/relax/test_vm_build.py @@ -110,10 +110,13 @@ def main( def test_match_check(exec_mode): + n = T.dynamic("n") + m = T.dynamic("m") + @tvm.script.ir_module class TestMatchCheck: @R.function - def foo(x: R.Tensor(["n", "m"], "int32"), y: R.Any) -> R.Tensor(["m", "n"], dtype=None): + def foo(x: R.Tensor([n, m], "int32"), y: R.Any) -> R.Tensor([m, n], dtype=None): return y mod = TestMatchCheck @@ -135,11 +138,13 @@ def foo(x: R.Tensor(["n", "m"], "int32"), y: R.Any) -> R.Tensor(["m", "n"], dtyp def test_vm_compile_stage2(exec_mode): + n = T.dynamic("n") + m = T.dynamic("m") + @tvm.script.ir_module class TestVMCompileStage2: @R.function def foo(x: R.Tensor(dtype="float32")) -> R.Shape: - n, m = T.int64(), T.int64() _ = R.match_cast(x, R.Tensor((n, m), "float32")) return R.shape([n * 2, m * 3]) @@ -189,12 +194,14 @@ def foo(x: R.Tensor((32, 16), "float32")) -> R.Tensor: def test_vm_compile_e2e(exec_mode): + n = T.dynamic("n") + m = T.dynamic("m") + @tvm.script.ir_module class TestVMCompileE2E: @R.function def foo(x: R.Tensor(dtype="float32")) -> R.Tensor: with R.dataflow(): - n, m = T.int64(), T.int64() _ = R.match_cast(x, R.Tensor((n, m), "float32")) y = R.call_dps_packed("test.vm.tile", (x), R.Tensor((n, m * 2), dtype="float32")) R.output(y) @@ -213,32 +220,35 @@ def foo(x: R.Tensor(dtype="float32")) -> R.Tensor: def test_vm_compile_e2e_func_param_with_shape(exec_mode): + m_tir_matmul = T.dynamic("m", "int32") + n_tir_matmul = T.dynamic("n", "int32") + k_tir_matmul = T.dynamic("k", "int32") + m_func = T.dynamic("m") + k_func = T.dynamic("k") + n_func = T.dynamic("n") + @tvm.script.ir_module class TestVMCompileE2E2: @Ts.prim_func def tir_matmul(x: T.handle, y: T.handle, z: T.handle) -> None: T.func_attr({"global_symbol": "tir_matmul"}) - m = T.int32() - n = T.int32() - k = T.int32() - A = T.match_buffer(x, (m, n)) - B = T.match_buffer(y, (n, k)) - C = T.match_buffer(z, (m, k)) + A = T.match_buffer(x, (m_tir_matmul, n_tir_matmul)) + B = T.match_buffer(y, (n_tir_matmul, k_tir_matmul)) + C = T.match_buffer(z, (m_tir_matmul, k_tir_matmul)) - for i, j, k in T.grid(m, k, n): + for i, j, k_tir_matmul_index in T.grid(m_tir_matmul, k_tir_matmul, n_tir_matmul): with Ts.sblock("matmul"): - vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) + vi, vj, vk = Ts.axis.remap("SSR", [i, j, k_tir_matmul_index]) with Ts.init(): C[vi, vj] = T.float32(0) C[vi, vj] = C[vi, vj] + A[vi, vk] * B[vk, vj] @R.function def func( - x: R.Tensor(("m", "n"), "float32"), w: R.Tensor(("n", "k"), "float32") + x: R.Tensor((m_func, n_func), "float32"), w: R.Tensor((n_func, k_func), "float32") ) -> R.Tensor: - m, k = T.int64(), T.int64() cls = TestVMCompileE2E2 - gv0 = R.call_tir(cls.tir_matmul, (x, w), R.Tensor((m, k), dtype="float32")) + gv0 = R.call_tir(cls.tir_matmul, (x, w), R.Tensor((m_func, k_func), dtype="float32")) return gv0 mod = TestVMCompileE2E2 @@ -418,7 +428,7 @@ def te_func(A, B): def test_vm_emit_te_dtype_change(exec_mode): bb = relax.BlockBuilder() - n = tirx.Var("n", "int64") + n = T.dynamic("n", "int64") x = relax.Var("x", R.Tensor([n], "float32")) # convert a tensor with dtype of float32 to int16 @@ -447,7 +457,7 @@ def te_func(A): def test_vm_emit_te_floor_symbolic_shape(exec_mode): bb = relax.BlockBuilder() - n = tirx.Var("n", "int64") + n = T.dynamic("n", "int64") x = relax.Var("x", R.Tensor([n], "float32")) def te_func(A): @@ -530,7 +540,7 @@ def run_and_check(): def test_vm_relax_symbolic_shape(exec_mode): bb = relax.BlockBuilder() - n = tirx.Var("n", "int64") + n = T.dynamic("n", "int64") x = relax.Var("x", R.Tensor([n], "float32")) y = relax.Var("y", R.Tensor([(n // 2) + 1], "float32")) @@ -561,12 +571,13 @@ def expected_output(): def test_vm_relax_symbolic_shape_tuple(exec_mode): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class mod: @R.function - def main(shape: R.Shape(["m", "n"])): - m = T.int64() - n = T.int64() + def main(shape: R.Shape([m, n])): return R.shape([2 * m, 3 * n]) target = tvm.target.Target("llvm", host="llvm") @@ -587,7 +598,7 @@ def main(shape: R.Shape(["m", "n"])): def test_vm_relax_dyn_tir_shape(exec_mode): # case where TIR variables are unbound in generated PrimFunc bb = relax.BlockBuilder() - n = tirx.Var("n", "int64") + n = T.dynamic("n", "int64") def te_func(A): C = te.compute((n + 1), lambda i: A[i]) @@ -619,7 +630,7 @@ def te_func(A): def test_vm_tuple(exec_mode): bb = relax.BlockBuilder() - n = tirx.Var("n", "int64") + n = T.dynamic("n", "int64") with bb.function("rx_func"): x = nn.Placeholder((n,), dtype="float32", name="x") @@ -700,21 +711,22 @@ def copy(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")): def test_sub_func_call(exec_mode): + m = T.dynamic("m", "int32") + n = T.dynamic("n", "int32") + k = T.dynamic("k", "int32") + @tvm.script.ir_module class TestVMSubFunction: @Ts.prim_func def tir_matmul(x: T.handle, y: T.handle, z: T.handle) -> None: T.func_attr({"global_symbol": "tir_matmul"}) - m = T.int32() - n = T.int32() - k = T.int32() A = T.match_buffer(x, (m, n)) B = T.match_buffer(y, (n, k)) C = T.match_buffer(z, (m, k)) - for i, j, k in T.grid(m, k, n): + for i, j, k_index in T.grid(m, k, n): with Ts.sblock("matmul"): - vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) + vi, vj, vk = Ts.axis.remap("SSR", [i, j, k_index]) with Ts.init(): C[vi, vj] = T.float32(0) C[vi, vj] = C[vi, vj] + A[vi, vk] * B[vk, vj] @@ -881,13 +893,15 @@ def main(x: R.Tensor((1,), "float32"), y: R.Tensor((1,), "float32")): assert timing_res.results +m = T.dynamic("m", "int32") +n = T.dynamic("n", "int32") + + @tvm.script.ir_module class TestVMSetInput: @Ts.prim_func def test_vm_mul(x: T.handle, y: T.handle, z: T.handle): T.func_attr({"global_symbol": "test_vm_mul"}) - m = T.int32() - n = T.int32() A = T.match_buffer(x, (m, n)) B = T.match_buffer(y, (m, n)) C = T.match_buffer(z, (m, n)) @@ -930,37 +944,39 @@ def main(x: R.Tensor((32, 32), "float32"), w: R.Tensor((32, 32), "float32")) -> def test_multi_systemlib(exec_mode): pytest.importorskip("cloudpickle") # needed by popen_pool.PopenWorker + N = T.dynamic("N") + m = T.dynamic("m") + @tvm.script.ir_module class ModA: I.module_attrs({"system_lib_prefix": "libA_"}) @Ts.prim_func def tir_init(x_handle: T.handle): - N = T.int64() x = T.match_buffer(x_handle, [N], "float32") for i in range(N): x[i] = T.float32(0) @R.function - def main(s: R.Shape(["m"])) -> R.Tensor: - m = T.int64() + def main(s: R.Shape([m])) -> R.Tensor: gv0 = R.call_tir(ModA.tir_init, (), R.Tensor((m + 1,), dtype="float32")) return gv0 + N = T.dynamic("N") + m = T.dynamic("m") + @tvm.script.ir_module class ModB: I.module_attrs({"system_lib_prefix": "libB_"}) @Ts.prim_func def tir_init(x_handle: T.handle): - N = T.int64() x = T.match_buffer(x_handle, [N], "float32") for i in range(N): x[i] = T.float32(1) @R.function - def main(s: R.Shape(["m"])) -> R.Tensor: - m = T.int64() + def main(s: R.Shape([m])) -> R.Tensor: gv0 = R.call_tir(ModB.tir_init, (), R.Tensor((m,), dtype="float32")) return gv0 diff --git a/tests/python/relax/test_vm_builtin_lower.py b/tests/python/relax/test_vm_builtin_lower.py index 8cdfe2addc24..7084de145cf4 100644 --- a/tests/python/relax/test_vm_builtin_lower.py +++ b/tests/python/relax/test_vm_builtin_lower.py @@ -26,12 +26,14 @@ def test_vm_builtin_lower_mem_alloc_storage(): + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Before: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor: + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor: R.func_attr({"relax.force_pure": True}) - m, n = T.int64(), T.int64() storage = R.memory.alloc_storage(R.shape([m * n * 4]), 0, "global", "uint8") alloc = R.memory.alloc_tensor(storage, 0, R.shape([m, n]), "float32") @@ -41,13 +43,15 @@ def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor: gv0 = alloc return gv0 + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Expected: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor: + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor: # we expected RemovePurityChecking to have been called first R.func_attr({"relax.force_pure": True}) - m, n = T.int64(), T.int64() storage = R.vm.alloc_storage(R.shape([m * n * 4]), R.prim_value(0), "uint8", "global") alloc = R.vm.alloc_tensor(storage, R.prim_value(0), R.shape([m, n]), "float32") @@ -65,12 +69,14 @@ def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor: def test_vm_builtin_alloc_tensor_raises_error(): """R.builtin.alloc_tensor should be handled earlier""" + m = T.dynamic("m") + n = T.dynamic("n") + @I.ir_module class Before: @R.function - def main(x: R.Tensor(("m", "n"), "float32")) -> R.Tensor: + def main(x: R.Tensor((m, n), "float32")) -> R.Tensor: R.func_attr({"relax.force_pure": True}) - m, n = T.int64(), T.int64() alloc = R.builtin.alloc_tensor(R.shape([m, n]), runtime_device_index=0, dtype="float32") _ = R.call_packed( diff --git a/tests/python/relax/test_vm_codegen_only.py b/tests/python/relax/test_vm_codegen_only.py index 093c4771160c..609e9804187a 100644 --- a/tests/python/relax/test_vm_codegen_only.py +++ b/tests/python/relax/test_vm_codegen_only.py @@ -221,13 +221,15 @@ def test_shape_check_builtin(exec_mode): # 0: n, 1: m sindex = {"n": 0, "m": 1} + n = T.dynamic("n") + k = T.dynamic("k") + m = T.dynamic("m") + @tvm.script.ir_module class TestVMShapeCheck: @R.function(pure=False) - def main(x: R.Tensor(["n", "m"], "float32")) -> R.Shape(ndim=3): + def main(x: R.Tensor([n, m], "float32")) -> R.Shape(ndim=3): R.func_attr({"global_symbol": "main"}) - n = T.int64() - k = T.int64() shape_heap = R.call_builtin_with_ctx( "vm.builtin.alloc_shape_heap", [R.prim_value(3)], diff --git a/tests/python/runtime/test_runtime_extension.py b/tests/python/runtime/test_runtime_extension.py index 807ac9e458ed..70db1f792ac5 100644 --- a/tests/python/runtime/test_runtime_extension.py +++ b/tests/python/runtime/test_runtime_extension.py @@ -23,11 +23,12 @@ def test_dltensor_compatible(): + n = T.dynamic("n", "int32") + @I.ir_module class Module: @Ts.prim_func def arange(A: T.handle): - n = T.int32() Ab = T.match_buffer(A, (n,), "int64") for i in T.serial(n - 1): Ab[i + 1] = Ab[i] + T.int64(1) diff --git a/tests/python/s_tir/dlight/test_benchmark.py b/tests/python/s_tir/dlight/test_benchmark.py index 2c1884e0d3b3..5d2e753b4a2e 100644 --- a/tests/python/s_tir/dlight/test_benchmark.py +++ b/tests/python/s_tir/dlight/test_benchmark.py @@ -39,20 +39,23 @@ from tvm.script import s_tir as Ts from tvm.script import tirx as T - # The test function uses an undefined symbolic var in Relax. # In principle, this should be attached to an argument. # pylint: disable=no-self-argument,invalid-name,line-too-long,no-method-argument # fmt: off +full1_n = T.dynamic("n") +full2_n = T.dynamic("n") +matmul1_n = T.dynamic("n") +_test_n = T.dynamic("n") + @I.ir_module(check_well_formed=False) class Module: @Ts.prim_func def full1(var_T_full: T.handle): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) - n = T.int64() - T_full = T.match_buffer(var_T_full, (T.int64(1), T.int64(32), T.int64(1), n), "float16") + T_full = T.match_buffer(var_T_full, (T.int64(1), T.int64(32), T.int64(1), full1_n), "float16") # with Ts.sblock("root"): - for ax0, ax1, ax2, ax3 in T.grid(T.int64(1), T.int64(32), T.int64(1), n): + for ax0, ax1, ax2, ax3 in T.grid(T.int64(1), T.int64(32), T.int64(1), full1_n): with Ts.sblock("T_full"): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) Ts.reads() @@ -62,10 +65,9 @@ def full1(var_T_full: T.handle): @Ts.prim_func def full2(var_T_full: T.handle): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) - n = T.int64() - T_full = T.match_buffer(var_T_full, (T.int64(1), T.int64(32), n, T.int64(128)), "float16") + T_full = T.match_buffer(var_T_full, (T.int64(1), T.int64(32), full2_n, T.int64(128)), "float16") # with Ts.sblock("root"): - for ax0, ax1, ax2, ax3 in T.grid(T.int64(1), T.int64(32), n, T.int64(128)): + for ax0, ax1, ax2, ax3 in T.grid(T.int64(1), T.int64(32), full2_n, T.int64(128)): with Ts.sblock("T_full"): v_ax0, v_ax1, v_ax2, v_ax3 = Ts.axis.remap("SSSS", [ax0, ax1, ax2, ax3]) Ts.reads() @@ -75,11 +77,10 @@ def full2(var_T_full: T.handle): @Ts.prim_func def matmul1(var_A: T.handle, var_B: T.handle, matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16")): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) - n = T.int64() - A = T.match_buffer(var_A, (T.int64(1), T.int64(32), T.int64(1), n), "float16") - B = T.match_buffer(var_B, (T.int64(1), T.int64(32), n, T.int64(128)), "float16") + A = T.match_buffer(var_A, (T.int64(1), T.int64(32), T.int64(1), matmul1_n), "float16") + B = T.match_buffer(var_B, (T.int64(1), T.int64(32), matmul1_n, T.int64(128)), "float16") # with Ts.sblock("root"): - for i0, i1, i2, i3, k in T.grid(T.int64(1), T.int64(32), T.int64(1), T.int64(128), n): + for i0, i1, i2, i3, k in T.grid(T.int64(1), T.int64(32), T.int64(1), T.int64(128), matmul1_n): with Ts.sblock("matmul"): v_i0, v_i1, v_i2, v_i3, v_k = Ts.axis.remap("SSSSR", [i0, i1, i2, i3, k]) Ts.reads(A[v_i0, v_i1, v_i2, v_k], B[v_i0, v_i1, v_k, v_i3]) @@ -90,23 +91,23 @@ def matmul1(var_A: T.handle, var_B: T.handle, matmul: T.Buffer((T.int64(1), T.in @R.function def test(): - n = T.int64() R.func_attr({"tir_var_upper_bound": {"n": 2048}}) cls = Module with R.dataflow(): - lv1 = R.call_tir(cls.full1,(), out_ty=R.Tensor((1, 32, 1, n), dtype="float16")) - lv1_1 = R.call_tir(cls.full1,(), out_ty=R.Tensor((1, 32, 1, n), dtype="float16")) - lv1_2 = R.call_tir(cls.full1,(), out_ty=R.Tensor((1, 32, 1, n), dtype="float16")) - lv2 = R.call_tir(cls.full2,(), out_ty=R.Tensor((1, 32, n, 128), dtype="float16")) - lv2_1 = R.call_tir(cls.full2,(), out_ty=R.Tensor((1, 32, n, 128), dtype="float16")) + lv1 = R.call_tir(cls.full1,(), out_ty=R.Tensor((1, 32, 1, _test_n), dtype="float16")) + lv1_1 = R.call_tir(cls.full1,(), out_ty=R.Tensor((1, 32, 1, _test_n), dtype="float16")) + lv1_2 = R.call_tir(cls.full1,(), out_ty=R.Tensor((1, 32, 1, _test_n), dtype="float16")) + lv2 = R.call_tir(cls.full2,(), out_ty=R.Tensor((1, 32, _test_n, 128), dtype="float16")) + lv2_1 = R.call_tir(cls.full2,(), out_ty=R.Tensor((1, 32, _test_n, 128), dtype="float16")) lv3 = R.call_tir(cls.matmul1, (lv1, lv2), out_ty=R.Tensor((1, 32, 1, 128), dtype="float16")) R.output(lv3) return lv3 +m = T.dynamic("m") + @Ts.prim_func def cuda_workload(var_inp0: T.handle, inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), var_matmul: T.handle): T.func_attr({"tirx.is_scheduled": True}) - m = T.int64() inp0 = T.match_buffer(var_inp0, (T.int64(1), m, T.int64(4096))) matmul = T.match_buffer(var_matmul, (T.int64(1), m, T.int64(4096))) # with Ts.sblock("root"): diff --git a/tests/python/s_tir/dlight/test_cpu_gemv.py b/tests/python/s_tir/dlight/test_cpu_gemv.py index fa3e25958b7d..9658db089ac9 100644 --- a/tests/python/s_tir/dlight/test_cpu_gemv.py +++ b/tests/python/s_tir/dlight/test_cpu_gemv.py @@ -27,10 +27,11 @@ def test_gemv_basic(): # fmt: off + n = T.dynamic("n", "int32") + @Ts.prim_func(private=True) def before(lv1637: T.Buffer((1, 32, 1, 128), "float16"), p_lv1638: T.handle, p_lv1614: T.handle, p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() lv1638 = T.match_buffer(p_lv1638, (1, 32, n, 128), "float16") lv1614 = T.match_buffer(p_lv1614, (1, 1, 1, n), "float16") var_compute_intermediate = T.match_buffer(p_output0, (1, 32, 1, n)) @@ -72,10 +73,11 @@ def before(lv1637: T.Buffer((1, 32, 1, 128), "float16"), p_lv1638: T.handle, p_l Ts.writes(var_compute_intermediate[v_i0, v_i1, v_i2, v_i3]) var_compute_intermediate[v_i0, v_i1, v_i2, v_i3] = T.Cast("float32", var_T_minimum_intermediate[v_i0, v_i1, v_i2, v_i3]) + n = T.dynamic("n", "int32") + @Ts.prim_func(private=True) def expected(lv1637: T.Buffer((1, 32, 1, 128), "float16"), p_lv1638: T.handle, p_lv1614: T.handle, p_output0: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int32() lv1638 = T.match_buffer(p_lv1638, (1, 32, n, 128), "float16") lv1614 = T.match_buffer(p_lv1614, (1, 1, 1, n), "float16") var_compute_intermediate = T.match_buffer(p_output0, (1, 32, 1, n)) @@ -433,10 +435,11 @@ def expected(lv575: T.Buffer((1376, 4096), "uint32"), lv576: T.Buffer((344, 4096 def test_outer_reduction_adreno_dynamic(): # fmt: off + v = T.dynamic("v") + @Ts.prim_func(private=True) def before(p_lv612: T.handle, p_lv613: T.handle, lv1607: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - v = T.int64() lv612 = T.match_buffer(p_lv612, (T.int64(512), v), "uint32") lv613 = T.match_buffer(p_lv613, (T.int64(128), v), "float16") p_output0_intermediate = T.match_buffer(p_output0, (T.int64(1), T.int64(1), v)) @@ -464,10 +467,11 @@ def before(p_lv612: T.handle, p_lv613: T.handle, lv1607: T.Buffer((T.int64(1), T Ts.writes(p_output0_intermediate[v_i0, v_i1, v_i2]) p_output0_intermediate[v_i0, v_i1, v_i2] = T.Cast("float32", var_matmul_intermediate[v_i0, v_i1, v_i2]) + v = T.dynamic("v") + @Ts.prim_func(private=True) def expected(p_lv612: T.handle, p_lv613: T.handle, lv1607: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - v = T.int64() lv612 = T.match_buffer(p_lv612, (T.int64(512), v), "uint32") lv613 = T.match_buffer(p_lv613, (T.int64(128), v), "float16") p_output0_intermediate = T.match_buffer(p_output0, (T.int64(1), T.int64(1), v)) diff --git a/tests/python/s_tir/dlight/test_gpu_fallback.py b/tests/python/s_tir/dlight/test_gpu_fallback.py index 01dec2daf666..d371651245bc 100644 --- a/tests/python/s_tir/dlight/test_gpu_fallback.py +++ b/tests/python/s_tir/dlight/test_gpu_fallback.py @@ -132,6 +132,14 @@ def main(A: T.Buffer((1, 6144), "float32"), B: T.Buffer((1,), "float32")): def test_fallback_irregular_spatial(): + nhead = T.dynamic("nhead", "int32") + nlayer = T.dynamic("nlayer", "int32") + seqlen = T.dynamic("seqlen", "int32") + npage = T.dynamic("npage", "int32") + page_size = T.dynamic("page_size", "int32") + num_total_pages = T.dynamic("num_total_pages", "int32") + num_total_seqs_plus_1 = T.dynamic("num_total_seqs_plus_1", "int32") + @Ts.prim_func(private=True) def func( var_pages: T.handle, @@ -140,14 +148,6 @@ def func( var_values: T.handle, seq_id: T.int32, ): - nhead = T.int32() - nlayer = T.int32() - seqlen = T.int32() - npage = T.int32() - page_size = T.int32() - num_total_pages = T.int32() - num_total_seqs_plus_1 = T.int32() - pages = T.match_buffer(var_pages, (num_total_pages, nlayer, nhead, page_size), "float16") page_table_indptr = T.match_buffer(var_page_table_indptr, (num_total_seqs_plus_1,), "int32") page_table_values = T.match_buffer(var_page_table_values, (npage,), "int32") @@ -164,16 +164,17 @@ def func( ] # fmt: off + nhead = T.dynamic("nhead", "int32") + nlayer = T.dynamic("nlayer", "int32") + seqlen = T.dynamic("seqlen", "int32") + npage = T.dynamic("npage", "int32") + page_size = T.dynamic("page_size", "int32") + num_total_pages = T.dynamic("num_total_pages", "int32") + num_total_seqs_plus_1 = T.dynamic("num_total_seqs_plus_1", "int32") + @Ts.prim_func(private=True) def expected(var_pages: T.handle, var_page_table_indptr: T.handle, var_page_table_values: T.handle, var_values: T.handle, seq_id: T.int32): T.func_attr({"tirx.is_scheduled": True}) - nhead = T.int32() - nlayer = T.int32() - seqlen = T.int32() - npage = T.int32() - page_size = T.int32() - num_total_pages = T.int32() - num_total_seqs_plus_1 = T.int32() pages = T.match_buffer(var_pages, (num_total_pages, nlayer, nhead, page_size), "float16") page_table_indptr = T.match_buffer(var_page_table_indptr, (num_total_seqs_plus_1,), "int32") diff --git a/tests/python/s_tir/dlight/test_gpu_gemv.py b/tests/python/s_tir/dlight/test_gpu_gemv.py index bd87f4a5d2c4..1ce652bba55b 100644 --- a/tests/python/s_tir/dlight/test_gpu_gemv.py +++ b/tests/python/s_tir/dlight/test_gpu_gemv.py @@ -54,10 +54,11 @@ def before( def test_gemv_basic(): # fmt: off + n = T.dynamic("n", "int32") + @Ts.prim_func(private=True) def before(lv1637: T.Buffer((1, 32, 1, 128), "float16"), p_lv1638: T.handle, p_lv1614: T.handle, p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() lv1638 = T.match_buffer(p_lv1638, (1, 32, n, 128), "float16") lv1614 = T.match_buffer(p_lv1614, (1, 1, 1, n), "float16") var_compute_intermediate = T.match_buffer(p_output0, (1, 32, 1, n)) @@ -99,10 +100,11 @@ def before(lv1637: T.Buffer((1, 32, 1, 128), "float16"), p_lv1638: T.handle, p_l Ts.writes(var_compute_intermediate[v_i0, v_i1, v_i2, v_i3]) var_compute_intermediate[v_i0, v_i1, v_i2, v_i3] = T.Cast("float32", var_T_minimum_intermediate[v_i0, v_i1, v_i2, v_i3]) + n = T.dynamic("n", "int32") + @Ts.prim_func(private=True) def expected(lv1637: T.Buffer((1, 32, 1, 128), "float16"), p_lv1638: T.handle, p_lv1614: T.handle, p_output0: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int32() lv1638 = T.match_buffer(p_lv1638, (1, 32, n, 128), "float16") lv1614 = T.match_buffer(p_lv1614, (1, 1, 1, n), "float16") var_compute_intermediate = T.match_buffer(p_output0, (1, 32, 1, n)) @@ -808,10 +810,11 @@ def expected(lv575: T.Buffer((1376, 4096), "uint32"), lv576: T.Buffer((344, 4096 def test_outer_reduction_adreno_dynamic(): # fmt: off + v = T.dynamic("v") + @Ts.prim_func(private=True) def before(p_lv612: T.handle, p_lv613: T.handle, lv1607: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - v = T.int64() lv612 = T.match_buffer(p_lv612, (T.int64(512), v), "uint32") lv613 = T.match_buffer(p_lv613, (T.int64(128), v), "float16") p_output0_intermediate = T.match_buffer(p_output0, (T.int64(1), T.int64(1), v)) @@ -839,10 +842,11 @@ def before(p_lv612: T.handle, p_lv613: T.handle, lv1607: T.Buffer((T.int64(1), T Ts.writes(p_output0_intermediate[v_i0, v_i1, v_i2]) p_output0_intermediate[v_i0, v_i1, v_i2] = T.Cast("float32", var_matmul_intermediate[v_i0, v_i1, v_i2]) + v = T.dynamic("v") + @Ts.prim_func(private=True) def expected(p_lv612: T.handle, p_lv613: T.handle, lv1607: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - v = T.int64() lv612 = T.match_buffer(p_lv612, (T.int64(512), v), "uint32") lv613 = T.match_buffer(p_lv613, (T.int64(128), v), "float16") p_output0_intermediate = T.match_buffer(p_output0, (T.int64(1), T.int64(1), v)) @@ -1145,6 +1149,8 @@ def before( def test_gemv_broadcast_epilogue(): # fmt: off + n = T.dynamic("n", "int32") + @Ts.prim_func(private=True) def before( A: T.Buffer((1, 32, 1, 128), "float16"), @@ -1152,7 +1158,6 @@ def before( p_C: T.handle, ): T.func_attr({"tirx.noalias": True}) - n = T.int32() B = T.match_buffer(p_B, (1, 32, n, 128), "float16") C = T.match_buffer(p_C, (1, 32, 2, 3, n), "float32") C_temp = Ts.sblock_alloc_buffer((1, 32, 1, n), "float16") diff --git a/tests/python/s_tir/dlight/test_gpu_general_reduction.py b/tests/python/s_tir/dlight/test_gpu_general_reduction.py index 80f1f817a606..f070bea0f015 100644 --- a/tests/python/s_tir/dlight/test_gpu_general_reduction.py +++ b/tests/python/s_tir/dlight/test_gpu_general_reduction.py @@ -91,12 +91,14 @@ def test_scalar_argmin_reduction_value_scope(): def test_softmax_1(): # fmt: off + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class Before: @Ts.prim_func def main(p_lv44: T.handle, p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - n, m = T.int64(), T.int64() lv44 = T.match_buffer(p_lv44, (T.int64(1), T.int64(32), n, m)) var_compute_intermediate = T.match_buffer(p_output0, (T.int64(1), T.int64(32), n, m), "float16") # with Ts.sblock("root"): @@ -140,12 +142,14 @@ def main(p_lv44: T.handle, p_output0: T.handle): Ts.writes(var_compute_intermediate[v_i0, v_i1, v_i2, v_i3]) var_compute_intermediate[v_i0, v_i1, v_i2, v_i3] = T.Cast("float16", var_T_softmax_norm_intermediate[v_i0, v_i1, v_i2, v_i3]) + n = T.dynamic("n") + m = T.dynamic("m") + @I.ir_module class After: @Ts.prim_func def main(p_lv44: T.handle, p_output0: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n, m = T.int64(), T.int64() lv44 = T.match_buffer(p_lv44, (T.int64(1), T.int64(32), n, m)) var_compute_intermediate = T.match_buffer(p_output0, (T.int64(1), T.int64(32), n, m), "float16") # with Ts.sblock("root"): @@ -371,12 +375,13 @@ def main(input: T.Buffer((T.int64(1), T.int64(4), T.int64(32), T.int64(8192)), " def test_layer_norm(): # fmt: off + n = T.dynamic("n") + @I.ir_module class Before: @Ts.prim_func def main(p_lv6: T.handle, weight1: T.Buffer((T.int64(2560),), "float32"), bias: T.Buffer((T.int64(2560),), "float32"), p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() lv6 = T.match_buffer(p_lv6, (T.int64(1), n, T.int64(2560))) var_compute_intermediate = T.match_buffer(p_output0, (T.int64(1), n, T.int64(2560)), "float16") # with Ts.sblock("root"): @@ -408,12 +413,13 @@ def main(p_lv6: T.handle, weight1: T.Buffer((T.int64(2560),), "float32"), bias: Ts.writes(var_compute_intermediate[v_i0, v_i1, v_i2]) var_compute_intermediate[v_i0, v_i1, v_i2] = T.Cast("float16", var_T_layer_norm_intermediate[v_i0, v_i1, v_i2]) + n = T.dynamic("n") + @I.ir_module class After: @Ts.prim_func def main(p_lv6: T.handle, weight1: T.Buffer((T.int64(2560),), "float32"), bias: T.Buffer((T.int64(2560),), "float32"), p_output0: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int64() lv6 = T.match_buffer(p_lv6, (T.int64(1), n, T.int64(2560))) var_compute_intermediate = T.match_buffer(p_output0, (T.int64(1), n, T.int64(2560)), "float16") # with Ts.sblock("root"): @@ -449,12 +455,13 @@ def main(p_lv6: T.handle, weight1: T.Buffer((T.int64(2560),), "float32"), bias: def test_rms_norm(): # fmt: off + n = T.dynamic("n") + @I.ir_module class Before: @Ts.prim_func def main(var_A: T.handle, B: T.Buffer((T.int64(4096),), "float16"), var_rms_norm: T.handle): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) - n = T.int64() A = T.match_buffer(var_A, (T.int64(1), n, T.int64(4096)), "float16") rms_norm_1 = T.match_buffer(var_rms_norm, (T.int64(1), n, T.int64(4096)), "float16") # with Ts.sblock("root"): @@ -474,12 +481,13 @@ def main(var_A: T.handle, B: T.Buffer((T.int64(4096),), "float16"), var_rms_norm Ts.writes(rms_norm_1[v_bsz, v_i, v_k]) rms_norm_1[v_bsz, v_i, v_k] = T.Cast("float16", T.Cast("float32", B[v_k]) * (T.Cast("float32", A[v_bsz, v_i, v_k]) / T.sqrt(Ared_temp[v_bsz, v_i] * T.float32(0.000244140625) + T.float32(9.9999999999999995e-07)))) + n = T.dynamic("n") + @I.ir_module class After: @Ts.prim_func def main(var_A: T.handle, B: T.Buffer((T.int64(4096),), "float16"), var_rms_norm: T.handle): T.func_attr({"op_pattern": 4, "tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int64() A = T.match_buffer(var_A, (T.int64(1), n, T.int64(4096)), "float16") rms_norm_1 = T.match_buffer(var_rms_norm, (T.int64(1), n, T.int64(4096)), "float16") # with Ts.sblock("root"): @@ -601,14 +609,15 @@ def main(A: T.Buffer((1, 2048), "float32"), B: T.Buffer((2048,), "float32"), C: def test_logsumexp(): + batch_size = T.dynamic("batch_size") + vocab_size = T.dynamic("vocab_size") + num_chunks = T.dynamic("num_chunks") + @I.ir_module class Before: @Ts.prim_func def compute_lse(var_A: T.handle, var_blocked_lse: T.handle): T.func_attr({"tirx.noalias": True}) - batch_size = T.int64() - vocab_size = T.int64() - num_chunks = T.int64() A = T.match_buffer(var_A, (batch_size, vocab_size), dtype="float32") blocked_lse = T.match_buffer(var_blocked_lse, (batch_size, num_chunks), dtype="float32") A_pad = Ts.sblock_alloc_buffer((batch_size, num_chunks, T.int64(4096)), dtype="float32") @@ -647,14 +656,16 @@ def compute_lse(var_A: T.handle, var_blocked_lse: T.handle): v0, v1, v2 = Ts.axis.remap("SSS", [l0, l1, l2]) blocked_lse[v0, v1] = T.log(temp_sum[v0, v1]) + temp_max[v0, v1] + batch_size = T.dynamic("batch_size") + vocab_size = T.dynamic("vocab_size") + num_chunks = T.dynamic("num_chunks") + @I.ir_module class After: @Ts.prim_func def compute_lse(var_A: T.handle, var_blocked_lse: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - batch_size, vocab_size = T.int64(), T.int64() A = T.match_buffer(var_A, (batch_size, vocab_size)) - num_chunks = T.int64() blocked_lse = T.match_buffer(var_blocked_lse, (batch_size, num_chunks)) temp_max_shared = Ts.sblock_alloc_buffer((batch_size, num_chunks), scope="shared") temp_sum_shared = Ts.sblock_alloc_buffer((batch_size, num_chunks), scope="shared") diff --git a/tests/python/s_tir/dlight/test_gpu_low_batch_gemv.py b/tests/python/s_tir/dlight/test_gpu_low_batch_gemv.py index c73ea0c2135c..ab2cf47826c6 100644 --- a/tests/python/s_tir/dlight/test_gpu_low_batch_gemv.py +++ b/tests/python/s_tir/dlight/test_gpu_low_batch_gemv.py @@ -30,10 +30,11 @@ def test_batch_decode_gemv(): # fmt: off + batch_size = T.dynamic("batch_size") + @Ts.prim_func(private=True) def before(lv429: T.Buffer((T.int64(4096), T.int64(3584)), "uint32"), lv430: T.Buffer((T.int64(4096), T.int64(896)), "float16"), p_lv807: T.handle, p_output0: T.handle): T.func_attr({"tirx.noalias": True, "tirx.HoistIfThenElseExprWithBlock": 1}) - batch_size = T.int64() lv807 = T.match_buffer(p_lv807, (batch_size, T.int64(1), T.int64(28672)), "float16") NT_matmul_intermediate = T.match_buffer(p_output0, (batch_size, T.int64(1), T.int64(4096)), "float16") # with Ts.sblock("root"): @@ -60,10 +61,11 @@ def before(lv429: T.Buffer((T.int64(4096), T.int64(3584)), "uint32"), lv430: T.B NT_matmul_intermediate[v_i0, v_i1, v_i2] = T.float16(0) NT_matmul_intermediate[v_i0, v_i1, v_i2] = NT_matmul_intermediate[v_i0, v_i1, v_i2] + lv807[v_i0, v_i1, v_k] * dequantize_intermediate_intermediate[v_i2, v_k] + batch_size = T.dynamic("batch_size") + @Ts.prim_func(private=True) def expected(lv429: T.Buffer((T.int64(4096), T.int64(3584)), "uint32"), lv430: T.Buffer((T.int64(4096), T.int64(896)), "float16"), p_lv807: T.handle, p_output0: T.handle): T.func_attr({"tirx.HoistIfThenElseExprWithBlock": 1, "tirx.is_scheduled": True, "tirx.noalias": True}) - batch_size = T.int64() lv807 = T.match_buffer(p_lv807, (batch_size, T.int64(1), T.int64(28672)), "float16") NT_matmul_intermediate = T.match_buffer(p_output0, (batch_size, T.int64(1), T.int64(4096)), "float16") # with Ts.sblock("root"): @@ -158,10 +160,11 @@ def test_batch_gemv(): K = 4096 # fmt: off + batch_size = T.dynamic("batch_size") + @Ts.prim_func(private=True) def before(var_A: T.handle, B: T.Buffer((T.int64(N), T.int64(K)), "float16"), var_NT_matmul: T.handle): T.func_attr({"tirx.noalias": True, "tirx.HoistIfThenElseExprWithBlock": 1}) - batch_size = T.int64() A = T.match_buffer(var_A, (batch_size, T.int64(1), T.int64(K)), "float16") NT_matmul = T.match_buffer(var_NT_matmul, (batch_size, T.int64(1), T.int64(N)), "float16") # with Ts.sblock("root"): @@ -174,10 +177,11 @@ def before(var_A: T.handle, B: T.Buffer((T.int64(N), T.int64(K)), "float16"), va NT_matmul[v_i0, v_i1, v_i2] = T.float16(0) NT_matmul[v_i0, v_i1, v_i2] = NT_matmul[v_i0, v_i1, v_i2] + A[v_i0, v_i1, v_k] * B[v_i2, v_k] + batch_size = T.dynamic("batch_size") + @Ts.prim_func(private=True) def expected(var_A: T.handle, B: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), var_NT_matmul: T.handle): T.func_attr({"tirx.HoistIfThenElseExprWithBlock": 1, "tirx.is_scheduled": True, "tirx.noalias": True}) - batch_size = T.int64() A = T.match_buffer(var_A, (batch_size, T.int64(1), T.int64(4096)), "float16") NT_matmul = T.match_buffer(var_NT_matmul, (batch_size, T.int64(1), T.int64(4096)), "float16") # with Ts.sblock("root"): @@ -259,10 +263,11 @@ def expected(var_A: T.handle, B: T.Buffer((T.int64(4096), T.int64(4096)), "float def test_reduction_symbolic_var(): # fmt: off + kv_seq_len = T.dynamic("kv_seq_len") + @Ts.prim_func(private=True) def before(var_A: T.handle, var_B: T.handle, matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float32")): T.func_attr({"tirx.noalias": True}) - kv_seq_len = T.int64() A = T.match_buffer(var_A, (T.int64(1), T.int64(32), T.int64(1), kv_seq_len)) B = T.match_buffer(var_B, (T.int64(1), T.int64(32), kv_seq_len, T.int64(128))) # with Ts.sblock("root"): @@ -282,10 +287,11 @@ def before(var_A: T.handle, var_B: T.handle, matmul: T.Buffer((T.int64(1), T.int def test_small_spatial_axis(): + batch_size = T.dynamic("batch_size") + @Ts.prim_func(private=True) def func(var_A: T.handle, B: T.Buffer((T.int64(8), T.int64(4096)), "float16"), var_C: T.handle): T.func_attr({"tirx.noalias": True}) - batch_size = T.int64() A = T.match_buffer(var_A, (batch_size, T.int64(4096)), "float16") C = T.match_buffer(var_C, (batch_size, T.int64(8)), "float16") for i0, i1, k in T.grid(batch_size, T.int64(8), T.int64(4096)): @@ -298,10 +304,11 @@ def func(var_A: T.handle, B: T.Buffer((T.int64(8), T.int64(4096)), "float16"), v C[v_i0, v_i1] = C[v_i0, v_i1] + A[v_i0, v_k] * B[v_i1, v_k] # fmt: off + batch_size = T.dynamic("batch_size") + @Ts.prim_func(private=True) def expected(var_A: T.handle, B: T.Buffer((T.int64(8), T.int64(4096)), "float16"), var_C: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - batch_size = T.int64() A = T.match_buffer(var_A, (batch_size, T.int64(4096)), "float16") C = T.match_buffer(var_C, (batch_size, T.int64(8)), "float16") # with Ts.sblock("root"): @@ -389,6 +396,8 @@ def expected(var_A: T.handle, B: T.Buffer((T.int64(8), T.int64(4096)), "float16" def test_outer_reduction(): # fmt: off + batch_size = T.dynamic("batch_size", "int32") + @Ts.prim_func(private=True) def before( B0: T.Buffer((512, 6144), "uint32"), @@ -396,7 +405,6 @@ def before( var_A: T.handle, var_C: T.handle ): - batch_size = T.int32() A = T.match_buffer(var_A, (batch_size, 1, 4096), "float16") C = T.match_buffer(var_C, (batch_size, 1, 6144), "float16") compute = Ts.sblock_alloc_buffer((4096, 6144), "float16") @@ -416,10 +424,11 @@ def before( C[v_i0, v_i1, v_i2] = T.float16(0) C[v_i0, v_i1, v_i2] = C[v_i0, v_i1, v_i2] + A[v_i0, v_i1, v_k] * B[v_k, v_i2] + batch_size = T.dynamic("batch_size", "int32") + @Ts.prim_func(private=True) def expected(B0: T.Buffer((512, 6144), "uint32"), B1: T.Buffer((128, 6144), "float16"), var_A: T.handle, var_C: T.handle): T.func_attr({"tirx.is_scheduled": True}) - batch_size = T.int32() A = T.match_buffer(var_A, (batch_size, 1, 4096), "float16") C = T.match_buffer(var_C, (batch_size, 1, 6144), "float16") # with Ts.sblock("root"): @@ -535,10 +544,11 @@ def expected(B0: T.Buffer((512, 6144), "uint32"), B1: T.Buffer((128, 6144), "flo def test_low_batch_gemv_cuda_target_without_max_shared_memory_per_block(): # fmt: off + batch_size = T.dynamic("batch_size") + @Ts.prim_func(private=True) def before(var_A: T.handle, B: T.Buffer((T.int64(128), T.int64(128)), "float16"), var_C: T.handle): T.func_attr({"tir.noalias": True}) - batch_size = T.int64() A = T.match_buffer(var_A, (batch_size, T.int64(1), T.int64(128)), "float16") C = T.match_buffer(var_C, (batch_size, T.int64(1), T.int64(128)), "float16") for i0, i1, i2, k in T.grid(batch_size, T.int64(1), T.int64(128), T.int64(128)): @@ -561,13 +571,14 @@ def before(var_A: T.handle, B: T.Buffer((T.int64(128), T.int64(128)), "float16") def test_low_batch_gemv_rejects_non_einsum_buffer_access(): + batch_size = T.dynamic("batch_size") + @Ts.prim_func(private=True) def before( var_A: T.handle, var_B: T.handle, var_C: T.handle, ): - batch_size = T.int64() A = T.match_buffer(var_A, (batch_size, 8), "float16") B = T.match_buffer(var_B, (4, batch_size + 8), "float16") C = T.match_buffer(var_C, (batch_size, 4), "float16") @@ -587,6 +598,8 @@ def before( def test_low_batch_gemv_broadcast_epilogue(): # fmt: off + batch_size = T.dynamic("batch_size") + @Ts.prim_func(private=True) def before( var_A: T.handle, @@ -594,7 +607,6 @@ def before( var_C: T.handle, ): T.func_attr({"tirx.noalias": True}) - batch_size = T.int64() A = T.match_buffer(var_A, (T.int64(1), batch_size, T.int64(1), T.int64(128)), "float16") C = T.match_buffer(var_C, (T.int64(1), batch_size, T.int64(2), T.int64(3), T.int64(128)), "float32") C_temp = Ts.sblock_alloc_buffer((T.int64(1), batch_size, T.int64(1), T.int64(128)), "float16") diff --git a/tests/python/s_tir/dlight/test_gpu_matmul.py b/tests/python/s_tir/dlight/test_gpu_matmul.py index 305dedf58317..0b4c02587bb8 100644 --- a/tests/python/s_tir/dlight/test_gpu_matmul.py +++ b/tests/python/s_tir/dlight/test_gpu_matmul.py @@ -26,9 +26,10 @@ def test_matmul(): # fmt: off + m = T.dynamic("m") + @Ts.prim_func(private=True) def before(var_inp0: T.handle, inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), var_matmul: T.handle): - m = T.int64() inp0 = T.match_buffer(var_inp0, (T.int64(1), m, T.int64(4096))) matmul = T.match_buffer(var_matmul, (T.int64(1), m, T.int64(4096))) for i0, i1, i2, k in T.grid(T.int64(1), m, T.int64(4096), T.int64(4096)): @@ -38,10 +39,11 @@ def before(var_inp0: T.handle, inp1: T.Buffer((T.int64(4096), T.int64(4096)), "f matmul[v_i0, v_i1, v_i2] = T.float32(0) matmul[v_i0, v_i1, v_i2] = matmul[v_i0, v_i1, v_i2] + inp0[v_i0, v_i1, v_k] * inp1[v_k, v_i2] + m = T.dynamic("m") + @Ts.prim_func(private=True) def expected(var_inp0: T.handle, inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), var_matmul: T.handle): T.func_attr({"tirx.is_scheduled": True}) - m = T.int64() inp0 = T.match_buffer(var_inp0, (T.int64(1), m, T.int64(4096))) matmul = T.match_buffer(var_matmul, (T.int64(1), m, T.int64(4096))) # with Ts.sblock("root"): @@ -118,9 +120,10 @@ def expected(var_inp0: T.handle, inp1: T.Buffer((T.int64(4096), T.int64(4096)), def test_matmul_int32(): # fmt: off + m = T.dynamic("m", "int32") + @Ts.prim_func(private=True) def func(var_inp0: T.handle, inp1: T.Buffer((4096, 4096), "float32"), var_matmul: T.handle): - m = T.int32() inp0 = T.match_buffer(var_inp0, (1, m, 4096)) matmul = T.match_buffer(var_matmul, (1, m, 4096)) for i0, i1, i2, k in T.grid(1, m, 4096, 4096): @@ -130,10 +133,11 @@ def func(var_inp0: T.handle, inp1: T.Buffer((4096, 4096), "float32"), var_matmul matmul[v_i0, v_i1, v_i2] = T.float32(0) matmul[v_i0, v_i1, v_i2] = matmul[v_i0, v_i1, v_i2] + inp0[v_i0, v_i1, v_k] * inp1[v_k, v_i2] + m = T.dynamic("m", "int32") + @Ts.prim_func(private=True) def expected(var_inp0: T.handle, inp1: T.Buffer((4096, 4096), "float32"), var_matmul: T.handle): T.func_attr({"tirx.is_scheduled": True}) - m = T.int32() inp0 = T.match_buffer(var_inp0, (1, m, 4096)) matmul = T.match_buffer(var_matmul, (1, m, 4096)) # with Ts.sblock("root"): @@ -350,10 +354,11 @@ def before(W: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), S: T.Buffer((T. def test_output_fp32(): # fmt: off + n = T.dynamic("n") + @Ts.prim_func(private=True) def before(lv13: T.Buffer((T.int64(4096), T.int64(512)), "uint32"), lv14: T.Buffer((T.int64(4096), T.int64(128)), "float16"), p_lv48: T.handle, lv13_1: T.Buffer((T.int64(4096),), "float16"), p_lv3: T.handle, p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() lv48 = T.match_buffer(p_lv48, (T.int64(1), n, T.int64(4096)), "float16") lv3 = T.match_buffer(p_lv3, (T.int64(1), n, T.int64(4096)), "float16") p_output0_intermediate = T.match_buffer(p_output0, (T.int64(1), n, T.int64(4096)), "float16") @@ -402,10 +407,11 @@ def before(lv13: T.Buffer((T.int64(4096), T.int64(512)), "uint32"), lv14: T.Buff Ts.writes(p_output0_intermediate[v_ax0, v_ax1, v_ax2]) p_output0_intermediate[v_ax0, v_ax1, v_ax2] = var_compute_intermediate_1[v_ax0, v_ax1, v_ax2] + lv3[v_ax0, v_ax1, v_ax2] + n = T.dynamic("n") + @Ts.prim_func(private=True) def expected(lv13: T.Buffer((T.int64(4096), T.int64(512)), "uint32"), lv14: T.Buffer((T.int64(4096), T.int64(128)), "float16"), p_lv48: T.handle, lv13_1: T.Buffer((T.int64(4096),), "float16"), p_lv3: T.handle, p_output0: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int64() lv48 = T.match_buffer(p_lv48, (T.int64(1), n, T.int64(4096)), "float16") lv3 = T.match_buffer(p_lv3, (T.int64(1), n, T.int64(4096)), "float16") p_output0_intermediate = T.match_buffer(p_output0, (T.int64(1), n, T.int64(4096)), "float16") @@ -484,10 +490,11 @@ def expected(lv13: T.Buffer((T.int64(4096), T.int64(512)), "uint32"), lv14: T.Bu def test_inline_consumer_chain(): # fmt: off + n = T.dynamic("n") + @Ts.prim_func(private=True) def before(p_lv26: T.handle, lv9: T.Buffer((T.int64(2048), T.int64(2048)), "float16"), p_lv52: T.handle, p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int64() lv26 = T.match_buffer(p_lv26, (n, T.int64(2048)), "float16") lv52 = T.match_buffer(p_lv52, (T.int64(1), n, T.int64(2048))) var_T_multiply_intermediate = T.match_buffer(p_output0, (n, T.int64(2048)), "float16") @@ -536,10 +543,11 @@ def before(p_lv26: T.handle, lv9: T.Buffer((T.int64(2048), T.int64(2048)), "floa Ts.writes(var_T_multiply_intermediate[v_ax0, v_ax1]) var_T_multiply_intermediate[v_ax0, v_ax1] = var_compute_intermediate[v_ax0, v_ax1] * var_T_multiply_intermediate_1[v_ax0, v_ax1] + n = T.dynamic("n") + @Ts.prim_func(private=True) def expected(p_lv26: T.handle, lv9: T.Buffer((T.int64(2048), T.int64(2048)), "float16"), p_lv52: T.handle, p_output0: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int64() lv26 = T.match_buffer(p_lv26, (n, T.int64(2048)), "float16") lv52 = T.match_buffer(p_lv52, (T.int64(1), n, T.int64(2048))) var_T_multiply_intermediate = T.match_buffer(p_output0, (n, T.int64(2048)), "float16") @@ -618,9 +626,10 @@ def expected(p_lv26: T.handle, lv9: T.Buffer((T.int64(2048), T.int64(2048)), "fl def test_matmul_android(): # fmt: off + m = T.dynamic("m") + @Ts.prim_func(private=True) def before(var_inp0: T.handle, inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), var_matmul: T.handle): - m = T.int64() inp0 = T.match_buffer(var_inp0, (T.int64(1), m, T.int64(4096))) matmul = T.match_buffer(var_matmul, (T.int64(1), m, T.int64(4096))) for i0, i1, i2, k in T.grid(T.int64(1), m, T.int64(4096), T.int64(4096)): @@ -630,10 +639,11 @@ def before(var_inp0: T.handle, inp1: T.Buffer((T.int64(4096), T.int64(4096)), "f matmul[v_i0, v_i1, v_i2] = T.float32(0) matmul[v_i0, v_i1, v_i2] = matmul[v_i0, v_i1, v_i2] + inp0[v_i0, v_i1, v_k] * inp1[v_k, v_i2] + m = T.dynamic("m") + @Ts.prim_func(private=True) def expected(var_inp0: T.handle, inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), var_matmul: T.handle): T.func_attr({"tirx.is_scheduled": True}) - m = T.int64() inp0 = T.match_buffer(var_inp0, (T.int64(1), m, T.int64(4096))) matmul = T.match_buffer(var_matmul, (T.int64(1), m, T.int64(4096))) # with Ts.sblock("root"): @@ -711,10 +721,11 @@ def expected(var_inp0: T.handle, inp1: T.Buffer((T.int64(4096), T.int64(4096)), def test_fused_dequant_matmul_android(): # fmt: off + seq_len = T.dynamic("seq_len") + @Ts.prim_func(private=True) def before(lv452: T.Buffer((T.int64(512), T.int64(12288)), "uint32"), lv453: T.Buffer((T.int64(128), T.int64(12288)), "float16"), p_rms_norm130: T.handle, transformer_h_0_attn_c_attn_bias3: T.Buffer((T.int64(12288),), "float16"), p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - seq_len = T.int64() rms_norm130 = T.match_buffer(p_rms_norm130, (T.int64(1), seq_len, T.int64(4096)), "float16") T_add_intermediate_intermediate = T.match_buffer(p_output0, (T.int64(1), seq_len, T.int64(12288)), "float16") # with Ts.sblock("root"): @@ -748,10 +759,11 @@ def before(lv452: T.Buffer((T.int64(512), T.int64(12288)), "uint32"), lv453: T.B Ts.writes(T_add_intermediate_intermediate[v_ax0, v_ax1, v_ax2]) T_add_intermediate_intermediate[v_ax0, v_ax1, v_ax2] = matmul_intermediate[v_ax0, v_ax1, v_ax2] + transformer_h_0_attn_c_attn_bias3[v_ax2] + seq_len = T.dynamic("seq_len") + @Ts.prim_func(private=True) def expected(lv452: T.Buffer((T.int64(512), T.int64(12288)), "uint32"), lv453: T.Buffer((T.int64(128), T.int64(12288)), "float16"), p_rms_norm130: T.handle, transformer_h_0_attn_c_attn_bias3: T.Buffer((T.int64(12288),), "float16"), p_output0: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - seq_len = T.int64() rms_norm130 = T.match_buffer(p_rms_norm130, (T.int64(1), seq_len, T.int64(4096)), "float16") T_add_intermediate_intermediate = T.match_buffer(p_output0, (T.int64(1), seq_len, T.int64(12288)), "float16") # with Ts.sblock("root"): diff --git a/tests/python/s_tir/dlight/test_gpu_matmul_tensorize.py b/tests/python/s_tir/dlight/test_gpu_matmul_tensorize.py index ddc80caaec20..bcb654daba49 100644 --- a/tests/python/s_tir/dlight/test_gpu_matmul_tensorize.py +++ b/tests/python/s_tir/dlight/test_gpu_matmul_tensorize.py @@ -40,6 +40,27 @@ def before(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float16" compute[v_i, v_j] = T.float16(0) compute[v_i, v_j] = compute[v_i, v_j] + X[v_i, v_k] * W[v_j, v_k] + A_3_s0 = T.dynamic("A_3_s0", "int32") + A_3_s1 = T.dynamic("A_3_s1", "int32") + C_4_s0 = T.dynamic("C_4_s0", "int32") + C_4_s1 = T.dynamic("C_4_s1", "int32") + C_s0 = T.dynamic("C_s0", "int32") + C_s1 = T.dynamic("C_s1", "int32") + A_s0 = T.dynamic("A_s0", "int32") + A_s1 = T.dynamic("A_s1", "int32") + C_1_s0 = T.dynamic("C_1_s0", "int32") + C_1_s1 = T.dynamic("C_1_s1", "int32") + A_1_s0 = T.dynamic("A_1_s0", "int32") + A_1_s1 = T.dynamic("A_1_s1", "int32") + C_2_s0 = T.dynamic("C_2_s0", "int32") + C_2_s1 = T.dynamic("C_2_s1", "int32") + A_2_s0 = T.dynamic("A_2_s0", "int32") + A_2_s1 = T.dynamic("A_2_s1", "int32") + B_s0 = T.dynamic("B_s0", "int32") + B_s1 = T.dynamic("B_s1", "int32") + C_3_s0 = T.dynamic("C_3_s0", "int32") + C_3_s1 = T.dynamic("C_3_s1", "int32") + @Ts.prim_func(private=True) def expected(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float16"), compute: T.Buffer((256, 256), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) @@ -66,7 +87,6 @@ def expected(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float1 v2_i_init_o = Ts.axis.spatial(1, 0) Ts.reads() Ts.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - C_s0, C_s1 = T.int32(), T.int32() C = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(C_s0, C_s1), scope="wmma.accumulator", offset_factor=16) T.tvm_fill_fragment(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.float32(0)) for ax3_0_0 in range(4, annotations={"software_pipeline_order": [0, 3, 1, 4, 5, 2, 6], "software_pipeline_stage": [0, 0, 0, 0, 0, 1, 1]}): @@ -103,9 +123,7 @@ def expected(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float1 v2_o = Ts.axis.spatial(16, ax3_0_0 * 4 + ax3_0_1 + ax1_0) Ts.reads(X_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(X_reindex_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_s0, A_s1 = T.int32(), T.int32() A = T.match_buffer(X_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(A_s0, A_s1), scope="shared.dyn", offset_factor=16) - C_1_s0, C_1_s1 = T.int32(), T.int32() C = T.match_buffer(X_reindex_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(C_1_s0, C_1_s1), scope="wmma.matrix_a", offset_factor=16) T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "row_major") for ax0_0 in T.unroll(2): @@ -116,9 +134,7 @@ def expected(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float1 v2_o = Ts.axis.spatial(16, ax3_0_0 * 4 + ax3_0_1 + ax1_0) Ts.reads(W_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(W_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_1_s0, A_1_s1 = T.int32(), T.int32() A = T.match_buffer(W_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(A_1_s0, A_1_s1), scope="shared.dyn", offset_factor=16) - C_2_s0, C_2_s1 = T.int32(), T.int32() C = T.match_buffer(W_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(C_2_s0, C_2_s1), scope="wmma.matrix_b", offset_factor=16) T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "col_major") for ax1_0_3, ax2_0_3 in T.grid(2, 2): @@ -135,11 +151,8 @@ def expected(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float1 v3_i_o = Ts.axis.reduce(1, 0) Ts.reads(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], X_reindex_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], W_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) Ts.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_2_s0, A_2_s1 = T.int32(), T.int32() A = T.match_buffer(X_reindex_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "float16", strides=(A_2_s0, A_2_s1), scope="wmma.matrix_a", offset_factor=16) - B_s0, B_s1 = T.int32(), T.int32() B = T.match_buffer(W_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "float16", strides=(B_s0, B_s1), scope="wmma.matrix_b", offset_factor=16) - C_3_s0, C_3_s1 = T.int32(), T.int32() C = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(C_3_s0, C_3_s1), scope="wmma.accumulator", offset_factor=16) T.tvm_mma_sync(C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, A.data, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, B.data, B.elem_offset // B.strides[0] // 16 * (B.strides[0] // 16) + B.elem_offset % B.strides[0] // 16, C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16) for ax0_0, ax1_0 in T.grid(2, 2): @@ -149,9 +162,7 @@ def expected(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float1 v2_o = Ts.axis.spatial(16, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax1_0) Ts.reads(compute_reindex_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(compute_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_3_s0, A_3_s1 = T.int32(), T.int32() A = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(A_3_s0, A_3_s1), scope="wmma.accumulator", offset_factor=16) - C_4_s0, C_4_s1 = T.int32(), T.int32() C = T.match_buffer(compute_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(C_4_s0, C_4_s1), scope="shared.dyn", offset_factor=16) T.tvm_store_matrix_sync(A.data, 16, 16, 16, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), C.data, C.elem_offset, C.strides[0] * 16, 2), C.strides[0], "row_major") for ax0_ax1_fused_0 in range(8): @@ -176,10 +187,11 @@ def expected(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float1 def test_matmul_tensorize_too_small(): # fmt: off + m = T.dynamic("m", "int32") + @Ts.prim_func(private=True) def before(var_X: T.handle, W: T.Buffer((15, 256), "float16"), var_compute: T.handle): T.func_attr({"tirx.noalias": True}) - m = T.int32() X = T.match_buffer(var_X, (m, 256), "float16") compute = T.match_buffer(var_compute, (m, 15)) # with Ts.sblock("root"): @@ -192,10 +204,11 @@ def before(var_X: T.handle, W: T.Buffer((15, 256), "float16"), var_compute: T.ha compute[v_i, v_j] = T.float32(0) compute[v_i, v_j] = compute[v_i, v_j] + T.Cast("float32", X[v_i, v_k]) * T.Cast("float32", W[v_j, v_k]) + m = T.dynamic("m", "int32") + @Ts.prim_func(private=True) def expected(var_X: T.handle, W: T.Buffer((15, 256), "float16"), var_compute: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - m = T.int32() X = T.match_buffer(var_X, (m, 256), "float16") compute = T.match_buffer(var_compute, (m, 15)) # with Ts.sblock("root"): @@ -272,10 +285,11 @@ def expected(var_X: T.handle, W: T.Buffer((15, 256), "float16"), var_compute: T. def test_matmul_tensorize_epilogue(): # fmt: off + n = T.dynamic("n", "int32") + @Ts.prim_func(private=True) def before(lv686: T.Buffer((T.int32(4096), T.int32(256)), "uint32"), lv687: T.Buffer((T.int32(4096), T.int32(64)), "float16"), p_lv42: T.handle, p_lv3: T.handle, p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() lv42 = T.match_buffer(p_lv42, (T.int32(1), n, T.int32(2048)), "float16") lv3 = T.match_buffer(p_lv3, (T.int32(1), n, T.int32(4096)), "float16") p_output0_intermediate = T.match_buffer(p_output0, (T.int32(1), n, T.int32(4096)), "float16") @@ -310,10 +324,31 @@ def before(lv686: T.Buffer((T.int32(4096), T.int32(256)), "uint32"), lv687: T.Bu Ts.writes(p_output0_intermediate[v_ax0, v_ax1, v_ax2]) p_output0_intermediate[v_ax0, v_ax1, v_ax2] = var_T_divide_intermediate[v_ax0, v_ax1, v_ax2] + var_NT_matmul_intermediate[v_ax0, v_ax1, v_ax2] + n = T.dynamic("n", "int32") + A_3_s0 = T.dynamic("A_3_s0", "int32") + A_3_s1 = T.dynamic("A_3_s1", "int32") + C_4_s0 = T.dynamic("C_4_s0", "int32") + C_4_s1 = T.dynamic("C_4_s1", "int32") + C_s0 = T.dynamic("C_s0", "int32") + C_s1 = T.dynamic("C_s1", "int32") + A_s0 = T.dynamic("A_s0", "int32") + A_s1 = T.dynamic("A_s1", "int32") + C_1_s0 = T.dynamic("C_1_s0", "int32") + C_1_s1 = T.dynamic("C_1_s1", "int32") + A_1_s0 = T.dynamic("A_1_s0", "int32") + A_1_s1 = T.dynamic("A_1_s1", "int32") + C_2_s0 = T.dynamic("C_2_s0", "int32") + C_2_s1 = T.dynamic("C_2_s1", "int32") + A_2_s0 = T.dynamic("A_2_s0", "int32") + A_2_s1 = T.dynamic("A_2_s1", "int32") + B_s0 = T.dynamic("B_s0", "int32") + B_s1 = T.dynamic("B_s1", "int32") + C_3_s0 = T.dynamic("C_3_s0", "int32") + C_3_s1 = T.dynamic("C_3_s1", "int32") + @Ts.prim_func(private=True) def expected(lv686: T.Buffer((4096, 256), "uint32"), lv687: T.Buffer((4096, 64), "float16"), p_lv42: T.handle, p_lv3: T.handle, p_output0: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int32() lv42 = T.match_buffer(p_lv42, (1, n, 2048), "float16") lv3 = T.match_buffer(p_lv3, (1, n, 4096), "float16") p_output0_intermediate = T.match_buffer(p_output0, (1, n, 4096), "float16") @@ -340,7 +375,6 @@ def expected(lv686: T.Buffer((4096, 256), "uint32"), lv687: T.Buffer((4096, 64), v2_i_init_o = Ts.axis.spatial(1, 0) Ts.reads() Ts.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - C_s0, C_s1 = T.int32(), T.int32() C = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(C_s0, C_s1), scope="wmma.accumulator", offset_factor=16) T.tvm_fill_fragment(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.float32(0)) for ax3_0_0 in range(32, annotations={"software_pipeline_order": [0, 3, 1, 4, 5, 2, 6], "software_pipeline_stage": [0, 0, 0, 0, 0, 1, 1]}): @@ -377,9 +411,7 @@ def expected(lv686: T.Buffer((4096, 256), "uint32"), lv687: T.Buffer((4096, 64), v2_o = Ts.axis.spatial(128, ax3_0_0 * 4 + ax3_0_1 + ax1_0) Ts.reads(lv42_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(lv42_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_s0, A_s1 = T.int32(), T.int32() A = T.match_buffer(lv42_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(A_s0, A_s1), scope="shared.dyn", offset_factor=16) - C_1_s0, C_1_s1 = T.int32(), T.int32() C = T.match_buffer(lv42_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(C_1_s0, C_1_s1), scope="wmma.matrix_a", offset_factor=16) T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "row_major") for ax0_0 in T.unroll(2): @@ -390,9 +422,7 @@ def expected(lv686: T.Buffer((4096, 256), "uint32"), lv687: T.Buffer((4096, 64), v2_o = Ts.axis.spatial(128, ax3_0_0 * 4 + ax3_0_1 + ax1_0) Ts.reads(p_output0_intermediate_1_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_1_s0, A_1_s1 = T.int32(), T.int32() A = T.match_buffer(p_output0_intermediate_1_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(A_1_s0, A_1_s1), scope="shared.dyn", offset_factor=16) - C_2_s0, C_2_s1 = T.int32(), T.int32() C = T.match_buffer(p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(C_2_s0, C_2_s1), scope="wmma.matrix_b", offset_factor=16) T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "col_major") for ax1_0_3, ax2_0_3 in T.grid(2, 2): @@ -409,11 +439,8 @@ def expected(lv686: T.Buffer((4096, 256), "uint32"), lv687: T.Buffer((4096, 64), v3_i_o = Ts.axis.reduce(1, 0) Ts.reads(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], lv42_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) Ts.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_2_s0, A_2_s1 = T.int32(), T.int32() A = T.match_buffer(lv42_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "float16", strides=(A_2_s0, A_2_s1), scope="wmma.matrix_a", offset_factor=16) - B_s0, B_s1 = T.int32(), T.int32() B = T.match_buffer(p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "float16", strides=(B_s0, B_s1), scope="wmma.matrix_b", offset_factor=16) - C_3_s0, C_3_s1 = T.int32(), T.int32() C = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(C_3_s0, C_3_s1), scope="wmma.accumulator", offset_factor=16) T.tvm_mma_sync(C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, A.data, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, B.data, B.elem_offset // B.strides[0] // 16 * (B.strides[0] // 16) + B.elem_offset % B.strides[0] // 16, C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16) for ax0_0, ax1_0 in T.grid(2, 2): @@ -423,9 +450,7 @@ def expected(lv686: T.Buffer((4096, 256), "uint32"), lv687: T.Buffer((4096, 64), v2_o = Ts.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax1_0) Ts.reads(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_3_s0, A_3_s1 = T.int32(), T.int32() A = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(A_3_s0, A_3_s1), scope="wmma.accumulator", offset_factor=16) - C_4_s0, C_4_s1 = T.int32(), T.int32() C = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=(C_4_s0, C_4_s1), scope="shared.dyn", offset_factor=16) T.tvm_store_matrix_sync(A.data, 16, 16, 16, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), C.data, C.elem_offset, C.strides[0] * 16, 2), C.strides[0], "row_major") for ax0_ax1_fused_0 in range(8): @@ -463,6 +488,27 @@ def before(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), com compute[v_i, v_j] = 0 compute[v_i, v_j] = compute[v_i, v_j] + T.Cast("int32", X[v_i, v_k]) * T.Cast("int32", W[v_j, v_k]) + A_3_s0 = T.dynamic("A_3_s0", "int32") + A_3_s1 = T.dynamic("A_3_s1", "int32") + C_4_s0 = T.dynamic("C_4_s0", "int32") + C_4_s1 = T.dynamic("C_4_s1", "int32") + C_s0 = T.dynamic("C_s0", "int32") + C_s1 = T.dynamic("C_s1", "int32") + A_s0 = T.dynamic("A_s0", "int32") + A_s1 = T.dynamic("A_s1", "int32") + C_1_s0 = T.dynamic("C_1_s0", "int32") + C_1_s1 = T.dynamic("C_1_s1", "int32") + A_1_s0 = T.dynamic("A_1_s0", "int32") + A_1_s1 = T.dynamic("A_1_s1", "int32") + C_2_s0 = T.dynamic("C_2_s0", "int32") + C_2_s1 = T.dynamic("C_2_s1", "int32") + A_2_s0 = T.dynamic("A_2_s0", "int32") + A_2_s1 = T.dynamic("A_2_s1", "int32") + B_s0 = T.dynamic("B_s0", "int32") + B_s1 = T.dynamic("B_s1", "int32") + C_3_s0 = T.dynamic("C_3_s0", "int32") + C_3_s1 = T.dynamic("C_3_s1", "int32") + @Ts.prim_func(private=True) def expected(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), compute: T.Buffer((256, 256), "int32")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) @@ -489,7 +535,6 @@ def expected(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), c v2_i_init_o = Ts.axis.spatial(1, 0) Ts.reads() Ts.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - C_s0, C_s1 = T.int32(), T.int32() C = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=(C_s0, C_s1), scope="wmma.accumulator", offset_factor=16) T.tvm_fill_fragment(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.float32(0)) for ax3_0_0 in T.serial(16, annotations={"software_pipeline_order": [0, 3, 1, 4, 5, 2, 6], "software_pipeline_stage": [0, 0, 0, 0, 0, 1, 1]}): @@ -526,9 +571,7 @@ def expected(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), c v2_o = Ts.axis.spatial(16, ax3_0_0 + ax1_0) Ts.reads(X_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(X_reindex_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_s0, A_s1 = T.int32(), T.int32() A = T.match_buffer(X_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=(A_s0, A_s1), scope="shared.dyn", offset_factor=16) - C_1_s0, C_1_s1 = T.int32(), T.int32() C = T.match_buffer(X_reindex_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=(C_1_s0, C_1_s1), scope="wmma.matrix_a", offset_factor=16) T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int8"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "row_major") for ax0_0 in T.unroll(2): @@ -539,9 +582,7 @@ def expected(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), c v2_o = Ts.axis.spatial(16, ax3_0_0 + ax1_0) Ts.reads(W_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(W_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_1_s0, A_1_s1 = T.int32(), T.int32() A = T.match_buffer(W_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=(A_1_s0, A_1_s1), scope="shared.dyn", offset_factor=16) - C_2_s0, C_2_s1 = T.int32(), T.int32() C = T.match_buffer(W_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=(C_2_s0, C_2_s1), scope="wmma.matrix_b", offset_factor=16) T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int8"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "col_major") for ax1_0_3, ax2_0_3 in T.grid(2, 2): @@ -558,11 +599,8 @@ def expected(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), c v3_i_o = Ts.axis.reduce(1, 0) Ts.reads(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], X_reindex_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], W_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) Ts.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_2_s0, A_2_s1 = T.int32(), T.int32() A = T.match_buffer(X_reindex_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=(A_2_s0, A_2_s1), scope="wmma.matrix_a", offset_factor=16) - B_s0, B_s1 = T.int32(), T.int32() B = T.match_buffer(W_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=(B_s0, B_s1), scope="wmma.matrix_b", offset_factor=16) - C_3_s0, C_3_s1 = T.int32(), T.int32() C = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=(C_3_s0, C_3_s1), scope="wmma.accumulator", offset_factor=16) T.tvm_mma_sync(C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, A.data, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, B.data, B.elem_offset // B.strides[0] // 16 * (B.strides[0] // 16) + B.elem_offset % B.strides[0] // 16, C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16) for ax0_0, ax1_0 in T.grid(2, 2): @@ -572,9 +610,7 @@ def expected(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), c v2_o = Ts.axis.spatial(16, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax1_0) Ts.reads(compute_reindex_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(compute_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_3_s0, A_3_s1 = T.int32(), T.int32() A = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=(A_3_s0, A_3_s1), scope="wmma.accumulator", offset_factor=16) - C_4_s0, C_4_s1 = T.int32(), T.int32() C = T.match_buffer(compute_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=(C_4_s0, C_4_s1), scope="shared.dyn", offset_factor=16) T.tvm_store_matrix_sync(A.data, 16, 16, 16, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int32"), C.data, C.elem_offset, C.strides[0] * 16, 2), C.strides[0], "row_major") for ax0_ax1_fused_0 in range(8): @@ -598,10 +634,11 @@ def expected(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), c def test_matmul_int8_tensorize_3d2d_dyn(): # fmt: off + m = T.dynamic("m", "int32") + @Ts.prim_func(private=True) def before(var_A: T.handle, B: T.Buffer((4096, 22016), "int8"), var_matmul: T.handle): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (1, m, 22016), "int8") matmul_1 = T.match_buffer(var_matmul, (1, m, 4096), "int32") # with Ts.sblock("root"): @@ -614,10 +651,31 @@ def before(var_A: T.handle, B: T.Buffer((4096, 22016), "int8"), var_matmul: T.ha matmul_1[v_i0, v_i1, v_i2] = 0 matmul_1[v_i0, v_i1, v_i2] = matmul_1[v_i0, v_i1, v_i2] + T.Cast("int32", A[v_i0, v_i1, v_k]) * T.Cast("int32", B[v_i2, v_k]) + m = T.dynamic("m", "int32") + A_3_s0 = T.dynamic("A_3_s0", "int32") + A_3_s1 = T.dynamic("A_3_s1", "int32") + C_4_s0 = T.dynamic("C_4_s0", "int32") + C_4_s1 = T.dynamic("C_4_s1", "int32") + C_s0 = T.dynamic("C_s0", "int32") + C_s1 = T.dynamic("C_s1", "int32") + A_s0 = T.dynamic("A_s0", "int32") + A_s1 = T.dynamic("A_s1", "int32") + C_1_s0 = T.dynamic("C_1_s0", "int32") + C_1_s1 = T.dynamic("C_1_s1", "int32") + A_1_s0 = T.dynamic("A_1_s0", "int32") + A_1_s1 = T.dynamic("A_1_s1", "int32") + C_2_s0 = T.dynamic("C_2_s0", "int32") + C_2_s1 = T.dynamic("C_2_s1", "int32") + A_2_s0 = T.dynamic("A_2_s0", "int32") + A_2_s1 = T.dynamic("A_2_s1", "int32") + B_s0 = T.dynamic("B_s0", "int32") + B_s1 = T.dynamic("B_s1", "int32") + C_3_s0 = T.dynamic("C_3_s0", "int32") + C_3_s1 = T.dynamic("C_3_s1", "int32") + @Ts.prim_func(private=True) def expected(var_A: T.handle, B: T.Buffer((4096, 22016), "int8"), var_matmul: T.handle): T.func_attr({"op_pattern": 4, "tirx.is_scheduled": True, "tirx.noalias": True}) - m = T.int32() A = T.match_buffer(var_A, (1, m, 22016), "int8") matmul_1 = T.match_buffer(var_matmul, (1, m, 4096), "int32") # with Ts.sblock("root"): @@ -643,7 +701,6 @@ def expected(var_A: T.handle, B: T.Buffer((4096, 22016), "int8"), var_matmul: T. v2_i_init_o = Ts.axis.spatial(1, 0) Ts.reads() Ts.writes(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - C_s0, C_s1 = T.int32(), T.int32() C = T.match_buffer(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=(C_s0, C_s1), scope="wmma.accumulator", offset_factor=16) T.tvm_fill_fragment(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.float32(0)) for ax3_0_0 in T.serial(1376, annotations={"software_pipeline_order": [0, 3, 1, 4, 5, 2, 6], "software_pipeline_stage": [0, 0, 0, 0, 0, 1, 1]}): @@ -680,9 +737,7 @@ def expected(var_A: T.handle, B: T.Buffer((4096, 22016), "int8"), var_matmul: T. v2_o = Ts.axis.spatial(1376, ax3_0_0 + ax1_0) Ts.reads(A_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(A_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_s0, A_s1 = T.int32(), T.int32() A_1 = T.match_buffer(A_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=(A_s0, A_s1), scope="shared.dyn", offset_factor=16) - C_1_s0, C_1_s1 = T.int32(), T.int32() C = T.match_buffer(A_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=(C_1_s0, C_1_s1), scope="wmma.matrix_a", offset_factor=16) T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int8"), A_1.data, A_1.elem_offset, A_1.strides[0] * 16, 1), A_1.strides[0], "row_major") for ax0_0 in T.unroll(2): @@ -693,9 +748,7 @@ def expected(var_A: T.handle, B: T.Buffer((4096, 22016), "int8"), var_matmul: T. v2_o = Ts.axis.spatial(1376, ax3_0_0 + ax1_0) Ts.reads(B_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(B_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_1_s0, A_1_s1 = T.int32(), T.int32() A_1 = T.match_buffer(B_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=(A_1_s0, A_1_s1), scope="shared.dyn", offset_factor=16) - C_2_s0, C_2_s1 = T.int32(), T.int32() C = T.match_buffer(B_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=(C_2_s0, C_2_s1), scope="wmma.matrix_b", offset_factor=16) T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int8"), A_1.data, A_1.elem_offset, A_1.strides[0] * 16, 1), A_1.strides[0], "col_major") for ax1_0_3, ax2_0_3 in T.grid(2, 2): @@ -712,11 +765,8 @@ def expected(var_A: T.handle, B: T.Buffer((4096, 22016), "int8"), var_matmul: T. v3_i_o = Ts.axis.reduce(1, 0) Ts.reads(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], A_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], B_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) Ts.writes(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_2_s0, A_2_s1 = T.int32(), T.int32() A_1 = T.match_buffer(A_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=(A_2_s0, A_2_s1), scope="wmma.matrix_a", offset_factor=16) - B_s0, B_s1 = T.int32(), T.int32() B_1 = T.match_buffer(B_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=(B_s0, B_s1), scope="wmma.matrix_b", offset_factor=16) - C_3_s0, C_3_s1 = T.int32(), T.int32() C = T.match_buffer(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=(C_3_s0, C_3_s1), scope="wmma.accumulator", offset_factor=16) T.tvm_mma_sync(C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, A_1.data, A_1.elem_offset // A_1.strides[0] // 16 * (A_1.strides[0] // 16) + A_1.elem_offset % A_1.strides[0] // 16, B_1.data, B_1.elem_offset // B_1.strides[0] // 16 * (B_1.strides[0] // 16) + B_1.elem_offset % B_1.strides[0] // 16, C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16) for ax0_0, ax1_0 in T.grid(2, 2): @@ -726,9 +776,7 @@ def expected(var_A: T.handle, B: T.Buffer((4096, 22016), "int8"), var_matmul: T. v2_o = Ts.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax1_0) Ts.reads(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) Ts.writes(matmul_1_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) - A_3_s0, A_3_s1 = T.int32(), T.int32() A_1 = T.match_buffer(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=(A_3_s0, A_3_s1), scope="wmma.accumulator", offset_factor=16) - C_4_s0, C_4_s1 = T.int32(), T.int32() C = T.match_buffer(matmul_1_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=(C_4_s0, C_4_s1), scope="shared.dyn", offset_factor=16) T.tvm_store_matrix_sync(A_1.data, 16, 16, 16, A_1.elem_offset // A_1.strides[0] // 16 * (A_1.strides[0] // 16) + A_1.elem_offset % A_1.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int32"), C.data, C.elem_offset, C.strides[0] * 16, 2), C.strides[0], "row_major") for ax0_ax1_fused_0 in range(8): @@ -753,13 +801,14 @@ def expected(var_A: T.handle, B: T.Buffer((4096, 22016), "int8"), var_matmul: T. def test_matmul_metal(): # fmt: off + batch_size = T.dynamic("batch_size", "int32") + @Ts.prim_func(private=True) def before( var_A: T.handle, B: T.Buffer((28672, 4096), "float16"), var_C: T.handle, ): - batch_size = T.int32() A = T.match_buffer(var_A, (batch_size, 1, 4096), "float16") C = T.match_buffer(var_C, (batch_size, 1, 28672), "float16") for i0, i1, i2, k in T.grid(batch_size, 1, 28672, 4096): @@ -770,10 +819,31 @@ def before( C[v_i0, v_i1, v_i2] = T.float16(0) C[v_i0, v_i1, v_i2] += A[v_i0, v_i1, v_k] * B[v_i2, v_k] + batch_size = T.dynamic("batch_size", "int32") + A_s0 = T.dynamic("A_s0", "int32") + A_s1 = T.dynamic("A_s1", "int32") + A_4_s0 = T.dynamic("A_4_s0", "int32") + A_4_s1 = T.dynamic("A_4_s1", "int32") + C_3_s0 = T.dynamic("C_3_s0", "int32") + C_3_s1 = T.dynamic("C_3_s1", "int32") + A_1_s0 = T.dynamic("A_1_s0", "int32") + A_1_s1 = T.dynamic("A_1_s1", "int32") + C_s0 = T.dynamic("C_s0", "int32") + C_s1 = T.dynamic("C_s1", "int32") + A_2_s0 = T.dynamic("A_2_s0", "int32") + A_2_s1 = T.dynamic("A_2_s1", "int32") + C_1_s0 = T.dynamic("C_1_s0", "int32") + C_1_s1 = T.dynamic("C_1_s1", "int32") + A_3_s0 = T.dynamic("A_3_s0", "int32") + A_3_s1 = T.dynamic("A_3_s1", "int32") + B_s0 = T.dynamic("B_s0", "int32") + B_s1 = T.dynamic("B_s1", "int32") + C_2_s0 = T.dynamic("C_2_s0", "int32") + C_2_s1 = T.dynamic("C_2_s1", "int32") + @Ts.prim_func(private=True) def expected(var_A: T.handle, B: T.Buffer((28672, 4096), "float16"), var_C: T.handle): T.func_attr({"tirx.is_scheduled": True}) - batch_size = T.int32() A = T.match_buffer(var_A, (batch_size, 1, 4096), "float16") C = T.match_buffer(var_C, (batch_size, 1, 28672), "float16") # with Ts.sblock("root"): @@ -795,7 +865,6 @@ def expected(var_A: T.handle, B: T.Buffer((28672, 4096), "float16"), var_C: T.ha v2_o = Ts.axis.spatial(3584, ax2_0 * 8 + ax2_1 * 2 + ax2_2_init + ax2_3_init_0) Ts.reads() Ts.writes(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_s0, A_s1 = T.int32(), T.int32() A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_s0, A_s1), scope="metal.simdgroup", offset_factor=1) T.metal.make_filled_simdgroup_matrix(A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, T.float32(0), 8, 8) for ax3_0 in range(128): @@ -831,9 +900,7 @@ def expected(var_A: T.handle, B: T.Buffer((28672, 4096), "float16"), var_C: T.ha v2_o = Ts.axis.spatial(512, ax3_0 * 4 + ax3_1 + ax1_0_1) Ts.reads(A_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) Ts.writes(A_reindex_pad_shared_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_1_s0, A_1_s1 = T.int32(), T.int32() A_1 = T.match_buffer(A_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_1_s0, A_1_s1), scope="shared", offset_factor=1) - C_s0, C_s1 = T.int32(), T.int32() C_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_s0, C_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_load(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), A_1.data, A_1.elem_offset, A_1.strides[0] * 8, 1), A_1.strides[0], 8, 8, T.bool(False)) for ax0_0, ax1_0_1 in T.grid(2, 1): @@ -843,9 +910,7 @@ def expected(var_A: T.handle, B: T.Buffer((28672, 4096), "float16"), var_C: T.ha v2_o = Ts.axis.spatial(512, ax3_0 * 4 + ax3_1 + ax1_0_1) Ts.reads(B_reindex_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) Ts.writes(B_reindex_shared_metal_simdgroup[v0_o, v2_o * 8:v2_o * 8 + 8, v1_o * 8:v1_o * 8 + 8]) - A_2_s0, A_2_s1 = T.int32(), T.int32() A_1 = T.match_buffer(B_reindex_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_2_s0, A_2_s1), scope="shared", offset_factor=1) - C_1_s0, C_1_s1 = T.int32(), T.int32() C_1 = T.match_buffer(B_reindex_shared_metal_simdgroup[v0_o, v2_o * 8:v2_o * 8 + 8, v1_o * 8:v1_o * 8 + 8], (8, 8), "float16", strides=(C_1_s0, C_1_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_load(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), A_1.data, A_1.elem_offset, A_1.strides[0] * 8, 1), A_1.strides[0], 8, 8, T.bool(True)) for ax1_2, ax2_2 in T.grid(2, 2): @@ -856,11 +921,8 @@ def expected(var_A: T.handle, B: T.Buffer((28672, 4096), "float16"), var_C: T.ha v3_o = Ts.axis.reduce(512, ax3_0 * 4 + ax3_1) Ts.reads(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], A_reindex_pad_shared_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v3_o * 8:v3_o * 8 + 8], B_reindex_shared_metal_simdgroup[0, v3_o * 8:v3_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) Ts.writes(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_3_s0, A_3_s1 = T.int32(), T.int32() A_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v3_o * 8:v3_o * 8 + 8], (8, 8), "float16", strides=(A_3_s0, A_3_s1), scope="metal.simdgroup", offset_factor=1) - B_s0, B_s1 = T.int32(), T.int32() B_1 = T.match_buffer(B_reindex_shared_metal_simdgroup[0, v3_o * 8:v3_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(B_s0, B_s1), scope="metal.simdgroup", offset_factor=1) - C_2_s0, C_2_s1 = T.int32(), T.int32() C_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_2_s0, C_2_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_multiply_accumulate(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, B_1.data, B_1.elem_offset // B_1.strides[0] // 8 * (B_1.strides[0] // 8) + B_1.elem_offset % B_1.strides[0] // 8, C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8) for ax0_1, ax1_0_1, ax2_0_1 in T.grid(1, 2, 2): @@ -870,9 +932,7 @@ def expected(var_A: T.handle, B: T.Buffer((28672, 4096), "float16"), var_C: T.ha v2_o = Ts.axis.spatial(3584, ax2_0 * 8 + ax2_1 * 2 + ax2_0_1) Ts.reads(C_reindex_pad_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) Ts.writes(C_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_4_s0, A_4_s1 = T.int32(), T.int32() A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_4_s0, A_4_s1), scope="metal.simdgroup", offset_factor=1) - C_3_s0, C_3_s1 = T.int32(), T.int32() C_1 = T.match_buffer(C_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_3_s0, C_3_s1), scope="shared", offset_factor=1) T.metal.simdgroup_store(A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), C_1.data, C_1.elem_offset, C_1.strides[0] * 8, 2), C_1.strides[0], 8, 8, T.bool(False)) for ax0_1, ax1_ax2_fused_0 in T.grid(1, 2): @@ -898,6 +958,8 @@ def expected(var_A: T.handle, B: T.Buffer((28672, 4096), "float16"), var_C: T.ha def test_matmul_metal_int4_quant(): # fmt: off + batch_size = T.dynamic("batch_size", "int32") + @Ts.prim_func(private=True) def before( B0: T.Buffer((28672, 512), "uint32"), @@ -905,7 +967,6 @@ def before( var_A: T.handle, var_C: T.handle ): - batch_size = T.int32() A = T.match_buffer(var_A, (batch_size, 1, 4096), "float16") C = T.match_buffer(var_C, (batch_size, 1, 28672), "float16") compute = Ts.sblock_alloc_buffer((28672, 4096), "float16") @@ -925,10 +986,31 @@ def before( C[v_i0, v_i1, v_i2] = T.float16(0) C[v_i0, v_i1, v_i2] = C[v_i0, v_i1, v_i2] + A[v_i0, v_i1, v_k] * B[v_i2, v_k] + batch_size = T.dynamic("batch_size", "int32") + A_s0 = T.dynamic("A_s0", "int32") + A_s1 = T.dynamic("A_s1", "int32") + A_4_s0 = T.dynamic("A_4_s0", "int32") + A_4_s1 = T.dynamic("A_4_s1", "int32") + C_3_s0 = T.dynamic("C_3_s0", "int32") + C_3_s1 = T.dynamic("C_3_s1", "int32") + A_1_s0 = T.dynamic("A_1_s0", "int32") + A_1_s1 = T.dynamic("A_1_s1", "int32") + C_s0 = T.dynamic("C_s0", "int32") + C_s1 = T.dynamic("C_s1", "int32") + A_2_s0 = T.dynamic("A_2_s0", "int32") + A_2_s1 = T.dynamic("A_2_s1", "int32") + C_1_s0 = T.dynamic("C_1_s0", "int32") + C_1_s1 = T.dynamic("C_1_s1", "int32") + A_3_s0 = T.dynamic("A_3_s0", "int32") + A_3_s1 = T.dynamic("A_3_s1", "int32") + B_s0 = T.dynamic("B_s0", "int32") + B_s1 = T.dynamic("B_s1", "int32") + C_2_s0 = T.dynamic("C_2_s0", "int32") + C_2_s1 = T.dynamic("C_2_s1", "int32") + @Ts.prim_func(private=True) def expected(B0: T.Buffer((28672, 512), "uint32"), B1: T.Buffer((28672, 128), "float16"), var_A: T.handle, var_C: T.handle): T.func_attr({"tirx.is_scheduled": True}) - batch_size = T.int32() A = T.match_buffer(var_A, (batch_size, 1, 4096), "float16") C = T.match_buffer(var_C, (batch_size, 1, 28672), "float16") # with Ts.sblock("root"): @@ -950,7 +1032,6 @@ def expected(B0: T.Buffer((28672, 512), "uint32"), B1: T.Buffer((28672, 128), "f v2_o = Ts.axis.spatial(3584, ax2_0 * 8 + ax2_1 * 2 + ax2_2_init + ax2_3_init_0) Ts.reads() Ts.writes(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_s0, A_s1 = T.int32(), T.int32() A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_s0, A_s1), scope="metal.simdgroup", offset_factor=1) T.metal.make_filled_simdgroup_matrix(A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, T.float32(0), 8, 8) for ax3_0 in range(128): @@ -986,9 +1067,7 @@ def expected(B0: T.Buffer((28672, 512), "uint32"), B1: T.Buffer((28672, 128), "f v2_o = Ts.axis.spatial(512, ax3_0 * 4 + ax3_1 + ax1_0_1) Ts.reads(A_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) Ts.writes(A_reindex_pad_shared_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_1_s0, A_1_s1 = T.int32(), T.int32() A_1 = T.match_buffer(A_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_1_s0, A_1_s1), scope="shared", offset_factor=1) - C_s0, C_s1 = T.int32(), T.int32() C_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_s0, C_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_load(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), A_1.data, A_1.elem_offset, A_1.strides[0] * 8, 1), A_1.strides[0], 8, 8, T.bool(False)) for ax0_0, ax1_0_1 in T.grid(2, 1): @@ -998,9 +1077,7 @@ def expected(B0: T.Buffer((28672, 512), "uint32"), B1: T.Buffer((28672, 128), "f v2_o = Ts.axis.spatial(512, ax3_0 * 4 + ax3_1 + ax1_0_1) Ts.reads(B_reindex_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) Ts.writes(B_reindex_shared_metal_simdgroup[v0_o, v2_o * 8:v2_o * 8 + 8, v1_o * 8:v1_o * 8 + 8]) - A_2_s0, A_2_s1 = T.int32(), T.int32() A_1 = T.match_buffer(B_reindex_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_2_s0, A_2_s1), scope="shared", offset_factor=1) - C_1_s0, C_1_s1 = T.int32(), T.int32() C_1 = T.match_buffer(B_reindex_shared_metal_simdgroup[v0_o, v2_o * 8:v2_o * 8 + 8, v1_o * 8:v1_o * 8 + 8], (8, 8), "float16", strides=(C_1_s0, C_1_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_load(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), A_1.data, A_1.elem_offset, A_1.strides[0] * 8, 1), A_1.strides[0], 8, 8, T.bool(True)) for ax1_2, ax2_2 in T.grid(2, 2): @@ -1011,11 +1088,8 @@ def expected(B0: T.Buffer((28672, 512), "uint32"), B1: T.Buffer((28672, 128), "f v3_o = Ts.axis.reduce(512, ax3_0 * 4 + ax3_1) Ts.reads(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], A_reindex_pad_shared_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v3_o * 8:v3_o * 8 + 8], B_reindex_shared_metal_simdgroup[0, v3_o * 8:v3_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) Ts.writes(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_3_s0, A_3_s1 = T.int32(), T.int32() A_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v3_o * 8:v3_o * 8 + 8], (8, 8), "float16", strides=(A_3_s0, A_3_s1), scope="metal.simdgroup", offset_factor=1) - B_s0, B_s1 = T.int32(), T.int32() B = T.match_buffer(B_reindex_shared_metal_simdgroup[0, v3_o * 8:v3_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(B_s0, B_s1), scope="metal.simdgroup", offset_factor=1) - C_2_s0, C_2_s1 = T.int32(), T.int32() C_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_2_s0, C_2_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_multiply_accumulate(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, B.data, B.elem_offset // B.strides[0] // 8 * (B.strides[0] // 8) + B.elem_offset % B.strides[0] // 8, C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8) for ax0_1, ax1_0_1, ax2_0_1 in T.grid(1, 2, 2): @@ -1025,9 +1099,7 @@ def expected(B0: T.Buffer((28672, 512), "uint32"), B1: T.Buffer((28672, 128), "f v2_o = Ts.axis.spatial(3584, ax2_0 * 8 + ax2_1 * 2 + ax2_0_1) Ts.reads(C_reindex_pad_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) Ts.writes(C_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_4_s0, A_4_s1 = T.int32(), T.int32() A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_4_s0, A_4_s1), scope="metal.simdgroup", offset_factor=1) - C_3_s0, C_3_s1 = T.int32(), T.int32() C_1 = T.match_buffer(C_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_3_s0, C_3_s1), scope="shared", offset_factor=1) T.metal.simdgroup_store(A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), C_1.data, C_1.elem_offset, C_1.strides[0] * 8, 2), C_1.strides[0], 8, 8, T.bool(False)) for ax0_1, ax1_ax2_fused_0 in T.grid(1, 2): diff --git a/tests/python/s_tir/dlight/test_gpu_reduction.py b/tests/python/s_tir/dlight/test_gpu_reduction.py index cb41bc7ceb47..f1561cf398e5 100644 --- a/tests/python/s_tir/dlight/test_gpu_reduction.py +++ b/tests/python/s_tir/dlight/test_gpu_reduction.py @@ -862,12 +862,13 @@ def main(lv9: T.Buffer((2560, 320), "uint32"), lv10: T.Buffer((2560, 80), "float def test_reduction_inner_spatial_choose_perfect_factor(): # fmt: off + n = T.dynamic("n") + @I.ir_module class Module: @Ts.prim_func def main(var_A: T.handle, var_B: T.handle, matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(100)), "float16")): T.func_attr({"tirx.noalias": True}) - n = T.int64() A = T.match_buffer(var_A, (T.int64(1), T.int64(32), T.int64(1), n), "float16") B = T.match_buffer(var_B, (T.int64(1), T.int64(32), n, T.int64(100)), "float16") # with Ts.sblock("root"): @@ -879,12 +880,13 @@ def main(var_A: T.handle, var_B: T.handle, matmul: T.Buffer((T.int64(1), T.int64 with Ts.init(): matmul[v_i0, v_i1, v_i2, v_i3] = T.float16(0) matmul[v_i0, v_i1, v_i2, v_i3] = matmul[v_i0, v_i1, v_i2, v_i3] + A[v_i0, v_i1, v_i2, v_k] * B[v_i0, v_i1, v_k, v_i3] + n = T.dynamic("n") + @I.ir_module class Expected: @Ts.prim_func def main(var_A: T.handle, var_B: T.handle, matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(100)), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int64() A = T.match_buffer(var_A, (T.int64(1), T.int64(32), T.int64(1), n), "float16") B = T.match_buffer(var_B, (T.int64(1), T.int64(32), n, T.int64(100)), "float16") # with Ts.sblock("root"): @@ -1100,12 +1102,13 @@ def main( def test_repeat_transpose_gemv(): # fmt: off + kv_seq_len = T.dynamic("kv_seq_len") + @I.ir_module class Before: @Ts.prim_func(private=True) def fused_relax_repeat_relax_permute_dims_relax_matmul1(p_lv716: T.handle, p_astype66: T.handle, var_matmul_intermediate: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16")): T.func_attr({"tirx.noalias": True}) - kv_seq_len = T.int64() lv716 = T.match_buffer(p_lv716, (T.int64(1), kv_seq_len, T.int64(8), T.int64(128)), "float16") astype66 = T.match_buffer(p_astype66, (T.int64(1), T.int64(32), T.int64(1), kv_seq_len), "float16") # with Ts.sblock("root"): @@ -1131,12 +1134,13 @@ def fused_relax_repeat_relax_permute_dims_relax_matmul1(p_lv716: T.handle, p_ast with Ts.init(): var_matmul_intermediate[v_i0, v_i1, v_i2, v_i3] = T.float16(0) var_matmul_intermediate[v_i0, v_i1, v_i2, v_i3] = var_matmul_intermediate[v_i0, v_i1, v_i2, v_i3] + astype66[v_i0, v_i1, v_i2, v_k] * var_T_transpose_intermediate[v_i0, v_i1, v_k, v_i3] + kv_seq_len = T.dynamic("kv_seq_len") + @I.ir_module class Expected: @Ts.prim_func(private=True) def fused_relax_repeat_relax_permute_dims_relax_matmul1(p_lv716: T.handle, p_astype66: T.handle, var_matmul_intermediate: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - kv_seq_len = T.int64() lv716 = T.match_buffer(p_lv716, (T.int64(1), kv_seq_len, T.int64(8), T.int64(128)), "float16") astype66 = T.match_buffer(p_astype66, (T.int64(1), T.int64(32), T.int64(1), kv_seq_len), "float16") # with Ts.sblock("root"): @@ -1182,6 +1186,8 @@ def fused_relax_repeat_relax_permute_dims_relax_matmul1(p_lv716: T.handle, p_ast def test_gemv_dyn_shape_epilogue(): + vocab_size = T.dynamic("vocab_size") + @I.ir_module class Module: @Ts.prim_func(private=True) @@ -1191,7 +1197,6 @@ def main( var_C: T.handle, ): T.func_attr({"tirx.noalias": True}) - vocab_size = T.int64() A = T.match_buffer(var_A, (T.int64(4096), vocab_size), "float16") C = T.match_buffer(var_C, (T.int64(1), T.int64(1), vocab_size)) C_temp = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), vocab_size), "float16") @@ -1213,12 +1218,13 @@ def main( C[v_i0, v_i1, v_i2] = T.Cast("float32", C_temp[v_i0, v_i1, v_i2]) # fmt: off + vocab_size = T.dynamic("vocab_size") + @I.ir_module class Expected: @Ts.prim_func(private=True) def main(var_A: T.handle, B: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), var_C: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - vocab_size = T.int64() A = T.match_buffer(var_A, (T.int64(4096), vocab_size), "float16") C = T.match_buffer(var_C, (T.int64(1), T.int64(1), vocab_size)) # with Ts.sblock("root"): diff --git a/tests/python/s_tir/dlight/test_gpu_rmsnorm.py b/tests/python/s_tir/dlight/test_gpu_rmsnorm.py index 04e95d5f7321..238614bbc59c 100644 --- a/tests/python/s_tir/dlight/test_gpu_rmsnorm.py +++ b/tests/python/s_tir/dlight/test_gpu_rmsnorm.py @@ -36,12 +36,13 @@ def _check(mod_before: IRModule, mod_after: IRModule): def test_rms_norm_with_casting(): # fmt: off + n = T.dynamic("n", "int32") + @I.ir_module class Before: @Ts.prim_func def main(var_data: T.handle, weight: T.Buffer((4096,), "float16"), var_T_cast: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() data = T.match_buffer(var_data, (1, n, 4096), "float16") T_cast = T.match_buffer(var_T_cast, (1, n, 4096), "float16") # with Ts.sblock("root"): @@ -96,12 +97,13 @@ def main(var_data: T.handle, weight: T.Buffer((4096,), "float16"), var_T_cast: T Ts.writes(T_cast[v_ax0, v_ax1, v_ax2]) T_cast[v_ax0, v_ax1, v_ax2] = T.Cast("float16", T_rms_norm[v_ax0, v_ax1, v_ax2]) + n = T.dynamic("n", "int32") + @I.ir_module class After: @Ts.prim_func def main(var_data: T.handle, weight: T.Buffer((4096,), "float16"), var_T_cast: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int32() data = T.match_buffer(var_data, (1, n, 4096), "float16") T_cast = T.match_buffer(var_T_cast, (1, n, 4096), "float16") # with Ts.sblock("root"): @@ -168,12 +170,13 @@ def main(var_data: T.handle, weight: T.Buffer((4096,), "float16"), var_T_cast: T def test_rms_norm_without_casting(): # fmt: off + n = T.dynamic("n", "int32") + @I.ir_module class Before: @Ts.prim_func def main(var_data: T.handle, weight: T.Buffer((4096,), "float32"), var_T_cast: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() data = T.match_buffer(var_data, (1, n, 4096)) T_cast = T.match_buffer(var_T_cast, (1, n, 4096)) # with Ts.sblock("root"): @@ -214,12 +217,13 @@ def main(var_data: T.handle, weight: T.Buffer((4096,), "float32"), var_T_cast: T Ts.writes(T_cast[v_ax0, v_ax1, v_ax2]) T_cast[v_ax0, v_ax1, v_ax2] = T_rms_norm[v_ax0, v_ax1, v_ax2] + n = T.dynamic("n", "int32") + @I.ir_module class After: @Ts.prim_func def main(var_data: T.handle, weight: T.Buffer((4096,), "float32"), var_T_cast: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - n = T.int32() data = T.match_buffer(var_data, (1, n, 4096)) T_cast = T.match_buffer(var_T_cast, (1, n, 4096)) # with Ts.sblock("root"): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_parallel_vectorize_unroll.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_parallel_vectorize_unroll.py index b54f6a175938..eeca2655dde7 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_parallel_vectorize_unroll.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_parallel_vectorize_unroll.py @@ -182,12 +182,13 @@ def after_postproc_add( add_compute[v0, v1, v2, v3, v4] = lhs[v0, v1, v2, v3, v4] + rhs[v0, v1, v2, v3, v4] +n = T.dynamic("n") + @Ts.prim_func def before_postproc_dynamic_shape_vectorize( a: T.handle, b: T.handle, ) -> None: - n = T.int64() A = T.match_buffer(a, (n,), dtype="float32") B = T.match_buffer(b, (n,), dtype="float32") with Ts.sblock("root"): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_gpu_code.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_gpu_code.py index 188ce7452546..2338b9b0da1d 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_gpu_code.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_gpu_code.py @@ -394,6 +394,13 @@ def GmmCuda2(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), " Z[v0, v1, v2] = Z_local[v0, v1, v2] +s0 = T.dynamic("s0", "int32") +s0_1 = T.dynamic("s0_1", "int32") +s0_2 = T.dynamic("s0_2", "int32") +s1 = T.dynamic("s1", "int32") +s1_1 = T.dynamic("s1_1", "int32") +s1_2 = T.dynamic("s1_2", "int32") + @Ts.prim_func def GMMCUDATensorCore( X: T.Buffer((1024, 1024), "float16"), @@ -402,12 +409,6 @@ def GMMCUDATensorCore( ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - s0 = T.int32() - s0_1 = T.int32() - s0_2 = T.int32() - s1 = T.int32() - s1_1 = T.int32() - s1_2 = T.int32() # body # with Ts.sblock("root") Z_wmma_accumulator = Ts.sblock_alloc_buffer([1024, 1024], dtype="float32", scope="wmma.accumulator") diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py index 638d78fcf843..565bdd9cd2f8 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py @@ -633,6 +633,27 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " compute[i0_13, i1_13, i2_13, i3_13] = T.max(T.min(T_cast_4[i0_13, i1_13, i2_13, i3_13], T.uint8(255)), T.uint8(0)) +C_s0 = T.dynamic("C_s0", "int32") +C_s1 = T.dynamic("C_s1", "int32") +A_3_s0 = T.dynamic("A_3_s0", "int32") +A_3_s1 = T.dynamic("A_3_s1", "int32") +C_4_s0 = T.dynamic("C_4_s0", "int32") +C_4_s1 = T.dynamic("C_4_s1", "int32") +A_s0 = T.dynamic("A_s0", "int32") +A_s1 = T.dynamic("A_s1", "int32") +C_1_s0 = T.dynamic("C_1_s0", "int32") +C_1_s1 = T.dynamic("C_1_s1", "int32") +A_1_s0 = T.dynamic("A_1_s0", "int32") +A_1_s1 = T.dynamic("A_1_s1", "int32") +C_2_s0 = T.dynamic("C_2_s0", "int32") +C_2_s1 = T.dynamic("C_2_s1", "int32") +A_2_s0 = T.dynamic("A_2_s0", "int32") +A_2_s1 = T.dynamic("A_2_s1", "int32") +B_s0 = T.dynamic("B_s0", "int32") +B_s1 = T.dynamic("B_s1", "int32") +C_3_s0 = T.dynamic("C_3_s0", "int32") +C_3_s1 = T.dynamic("C_3_s1", "int32") + @tvm.script.ir_module class Conv2dInt8_tensorcore_scheduled: @Ts.prim_func @@ -657,7 +678,6 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " Ts.reads() Ts.writes(conv2d_nhwc_reindex_shared_wmma_accumulator[v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) Ts.sblock_attr({"meta_schedule.thread_extent_high_inclusive": 1024, "meta_schedule.thread_extent_low_inclusive": 32, "warp_execution": 1}) - C_s0, C_s1 = T.int32(), T.int32() C = T.match_buffer(conv2d_nhwc_reindex_shared_wmma_accumulator[v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int32", strides=(C_s0, C_s1), scope="wmma.accumulator", offset_factor=16) T.tvm_fill_fragment(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.float32(0)) for ax0_0, ax1_0, ax4_0_0 in T.grid(1, 1, 2): @@ -690,9 +710,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " v1_o = Ts.axis.spatial(4, ax4_0_0 * 2 + ax1_0_1) Ts.reads(pad_temp_reindex_shared[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16]) Ts.writes(pad_temp_reindex_shared_wmma_matrix_a[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16]) - A_s0, A_s1 = T.int32(), T.int32() A = T.match_buffer(pad_temp_reindex_shared[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16], (16, 16), "int8", strides=(A_s0, A_s1), scope="shared", offset_factor=16) - C_1_s0, C_1_s1 = T.int32(), T.int32() C = T.match_buffer(pad_temp_reindex_shared_wmma_matrix_a[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16], (16, 16), "int8", strides=(C_1_s0, C_1_s1), scope="wmma.matrix_a", offset_factor=16) T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int8"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "row_major") for ax0, ax1, ax2_0, ax3_0 in T.grid(1, 1, 1, 2): @@ -702,9 +720,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " v3_o = Ts.axis.spatial(4, ax4_0_0 * 2 + ax3_0) Ts.reads(p1_reindex_shared[v0_o, v1_o, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) Ts.writes(p1_reindex_shared_wmma_matrix_b[v0_o, v1_o, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) - A_1_s0, A_1_s1 = T.int32(), T.int32() A = T.match_buffer(p1_reindex_shared[v0_o, v1_o, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=(A_1_s0, A_1_s1), scope="shared", offset_factor=16) - C_2_s0, C_2_s1 = T.int32(), T.int32() C = T.match_buffer(p1_reindex_shared_wmma_matrix_b[v0_o, v1_o, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=(C_2_s0, C_2_s1), scope="wmma.matrix_b", offset_factor=16) T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int8"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "col_major") for ax2_0_3, ax3_0_3, ax0_2, ax1_2, ax4_0_2, ax2_0_4, ax3_0_4 in T.grid(1, 1, 1, 1, 2, 1, 1): @@ -717,11 +733,8 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " Ts.reads(conv2d_nhwc_reindex_shared_wmma_accumulator[v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], pad_temp_reindex_shared_wmma_matrix_a[v2_o * 16:v2_o * 16 + 16, v4_o * 16:v4_o * 16 + 16], p1_reindex_shared_wmma_matrix_b[v0_o, v1_o, v3_o * 16:v3_o * 16 + 16, v4_o * 16:v4_o * 16 + 16]) Ts.writes(conv2d_nhwc_reindex_shared_wmma_accumulator[v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) Ts.sblock_attr({"meta_schedule.thread_extent_high_inclusive": 1024, "meta_schedule.thread_extent_low_inclusive": 32, "warp_execution": 1}) - A_2_s0, A_2_s1 = T.int32(), T.int32() A = T.match_buffer(pad_temp_reindex_shared_wmma_matrix_a[v2_o * 16:v2_o * 16 + 16, v4_o * 16:v4_o * 16 + 16], (16, 16), "int8", strides=(A_2_s0, A_2_s1), scope="wmma.matrix_a", offset_factor=16) - B_s0, B_s1 = T.int32(), T.int32() B = T.match_buffer(p1_reindex_shared_wmma_matrix_b[v0_o, v1_o, v3_o * 16:v3_o * 16 + 16, v4_o * 16:v4_o * 16 + 16], (16, 16), "int8", strides=(B_s0, B_s1), scope="wmma.matrix_b", offset_factor=16) - C_3_s0, C_3_s1 = T.int32(), T.int32() C = T.match_buffer(conv2d_nhwc_reindex_shared_wmma_accumulator[v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int32", strides=(C_3_s0, C_3_s1), scope="wmma.accumulator", offset_factor=16) T.tvm_mma_sync(C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, A.data, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, B.data, B.elem_offset // B.strides[0] // 16 * (B.strides[0] // 16) + B.elem_offset % B.strides[0] // 16, C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16) for ax0_0, ax1_0 in T.grid(1, 1): @@ -730,9 +743,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " v1_o = Ts.axis.spatial(16, ax2_0_0_ax3_0_0_fused % 8 * 2 + ax2_0_2_ax3_0_2_fused % 2 + ax1_0) Ts.reads(conv2d_nhwc_reindex_shared_wmma_accumulator[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16]) Ts.writes(conv2d_nhwc_reindex_shared[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16]) - A_3_s0, A_3_s1 = T.int32(), T.int32() A = T.match_buffer(conv2d_nhwc_reindex_shared_wmma_accumulator[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16], (16, 16), "int32", strides=(A_3_s0, A_3_s1), scope="wmma.accumulator", offset_factor=16) - C_4_s0, C_4_s1 = T.int32(), T.int32() C = T.match_buffer(conv2d_nhwc_reindex_shared[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16], (16, 16), "int32", strides=(C_4_s0, C_4_s1), scope="shared", offset_factor=16) T.tvm_store_matrix_sync(A.data, 16, 16, 16, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int32"), C.data, C.elem_offset, C.strides[0] * 16, 2), C.strides[0], "row_major") for ax0, ax1_0 in T.grid(128, 2): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_compute_at.py b/tests/python/s_tir/schedule/test_tir_schedule_compute_at.py index 676765438909..ba9f75f63051 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_compute_at.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_compute_at.py @@ -1752,9 +1752,10 @@ def after(A: T.Buffer((1, 3, 5, 5, 16), "float32"), C: T.Buffer((1, 6, 5, 5, 8), def test_shape_var_as_bound(): # fmt: off + n = T.dynamic("n", "int32") + @Ts.prim_func def before(a: T.handle, b: T.handle, c: T.handle): - n = T.int32() A = T.match_buffer(a, (32, 1, 128)) B = T.match_buffer(b, (32, n, 128)) C = T.match_buffer(c, (32, 1, n)) @@ -1782,9 +1783,10 @@ def before(a: T.handle, b: T.handle, c: T.handle): C[v0, 0, v1] = T.float32(0) C[v0, 0, v1] = C[v0, 0, v1] + C_rf[vax2_fused_1, v0, 0, v1] + n = T.dynamic("n", "int32") + @Ts.prim_func def expected(A: T.Buffer((32, 1, 128), "float32"), b: T.handle, c: T.handle): - n = T.int32() B = T.match_buffer(b, (32, n, 128)) C = T.match_buffer(c, (32, 1, n)) # with Ts.sblock("root"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_compute_inline.py b/tests/python/s_tir/schedule/test_tir_schedule_compute_inline.py index 14524888cb79..903be55f8464 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_compute_inline.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_compute_inline.py @@ -1310,10 +1310,12 @@ def test_reverse_compute_inline_producer_is_reduction(): def test_compute_inline_softmax(): # fmt: off + n = T.dynamic("n") + m = T.dynamic("m") + @Ts.prim_func def before(p_lv44: T.handle, p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - n, m = T.int64(), T.int64() lv44 = T.match_buffer(p_lv44, (T.int64(1), T.int64(32), n, m)) var_compute_intermediate = T.match_buffer(p_output0, (T.int64(1), T.int64(32), n, m), "float16") T_softmax_maxelem = Ts.sblock_alloc_buffer((T.int64(1), T.int64(32), n)) @@ -1356,10 +1358,12 @@ def before(p_lv44: T.handle, p_output0: T.handle): Ts.writes(var_compute_intermediate[v_i0, v_i1, v_i2, v_i3]) var_compute_intermediate[v_i0, v_i1, v_i2, v_i3] = T.Cast("float16", var_T_softmax_norm_intermediate[v_i0, v_i1, v_i2, v_i3]) + n = T.dynamic("n") + m = T.dynamic("m") + @Ts.prim_func def after(p_lv44: T.handle, p_output0: T.handle): T.func_attr({"tirx.noalias": True}) - n, m = T.int64(), T.int64() lv44 = T.match_buffer(p_lv44, (T.int64(1), T.int64(32), n, m)) var_compute_intermediate = T.match_buffer(p_output0, (T.int64(1), T.int64(32), n, m), "float16") # with Ts.sblock("root"): @@ -1404,10 +1408,11 @@ def after(p_lv44: T.handle, p_output0: T.handle): def test_reverse_compute_inline_layer_norm(): # fmt: off + n = T.dynamic("n") + @Ts.prim_func def before(p_lv6: T.handle, weight1: T.Buffer((T.int64(2560),), "float32"), bias: T.Buffer((T.int64(2560),), "float32"), p_output0: T.handle): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - n = T.int64() lv6 = T.match_buffer(p_lv6, (T.int64(1), n, T.int64(2560))) var_compute_intermediate = T.match_buffer(p_output0, (T.int64(1), n, T.int64(2560)), "float16") A_red_temp_v0_shared = Ts.sblock_alloc_buffer((T.int64(1), n), scope="shared") @@ -1445,10 +1450,11 @@ def before(p_lv6: T.handle, weight1: T.Buffer((T.int64(2560),), "float32"), bias Ts.writes(var_compute_intermediate[v_i0, v_i1, v_i2]) var_compute_intermediate[v_i0, v_i1, v_i2] = T.Cast("float16", var_T_layer_norm_intermediate[v_i0, v_i1, v_i2]) + n = T.dynamic("n") + @Ts.prim_func def after(p_lv6: T.handle, weight1: T.Buffer((T.int64(2560),), "float32"), bias: T.Buffer((T.int64(2560),), "float32"), p_output0: T.handle): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - n = T.int64() lv6 = T.match_buffer(p_lv6, (T.int64(1), n, T.int64(2560))) var_compute_intermediate = T.match_buffer(p_output0, (T.int64(1), n, T.int64(2560)), "float16") # with Ts.sblock("root"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_pad_einsum.py b/tests/python/s_tir/schedule/test_tir_schedule_pad_einsum.py index 9af55dd1e288..95e40cc65325 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_pad_einsum.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_pad_einsum.py @@ -107,13 +107,14 @@ def matmul_expected( def test_pad_matmul(): # pylint: disable=no-member,invalid-name,unused-variable,unexpected-keyword-arg + n = T.dynamic("n", "int32") + @Ts.prim_func def matmul_before( a: T.handle, b: T.handle, c: T.handle, ) -> None: - n = T.int32() A = T.match_buffer(a, (128, 128), "float32") B = T.match_buffer(b, (n, 128), "float32") C = T.match_buffer(c, (128, n), "float32") @@ -124,13 +125,14 @@ def matmul_before( C[i, j] = T.float32(0) C[i, j] = C[i, j] + A[i, k] * B[j, k] + n = T.dynamic("n", "int32") + @Ts.prim_func def matmul_after( a: T.handle, b: T.handle, c: T.handle, ): - n = T.int32() A = T.match_buffer(a, (128, 128), "float32") B = T.match_buffer(b, (n, 128), "float32") C = T.match_buffer(c, (128, n), "float32") @@ -161,6 +163,8 @@ def matmul_after( def test_pad_matmul_2(): + n = T.dynamic("n", "int32") + @Ts.prim_func def before( a: T.handle, @@ -169,7 +173,6 @@ def before( d: T.handle, ): T.func_attr({"tirx.noalias": True}) - n = T.int32() A = T.match_buffer(a, (1, n, 4096)) B = T.match_buffer(b, (11008, 4096)) M = T.match_buffer(m, (1, n, 11008)) @@ -188,10 +191,11 @@ def before( v_ax0, v_ax1, v_ax2 = Ts.axis.remap("SSS", [ax0, ax1, ax2]) D[v_ax0, v_ax1, v_ax2] = M[v_ax0, v_ax1, v_ax2] * C[v_ax0, v_ax1, v_ax2] + n = T.dynamic("n", "int32") + @Ts.prim_func def after(a: T.handle, b: T.handle, m: T.handle, d: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() A = T.match_buffer(a, (1, n, 4096)) B = T.match_buffer(b, (11008, 4096)) M = T.match_buffer(m, (1, n, 11008)) @@ -231,6 +235,8 @@ def after(a: T.handle, b: T.handle, m: T.handle, d: T.handle): def test_pad_rms(): + n = T.dynamic("n", "int32") + @Ts.prim_func def before( a: T.handle, @@ -238,7 +244,6 @@ def before( r: T.handle, ): T.func_attr({"tirx.noalias": True}) - n = T.int32() A = T.match_buffer(a, (1, n, 4096)) W = T.match_buffer(w, (4096,), "float32") Result = T.match_buffer(r, (1, n, 4096), "float32") @@ -259,10 +264,11 @@ def before( / T.sqrt(S[v_bsz, v_i] * T.float32(0.000244140625) + T.float32(1e-6)) ) + n = T.dynamic("n", "int32") + @Ts.prim_func def after(a: T.handle, w: T.handle, r: T.handle): T.func_attr({"tirx.noalias": True}) - n = T.int32() A = T.match_buffer(a, (1, n, 4096)) W = T.match_buffer(w, (4096,), "float32") Result = T.match_buffer(r, (1, n, 4096)) diff --git a/tests/python/s_tir/schedule/test_tir_schedule_rfactor.py b/tests/python/s_tir/schedule/test_tir_schedule_rfactor.py index 2243ef4f8d6e..e8adb65ef897 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_rfactor.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_rfactor.py @@ -1000,6 +1000,9 @@ def argmax_split_body_bufferstore_value_not_var( # v_unbound is unbound +v_unbound = T.dynamic("v_unbound", "int32") + + @Ts.prim_func(check_well_formed=False) def argmax_split_body_bufferstore_value_unbound_var( idx: T.Buffer((128, 128), "int32"), @@ -1007,7 +1010,6 @@ def argmax_split_body_bufferstore_value_unbound_var( argmax_v0: T.Buffer((128,), "int32"), argmax_v1: T.Buffer((128,), "float32"), ) -> None: - v_unbound = T.int32() for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): i = Ts.axis.spatial(128, i0) diff --git a/tests/python/s_tir/schedule/test_tir_schedule_sampling.py b/tests/python/s_tir/schedule/test_tir_schedule_sampling.py index b7ed59fff7a2..9274accb33d0 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_sampling.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_sampling.py @@ -216,9 +216,10 @@ def test_sample_perfect_tile_after_copy(): def test_sample_perfect_tile_on_dynamic_loops(): """Currently dynamic loop is trivially tiled""" + n = T.dynamic("n", "int32") + @Ts.prim_func def workload(a: T.handle) -> None: - n = T.int32() A = T.match_buffer(a, (n, 1024)) for i, j in T.grid(n, 1024): with Ts.sblock("B"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_split_fuse.py b/tests/python/s_tir/schedule/test_tir_schedule_split_fuse.py index e051128a3ea7..cdb9aedacaaf 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_split_fuse.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_split_fuse.py @@ -393,10 +393,11 @@ def test_split_with_inferred_factor(): def test_split_with_dynamic_inferred_factor(): + N = T.dynamic("N", "int32") + M = T.dynamic("M", "int32") + @Ts.prim_func def before(a: T.handle, b: T.handle) -> None: - N = T.int32() - M = T.int32() A = T.match_buffer(a, (N, 128, M)) B = T.match_buffer(b, (N, 128, M)) for i, j, k in T.grid(N, 128, M): @@ -404,9 +405,11 @@ def before(a: T.handle, b: T.handle) -> None: vi, vj, vk = Ts.axis.remap("SSS", [i, j, k]) B[vi, vj, vk] = A[vi, vj, vk] * 2.0 + N = T.dynamic("N", "int32") + M = T.dynamic("M", "int32") + @Ts.prim_func def expected(a: T.handle, b: T.handle) -> None: - N, M = T.int32(), T.int32() A = T.match_buffer(a, (N, 128, M)) B = T.match_buffer(b, (N, 128, M)) for i_0, i_1, j_0, j_1, k_0, k_1 in T.grid((N + 15) // 16, 16, 4, 32, 16, (M + 15) // 16): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_tensorize.py b/tests/python/s_tir/schedule/test_tir_schedule_tensorize.py index cbd9722d28f0..da27923bf39d 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_tensorize.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_tensorize.py @@ -204,6 +204,10 @@ def matmul( C[vi, vj] = C[vi, vj] + A[vi, vk] * B[vj, vk] +A_elem_offset = T.dynamic("A_elem_offset", "int32") +B_elem_offset = T.dynamic("B_elem_offset", "int32") +C_elem_offset = T.dynamic("C_elem_offset", "int32") + @Ts.prim_func def tensorized_matmul(a: T.handle, b: T.handle, c: T.handle) -> None: C = T.match_buffer(c, [128, 128], elem_offset=0, align=64, offset_factor=1) @@ -227,9 +231,6 @@ def tensorized_matmul(a: T.handle, b: T.handle, c: T.handle) -> None: ] ) Ts.writes(C[vi * 16 : vi * 16 + 16, vj * 16 : vj * 16 + 16]) - A_elem_offset = T.int32() - B_elem_offset = T.int32() - C_elem_offset = T.int32() A_sub = T.match_buffer( A[vi * 16 : vi * 16 + 16, vk * 16 : vk * 16 + 16], [16, 16], @@ -277,6 +278,10 @@ def batch_matmul( C[vn, vi, vj] = C[vn, vi, vj] + A[vn, vi, vk] * B[vn, vj, vk] +A_elem_offset = T.dynamic("A_elem_offset", "int32") +B_elem_offset = T.dynamic("B_elem_offset", "int32") +C_elem_offset = T.dynamic("C_elem_offset", "int32") + @Ts.prim_func def tensorized_batch_matmul_mma( A: T.Buffer((16, 128, 128), "float32"), @@ -299,9 +304,6 @@ def tensorized_batch_matmul_mma( B[vn : vn + 1, vj * 16 : vj * 16 + 16, vk * 16 : vk * 16 + 16], ) Ts.writes(C[vn : vn + 1, vi * 16 : vi * 16 + 16, vj * 16 : vj * 16 + 16]) - A_elem_offset = T.int32() - B_elem_offset = T.int32() - C_elem_offset = T.int32() A_sub = T.match_buffer( A[vn : vn + 1, vi * 16 : vi * 16 + 16, vk * 16 : vk * 16 + 16], (16, 16), @@ -437,6 +439,10 @@ def annotated_matmul( C[vi, vj] = C[vi, vj] + A[vi, vk] * B[vj, vk] +A_elem_offset = T.dynamic("A_elem_offset", "int32") +B_elem_offset = T.dynamic("B_elem_offset", "int32") +C_elem_offset = T.dynamic("C_elem_offset", "int32") + @Ts.prim_func def annotated_tensorized_matmul(a: T.handle, b: T.handle, c: T.handle) -> None: C = T.match_buffer(c, [128, 128], elem_offset=0, align=64, offset_factor=1) @@ -461,9 +467,6 @@ def annotated_tensorized_matmul(a: T.handle, b: T.handle, c: T.handle) -> None: ] ) Ts.writes(C[vi * 16 : vi * 16 + 16, vj * 16 : vj * 16 + 16]) - A_elem_offset = T.int32() - B_elem_offset = T.int32() - C_elem_offset = T.int32() A_sub = T.match_buffer( A[vi * 16 : vi * 16 + 16, vk * 16 : vk * 16 + 16], [16, 16], @@ -776,6 +779,10 @@ def matmul_int64_shape( vk = Ts.axis.reduce(T.int64(128), k_0 * T.int64(16) + k_1) C[vi, vj] = C[vi, vj] + A[vi, vk] * B[vj, vk] + A_elem_offset = T.dynamic("A_elem_offset") + B_elem_offset = T.dynamic("B_elem_offset") + C_elem_offset = T.dynamic("C_elem_offset") + @Ts.prim_func def tensorized_matmul_int64_shape( A: T.Buffer((T.int64(128), T.int64(128)), "float32"), @@ -799,9 +806,6 @@ def tensorized_matmul_int64_shape( ] ) Ts.writes(C[vi * T.int64(16) : vi * T.int64(16) + T.int64(16), vj * T.int64(16) : vj * T.int64(16) + T.int64(16)]) - A_elem_offset = T.int64() - B_elem_offset = T.int64() - C_elem_offset = T.int64() A_sub = T.match_buffer( A[vi * T.int64(16) : vi * T.int64(16) + T.int64(16), vk * T.int64(16) : vk * T.int64(16) + T.int64(16)], [T.int64(16), T.int64(16)], diff --git a/tests/python/s_tir/schedule/test_tir_schedule_transform_layout.py b/tests/python/s_tir/schedule/test_tir_schedule_transform_layout.py index 1856dd78c462..55dc212a40e2 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_transform_layout.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_transform_layout.py @@ -1171,10 +1171,11 @@ def func(A: T.Buffer(T.int64(16), "int32")): def test_transform_layout_with_symbolic_bound(): # fmt: off # pylint: disable=invalid-name,line-too-long,too-many-locals + n = T.dynamic("n") + @Ts.prim_func def before(a: T.handle, b: T.handle, c: T.handle): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - n = T.int64() A = T.match_buffer(a, (T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16") B = T.match_buffer(b, (T.int64(1), T.int64(32), n, T.int64(128)), "float16") C = T.match_buffer(c, (T.int64(1), T.int64(32), T.int64(1), n), "float16") @@ -1187,10 +1188,11 @@ def before(a: T.handle, b: T.handle, c: T.handle): C[v_i0, v_i1, v_i2, v_i3] = T.float16(0) C[v_i0, v_i1, v_i2, v_i3] = C[v_i0, v_i1, v_i2, v_i3] + A[v_i0, v_i1, v_i2, v_k] * B[v_i0, v_i1, v_i3, v_k] + n = T.dynamic("n") + @Ts.prim_func def after(a: T.handle, b: T.handle, c: T.handle): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - n = T.int64() A = T.match_buffer(a, (T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16") B = T.match_buffer(b, (T.int64(1), T.int64(32), n, T.int64(128)), "float16") C = T.match_buffer(c, (n * T.int64(32),), "float16") @@ -1221,10 +1223,11 @@ def after(a: T.handle, b: T.handle, c: T.handle): def test_transform_block_layout_with_symbolic_bound(): # fmt: off # pylint: disable=invalid-name,line-too-long,too-many-locals + n = T.dynamic("n") + @Ts.prim_func def before(a: T.handle, b: T.handle, c: T.handle): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - n = T.int64() A = T.match_buffer(a, (T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16") B = T.match_buffer(b, (T.int64(1), T.int64(32), n, T.int64(128)), "float16") C = T.match_buffer(c, (n * T.int64(32),), "float16") @@ -1237,10 +1240,11 @@ def before(a: T.handle, b: T.handle, c: T.handle): C[v_i1 * n + v_i3] = T.float16(0) C[v_i1 * n + v_i3] = C[v_i1 * n + v_i3] + A[v_i0, v_i1, v_i2, v_k] * B[v_i0, v_i1, v_i3, v_k] + n = T.dynamic("n") + @Ts.prim_func def after(a: T.handle, b: T.handle, c: T.handle): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - n = T.int64() A = T.match_buffer(a, (T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16") B = T.match_buffer(b, (T.int64(1), T.int64(32), n, T.int64(128)), "float16") C = T.match_buffer(c, (n * T.int64(32),), "float16") diff --git a/tests/python/s_tir/test_arith_domain_touched.py b/tests/python/s_tir/test_arith_domain_touched.py index 28779d1266a6..20d3a4845964 100644 --- a/tests/python/s_tir/test_arith_domain_touched.py +++ b/tests/python/s_tir/test_arith_domain_touched.py @@ -21,10 +21,11 @@ from tvm.script import s_tir as Ts from tvm.script import tirx as T +m = T.dynamic("m", "int32") + @Ts.prim_func def scalar_func(a: T.handle, b: T.handle): - m = T.int32() A = T.match_buffer(a, (100, m)) B = T.match_buffer(b, (100, m)) diff --git a/tests/python/s_tir/test_s_tir_renew_defs.py b/tests/python/s_tir/test_s_tir_renew_defs.py index 852a4fdb8efc..09371e9286e2 100644 --- a/tests/python/s_tir/test_s_tir_renew_defs.py +++ b/tests/python/s_tir/test_s_tir_renew_defs.py @@ -85,12 +85,13 @@ def _get_sblock(f): def test_match_buffer(): # well-formed checker complains about multiple definitions for variable A0_s1, # likely stemming from strides=[s, s] + s = T.dynamic("s", "int32") + e = T.dynamic("e", "int32") + @Ts.prim_func(check_well_formed=False) # A and B should be remapped def func_match_buffer(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32")): with Ts.sblock("root"): - s = T.int32() - e = T.int32() # A0 should be remapped A0 = T.match_buffer( A[0:128, 0:128], @@ -156,9 +157,10 @@ def _get_buffer_store_buffer(f): def test_symbolic_func(): + m = T.dynamic("m", "int32") + @Ts.prim_func def symbolic_func(a: T.handle, b: T.handle, n: T.int32): - m = T.int32() A = T.match_buffer(a, (n, m)) B = T.match_buffer(b, (n, m * 2)) for i, j in T.grid(n, m): @@ -171,9 +173,10 @@ def symbolic_func(a: T.handle, b: T.handle, n: T.int32): def test_buffer_params(): + m = T.dynamic("m") + @Ts.prim_func def main(a: T.handle, b: T.handle): - m = T.int64() A = T.match_buffer(a, (m * 2,)) B = T.match_buffer(b, (m, 2)) for i, j in T.grid(m, 2): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_compact_buffer_region.py b/tests/python/s_tir/transform/test_s_tir_transform_compact_buffer_region.py index f2b65aa95708..b3f8bd49f0b5 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_compact_buffer_region.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_compact_buffer_region.py @@ -1290,9 +1290,10 @@ def expected(): def test_loop_var_does_not_escape_compacted_buffer_extent(): + n = T.dynamic("n") + @Ts.prim_func(private=True) def before(a: T.handle): - n = T.int64() A = T.match_buffer(a, (n,), "int32") tmp = T.alloc_buffer((n,), "int32") for i in range(n): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_default_gpu_schedule.py b/tests/python/s_tir/transform/test_s_tir_transform_default_gpu_schedule.py index 39d03910ec27..c2587d73e07e 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_default_gpu_schedule.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_default_gpu_schedule.py @@ -26,6 +26,9 @@ def test_broadcast_to_symbolic(): # pylint: disable=no-self-argument,missing-class-docstring,line-too-long # fmt: off + x_0 = T.dynamic("x_0") + x_1 = T.dynamic("x_1") + @tvm.script.ir_module class Before: @Ts.prim_func @@ -34,8 +37,6 @@ def broadcast_to( var_T_broadcast_to: T.handle, ): T.func_attr({"tirx.noalias": True}) - x_0 = T.int64() - x_1 = T.int64() T_broadcast_to = T.match_buffer(var_T_broadcast_to, (x_0, x_1)) # with Ts.sblock("root"): for ax0, ax1 in T.grid(x_0, x_1): @@ -45,12 +46,14 @@ def broadcast_to( Ts.writes(T_broadcast_to[v_ax0, v_ax1]) T_broadcast_to[v_ax0, v_ax1] = rxplaceholder[v_ax0, T.int64(0)] + x_0 = T.dynamic("x_0") + x_1 = T.dynamic("x_1") + @tvm.script.ir_module class Expected: @Ts.prim_func def broadcast_to(rxplaceholder: T.Buffer((T.int64(3), T.int64(1)), "float32"), var_T_broadcast_to: T.handle): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) - x_0, x_1 = T.int64(), T.int64() T_broadcast_to = T.match_buffer(var_T_broadcast_to, (x_0, x_1)) for ax0_ax1_fused_1 in T.thread_binding(T.int64(256), thread="blockIdx.x"): for ax0_ax1_fused_2 in T.thread_binding(T.int64(1024), thread="threadIdx.x"): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_force_narrow_index_to_i32.py b/tests/python/s_tir/transform/test_s_tir_transform_force_narrow_index_to_i32.py index e9edc83b45aa..2aeac18f50ad 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_force_narrow_index_to_i32.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_force_narrow_index_to_i32.py @@ -150,11 +150,12 @@ def func(A: T.Buffer((128,), "int32"), B: T.Buffer((128,), "int32")): def test_metal_simdgroup_matmul_builds(): """Narrowing a DLight-scheduled Metal matmul keeps its tensorized blocks consistent.""" + n = T.dynamic("n") + @Ts.prim_func def main( var_A: T.handle, B: T.Buffer((T.int64(256), T.int64(256)), "float16"), var_C: T.handle ): - n = T.int64() A = T.match_buffer(var_A, (T.int64(1), n, T.int64(256)), "float16") C = T.match_buffer(var_C, (T.int64(1), n, T.int64(256)), "float16") for i0, i1, i2, k in T.grid(T.int64(1), n, T.int64(256), T.int64(256)): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_hoist_if.py b/tests/python/s_tir/transform/test_s_tir_transform_hoist_if.py index f94e9c31a690..71c374ef1988 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_hoist_if.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_hoist_if.py @@ -350,7 +350,7 @@ def test_no_hoisting_4(): dshape_inner = (33, 63) # Create iter_var for tx (used inside loop with T.attr) - tx_var = tvm.tirx.Var("threadIdx.x", "int32") + tx_var = T.dynamic("threadIdx.x", "int32") tx_iter = tvm.tirx.IterVar( tvm.ir.Range(0, dshape_inner[0]), tx_var, tvm.tirx.IterVar.ThreadIndex, "threadIdx.x" ) @@ -445,7 +445,7 @@ def test_hoisting_block_scope_2(): dshape = (32, 64) # Create iter_var for bx (used inside loop with T.attr) - bx_var = tvm.tirx.Var("blockIdx.x", "int32") + bx_var = T.dynamic("blockIdx.x", "int32") bx_iter = tvm.tirx.IterVar( tvm.ir.Range(0, dshape[1]), bx_var, tvm.tirx.IterVar.ThreadIndex, "blockIdx.x" ) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_inject_ptx_async_copy.py b/tests/python/s_tir/transform/test_s_tir_transform_inject_ptx_async_copy.py index a20178687c1b..5a8760a27fd2 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_inject_ptx_async_copy.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_inject_ptx_async_copy.py @@ -437,6 +437,27 @@ def simple_compute( @pytest.mark.gpu @pytest.mark.skipif(not env.has_cuda(), reason="need cuda") def test_vectorize_cp_async_in_if_then_else(postproc_if_missing_async_support): + C_s0_0 = T.dynamic("C_s0_0", "int32") + C_s1_0 = T.dynamic("C_s1_0", "int32") + A_s0_3 = T.dynamic("A_s0_3", "int32") + A_s1_3 = T.dynamic("A_s1_3", "int32") + C_s0_4 = T.dynamic("C_s0_4", "int32") + C_s1_4 = T.dynamic("C_s1_4", "int32") + A_s0_0 = T.dynamic("A_s0_0", "int32") + A_s1_0 = T.dynamic("A_s1_0", "int32") + C_s0_1 = T.dynamic("C_s0_1", "int32") + C_s1_1 = T.dynamic("C_s1_1", "int32") + A_s0_1 = T.dynamic("A_s0_1", "int32") + A_s1_1 = T.dynamic("A_s1_1", "int32") + C_s0_2 = T.dynamic("C_s0_2", "int32") + C_s1_2 = T.dynamic("C_s1_2", "int32") + A_s0_2 = T.dynamic("A_s0_2", "int32") + A_s1_2 = T.dynamic("A_s1_2", "int32") + B_s0 = T.dynamic("B_s0", "int32") + B_s1 = T.dynamic("B_s1", "int32") + C_s0_3 = T.dynamic("C_s0_3", "int32") + C_s1_3 = T.dynamic("C_s1_3", "int32") + @Ts.prim_func def complex_compute( A: T.Buffer((2, 16, 16, 1280), "float16"), @@ -474,15 +495,13 @@ def complex_compute( v_x_o * 16 : v_x_o * 16 + 16, v_y_o * 16 : v_y_o * 16 + 16 ] ) - C_s0 = T.int32() - C_s1 = T.int32() C = T.match_buffer( Conv_reindex_wmma_accumulator[ v_x_o * 16 : v_x_o * 16 + 16, v_y_o * 16 : v_y_o * 16 + 16 ], (16, 16), "float16", - strides=(C_s0, C_s1), + strides=(C_s0_0, C_s1_0), scope="wmma.accumulator", offset_factor=16, ) @@ -491,8 +510,8 @@ def complex_compute( 16, 16, 16, - C.ty.elem_offset // C_s0 // 16 * (C_s0 // 16) - + C.ty.elem_offset % C_s0 // 16, + C.ty.elem_offset // C_s0_0 // 16 * (C_s0_0 // 16) + + C.ty.elem_offset % C_s0_0 // 16, T.float32(0), ) for k_0_0 in T.serial( @@ -656,8 +675,6 @@ def complex_compute( v1_o * 16 : v1_o * 16 + 16, ] ) - A_s0 = T.int32() - A_s1 = T.int32() A_1 = T.match_buffer( data_im2col_reindex_shared_dyn[ v0_o * 16 : v0_o * 16 + 16, @@ -665,12 +682,10 @@ def complex_compute( ], (16, 16), "float16", - strides=(A_s0, A_s1), + strides=(A_s0_0, A_s1_0), scope="shared.dyn", offset_factor=16, ) - C_s0 = T.int32() - C_s1 = T.int32() C = T.match_buffer( data_im2col_reindex_shared_dyn_wmma_matrix_a[ v0_o * 16 : v0_o * 16 + 16, @@ -678,7 +693,7 @@ def complex_compute( ], (16, 16), "float16", - strides=(C_s0, C_s1), + strides=(C_s0_1, C_s1_1), scope="wmma.matrix_a", offset_factor=16, ) @@ -687,16 +702,16 @@ def complex_compute( 16, 16, 16, - C.ty.elem_offset // C_s0 // 16 * (C_s0 // 16) - + C.ty.elem_offset % C_s0 // 16, + C.ty.elem_offset // C_s0_1 // 16 * (C_s0_1 // 16) + + C.ty.elem_offset % C_s0_1 // 16, T.tvm_access_ptr( T.type_annotation("float16"), A_1.data, A_1.ty.elem_offset, - A_s0 * 16, + A_s0_0 * 16, 1, ), - A_s0, + A_s0_0, "row_major", ) for ax0_0, ax1_0 in T.grid(2, 1): @@ -717,8 +732,6 @@ def complex_compute( v1_o * 16 : v1_o * 16 + 16, ] ) - A_s0 = T.int32() - A_s1 = T.int32() A_1 = T.match_buffer( weight_flatten_reindex_shared_dyn[ v0_o * 16 : v0_o * 16 + 16, @@ -726,12 +739,10 @@ def complex_compute( ], (16, 16), "float16", - strides=(A_s0, A_s1), + strides=(A_s0_1, A_s1_1), scope="shared.dyn", offset_factor=16, ) - C_s0 = T.int32() - C_s1 = T.int32() C = T.match_buffer( weight_flatten_reindex_shared_dyn_wmma_matrix_b[ v0_o * 16 : v0_o * 16 + 16, @@ -739,7 +750,7 @@ def complex_compute( ], (16, 16), "float16", - strides=(C_s0, C_s1), + strides=(C_s0_2, C_s1_2), scope="wmma.matrix_b", offset_factor=16, ) @@ -748,16 +759,16 @@ def complex_compute( 16, 16, 16, - C.ty.elem_offset // C_s0 // 16 * (C_s0 // 16) - + C.ty.elem_offset % C_s0 // 16, + C.ty.elem_offset // C_s0_2 // 16 * (C_s0_2 // 16) + + C.ty.elem_offset % C_s0_2 // 16, T.tvm_access_ptr( T.type_annotation("float16"), A_1.data, A_1.ty.elem_offset, - A_s0 * 16, + A_s0_1 * 16, 1, ), - A_s0, + A_s0_1, "col_major", ) for x_0_2, y_0_2 in T.grid(2, 2): @@ -785,8 +796,6 @@ def complex_compute( v_y_o * 16 : v_y_o * 16 + 16, ] ) - A_s0 = T.int32() - A_s1 = T.int32() A_1 = T.match_buffer( data_im2col_reindex_shared_dyn_wmma_matrix_a[ v_x_o * 16 : v_x_o * 16 + 16, @@ -794,12 +803,10 @@ def complex_compute( ], (16, 16), "float16", - strides=(A_s0, A_s1), + strides=(A_s0_2, A_s1_2), scope="wmma.matrix_a", offset_factor=16, ) - B_s0 = T.int32() - B_s1 = T.int32() B = T.match_buffer( weight_flatten_reindex_shared_dyn_wmma_matrix_b[ v_y_o * 16 : v_y_o * 16 + 16, @@ -811,8 +818,6 @@ def complex_compute( scope="wmma.matrix_b", offset_factor=16, ) - C_s0 = T.int32() - C_s1 = T.int32() C = T.match_buffer( Conv_reindex_wmma_accumulator[ v_x_o * 16 : v_x_o * 16 + 16, @@ -820,23 +825,23 @@ def complex_compute( ], (16, 16), "float16", - strides=(C_s0, C_s1), + strides=(C_s0_3, C_s1_3), scope="wmma.accumulator", offset_factor=16, ) T.tvm_mma_sync( C.data, - C.ty.elem_offset // C_s0 // 16 * (C_s0 // 16) - + C.ty.elem_offset % C_s0 // 16, + C.ty.elem_offset // C_s0_3 // 16 * (C_s0_3 // 16) + + C.ty.elem_offset % C_s0_3 // 16, A_1.data, - A_1.ty.elem_offset // A_s0 // 16 * (A_s0 // 16) - + A_1.ty.elem_offset % A_s0 // 16, + A_1.ty.elem_offset // A_s0_2 // 16 * (A_s0_2 // 16) + + A_1.ty.elem_offset % A_s0_2 // 16, B.data, B.ty.elem_offset // B_s0 // 16 * (B_s0 // 16) + B.ty.elem_offset % B_s0 // 16, C.data, - C.ty.elem_offset // C_s0 // 16 * (C_s0 // 16) - + C.ty.elem_offset % C_s0 // 16, + C.ty.elem_offset // C_s0_3 // 16 * (C_s0_3 // 16) + + C.ty.elem_offset % C_s0_3 // 16, ) for ax0_0, ax1_0 in T.grid(2, 2): with Ts.sblock("Conv_reindex_wmma.accumulator_o"): @@ -850,25 +855,21 @@ def complex_compute( Ts.writes( Conv[v0_o * 16 : v0_o * 16 + 16, v1_o * 16 : v1_o * 16 + 16] ) - A_s0 = T.int32() - A_s1 = T.int32() A_1 = T.match_buffer( Conv_reindex_wmma_accumulator[ v0_o * 16 : v0_o * 16 + 16, v1_o * 16 : v1_o * 16 + 16 ], (16, 16), "float16", - strides=(A_s0, A_s1), + strides=(A_s0_3, A_s1_3), scope="wmma.accumulator", offset_factor=16, ) - C_s0 = T.int32() - C_s1 = T.int32() C = T.match_buffer( Conv[v0_o * 16 : v0_o * 16 + 16, v1_o * 16 : v1_o * 16 + 16], (16, 16), "float16", - strides=(C_s0, C_s1), + strides=(C_s0_4, C_s1_4), offset_factor=16, ) T.tvm_store_matrix_sync( @@ -876,16 +877,16 @@ def complex_compute( 16, 16, 16, - A_1.ty.elem_offset // A_s0 // 16 * (A_s0 // 16) - + A_1.ty.elem_offset % A_s0 // 16, + A_1.ty.elem_offset // A_s0_3 // 16 * (A_s0_3 // 16) + + A_1.ty.elem_offset % A_s0_3 // 16, T.tvm_access_ptr( T.type_annotation("float16"), C.data, C.ty.elem_offset, - C_s0 * 16, + C_s0_4 * 16, 2, ), - C_s0, + C_s0_4, "row_major", ) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_inject_software_pipeline.py b/tests/python/s_tir/transform/test_s_tir_transform_inject_software_pipeline.py index b796641af16b..91a3e40d365f 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_inject_software_pipeline.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_inject_software_pipeline.py @@ -157,9 +157,11 @@ def transformed_simple_compute( C[tx, 15] = B[1, tx, 0] + T.float32(1) +k = T.dynamic("k", "int32") + + @Ts.prim_func def dynamic_compute(a_handle: T.handle, c_handle: T.handle): - k = T.int32() A = T.match_buffer(a_handle, (16, k), "float32") C = T.match_buffer(c_handle, (16, k), "float32") for tx in T.thread_binding(0, 16, thread="threadIdx.x"): @@ -185,9 +187,11 @@ def dynamic_compute(a_handle: T.handle, c_handle: T.handle): C[tx, i] = B[tx, 0] + T.float32(1) +k = T.dynamic("k", "int32") + + @Ts.prim_func def transformed_dynamic_compute(a_handle: T.handle, c_handle: T.handle): - k = T.int32() A = T.match_buffer(a_handle, (16, k), "float32") C = T.match_buffer(c_handle, (16, k), "float32") for tx in T.thread_binding(0, 16, thread="threadIdx.x"): @@ -1719,9 +1723,10 @@ def after(A: T.Buffer((2,), "float32"), E: T.Buffer((2,), "float32")): def test_less_loop_than_num_stage_dynamic(): + K = T.dynamic("K", "int32") + @Ts.prim_func def before(a: T.handle, b: T.handle): - K = T.int32() A = T.match_buffer(a, [K], "float32") E = T.match_buffer(b, [K], "float32") for i in T.serial( @@ -1745,9 +1750,10 @@ def before(a: T.handle, b: T.handle): with Ts.sblock(): E[i] = D[0] + T.float32(5) + K = T.dynamic("K", "int32") + @Ts.prim_func def after(a: T.handle, b: T.handle): - K = T.int32() A = T.match_buffer(a, [K], "float32") E = T.match_buffer(b, [K], "float32") with Ts.sblock("root"): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lift_thread_binding.py b/tests/python/s_tir/transform/test_s_tir_transform_lift_thread_binding.py index 2fde8fe4eee2..ab766b136a61 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lift_thread_binding.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lift_thread_binding.py @@ -23,9 +23,10 @@ def test_lift_tx_beyond_local(): # fmt: off + n = T.dynamic("n", "int32") + @Ts.prim_func def before(a: T.handle, b: T.handle, c: T.handle): - n = T.int32() A = T.match_buffer(a, (32, 1, 128)) B = T.match_buffer(b, (32, n, 128)) C = T.match_buffer(c, (32, 1, n)) @@ -78,9 +79,10 @@ def before(a: T.handle, b: T.handle, c: T.handle): Ts.writes(C[ax0_ax1_fused // n, 0, ax0_ax1_fused % n]) C[ax0_ax1_fused // n, 0, ax0_ax1_fused % n] = D_local[ax0_ax1_fused // n, 0, ax0_ax1_fused % n] * T.float32(0.088397790055248615) + n = T.dynamic("n", "int32") + @Ts.prim_func def expected(A: T.Buffer((32, 1, 128), "float32"), b: T.handle, c: T.handle): - n = T.int32() B = T.match_buffer(b, (32, n, 128)) C = T.match_buffer(c, (32, 1, n)) # with Ts.sblock("root"): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py b/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py index b82e5e5a66ad..6ae64430320d 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lower_cross_thread_reduction.py @@ -737,6 +737,9 @@ def spatial_reduction_loop_predicate(A: T.Buffer((2, 32), "float32"), B: T.Buffe B[vi] = B[vi] + A[vi, vk] +k_0 = T.dynamic("k_0", "int32") + + @Ts.prim_func def lowered_reduction_spatial_loop_predicate( A: T.Buffer((2, 32), "float32"), B: T.Buffer((2,), "float32") @@ -769,7 +772,6 @@ def lowered_reduction_spatial_loop_predicate( T.tvm_thread_allreduce( T.uint32(1), in_thread_B[0], T.bool(True), cross_thread_B[0], k_1 ) - k_0 = T.int32() with Ts.sblock("block_write_back"): vi = Ts.axis.spatial(2, i_0 * 16 + i_1) Ts.where(i_0 * 16 + i_1 < 2 and k_1 == 0) @@ -1627,9 +1629,10 @@ def lowered_thread_broadcast_1(A: T.Buffer((256, 256), "float32"), B: T.Buffer(( # fmt: off +n = T.dynamic("n") + @Ts.prim_func def thread_broadcast_2(lv1605: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16"), p_lv1606: T.handle, p_lv1582: T.handle, p_output0: T.handle): - n = T.int64() lv1606 = T.match_buffer(p_lv1606, (T.int64(1), T.int64(32), n, T.int64(128)), "float16") lv1582 = T.match_buffer(p_lv1582, (T.int64(1), T.int64(1), T.int64(1), n), "float16") var_compute_intermediate = T.match_buffer(p_output0, (T.int64(1), T.int64(32), T.int64(1), n)) @@ -1675,9 +1678,10 @@ def thread_broadcast_2(lv1605: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T. var_compute_intermediate[T.int64(0), v0, T.int64(0), v1] = T.Cast("float32", T.min(T.max(var_NT_matmul_intermediate_local[T.int64(0), v0, T.int64(0), v1] * T.float16(0.088397790055248615), T.float16(-65504)), lv1582[T.int64(0), T.int64(0), T.int64(0), v1])) +n = T.dynamic("n") + @Ts.prim_func def lowered_thread_broadcast_2(lv1605: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16"), p_lv1606: T.handle, p_lv1582: T.handle, p_output0: T.handle): - n = T.int64() lv1606 = T.match_buffer(p_lv1606, (T.int64(1), T.int64(32), n, T.int64(128)), "float16") lv1582 = T.match_buffer(p_lv1582, (T.int64(1), T.int64(1), T.int64(1), n), "float16") var_compute_intermediate = T.match_buffer(p_output0, (T.int64(1), T.int64(32), T.int64(1), n)) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lower_match_buffer.py b/tests/python/s_tir/transform/test_s_tir_transform_lower_match_buffer.py index c1239437a625..3ae389ceddd0 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lower_match_buffer.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lower_match_buffer.py @@ -73,6 +73,10 @@ def intrin_test(data, elem_offset, stride_0, stride_1, shape_0, shape_1): return 0 +Bs_0 = T.dynamic("Bs_0", "int32") +Bs_1 = T.dynamic("Bs_1", "int32") + + @Ts.prim_func def opaque_access(a: T.handle, b: T.handle) -> None: A = T.match_buffer(a, (32, 64, 128)) @@ -99,8 +103,6 @@ def opaque_access(a: T.handle, b: T.handle) -> None: ) for i, j, k in T.grid(64, 2, 8): with Ts.sblock(): - Bs_0 = T.int32() - Bs_1 = T.int32() Ts.reads([]) Ts.writes(B[i, j * 32 : j * 32 + 32, k * 8 : k * 8 + 8]) sub_B = T.match_buffer( @@ -174,13 +176,15 @@ def transformed_opaque_buffer_data_projection(a: T.handle) -> None: T.evaluate(T.call_extern("consume", A.data, 4, dtype="int32")) +As_0 = T.dynamic("As_0", "int32") +As_1 = T.dynamic("As_1", "int32") + + @Ts.prim_func def high_dim_opaque_access(a: T.handle) -> None: A = T.match_buffer(a, (16, 32, 64)) for i, j, k in T.grid(16, 2, 4): with Ts.sblock(): - As_0 = T.int32() - As_1 = T.int32() Ts.reads([]) Ts.writes(A[i, j * 16 : j * 16 + 16, k * 16 : k * 16 + 16]) sub_A = T.match_buffer( @@ -220,13 +224,15 @@ def transformed_high_dim_opaque_access(a: T.handle) -> None: ) +As_0 = T.dynamic("As_0", "int32") +As_1 = T.dynamic("As_1", "int32") + + @Ts.prim_func def high_dim_opaque_access_with_source_strides(a: T.handle) -> None: A = T.match_buffer(a, (16, 32, 64), strides=[2576, 80, 1]) for i, j, k in T.grid(16, 2, 4): with Ts.sblock(): - As_0 = T.int32() - As_1 = T.int32() Ts.reads([]) Ts.writes(A[i, j * 16 : j * 16 + 16, k * 16 : k * 16 + 16]) sub_A = T.match_buffer( @@ -266,6 +272,12 @@ def transformed_high_dim_opaque_access_with_source_strides(a: T.handle) -> None: ) +As_0 = T.dynamic("As_0", "int32") +As_1 = T.dynamic("As_1", "int32") +Ass_0 = T.dynamic("Ass_0", "int32") +Ass_1 = T.dynamic("Ass_1", "int32") + + @Ts.prim_func def recursive_match(a: T.handle, b: T.handle) -> None: A = T.match_buffer(a, (64, 64, 64)) @@ -279,8 +291,6 @@ def recursive_match(a: T.handle, b: T.handle) -> None: B[i, j * 16 : j * 16 + 16, k * 16 : k * 16 + 16], ] ) - As_0 = T.int32() - As_1 = T.int32() sub_A = T.match_buffer( A[i, j * 16 : j * 16 + 16, k * 16 : k * 16 + 16], (16, 16), @@ -301,8 +311,6 @@ def recursive_match(a: T.handle, b: T.handle) -> None: sub_B[jj * 4 : jj * 4 + 4, kk * 4 : kk * 4 + 4], ] ) - Ass_0 = T.int32() - Ass_1 = T.int32() sub_sub_A = T.match_buffer( sub_A[jj * 4 : jj * 4 + 4, kk * 4 : kk * 4 + 4], (4, 4), @@ -372,6 +380,10 @@ def transformed_recursive_match(a: T.handle, b: T.handle) -> None: B[i, j * 16 + jj * 4 + jjj, k * 16 + kk * 4 + kkk] = 1 +Bs_0 = T.dynamic("Bs_0", "int32") +Bs_1 = T.dynamic("Bs_1", "int32") + + @Ts.prim_func def symbolic_match(a: T.handle, b: T.handle, n: T.int32, m: T.int32) -> None: A = T.match_buffer(a, (n * m, m)) @@ -380,8 +392,6 @@ def symbolic_match(a: T.handle, b: T.handle, n: T.int32, m: T.int32) -> None: with Ts.sblock(): Ts.reads([]) Ts.writes([A[i * m : i * m + n, 0:m], B[i * n : i * n + 2, 0 : m * 4]]) - Bs_0 = T.int32() - Bs_1 = T.int32() sub_A = T.match_buffer(A[i * m : i * m + m, 0:m], (m, m), offset_factor=1) sub_B = T.match_buffer( B[i * n : i * n + 2, 0 : m * 4], (2, m * 4), strides=[Bs_0, Bs_1], offset_factor=1 @@ -491,12 +501,14 @@ def fail_match_store(a: T.handle) -> None: # well-formed checker complains about redefinition of a stride variable +stride = T.dynamic("stride", "int32") + + @Ts.prim_func(check_well_formed=False) def fail_buffer_bind(a: T.handle) -> None: A = T.match_buffer(a, (8, 8)) for i, j in T.grid(8, 2): with Ts.sblock(): - stride = T.int32() sub_A = T.match_buffer( A[i, j * 4 : j * 4 + 4], (1, 4), strides=[stride, stride], offset_factor=1 ) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lower_opaque_block.py b/tests/python/s_tir/transform/test_s_tir_transform_lower_opaque_block.py index 1b748b252619..ee6ea3ed487f 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lower_opaque_block.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lower_opaque_block.py @@ -276,9 +276,11 @@ def transformed_strided_buffer_func( C[i0 * 4 + i1, j] = B[i1, j] * T.float32(2) +n = T.dynamic("n", "int32") + + @Ts.prim_func def compacted_symbolic_strided_buffer_func(a: T.handle) -> None: - n = T.int32() A = T.match_buffer(a, (1, n, 10240)) padded_size = T.meta_var(T.min((n + 63) // 64 * 64, 96)) # with Ts.sblock("root"): @@ -297,9 +299,11 @@ def compacted_symbolic_strided_buffer_func(a: T.handle) -> None: ) +n = T.dynamic("n", "int32") + + @Ts.prim_func def transformed_symbolic_strided_buffer_func(a: T.handle): - n = T.int32() A = T.match_buffer(a, (1, n, 10240)) padded_size = T.min((n + 63) // 64 * 64, 96) for i, j, k in T.grid(((n + 63) // 64 * 4 + 7) // 8, 2, 160): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_memhammer_lower_auto_copy.py b/tests/python/s_tir/transform/test_s_tir_transform_memhammer_lower_auto_copy.py index acffb958aed5..5478298c686f 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_memhammer_lower_auto_copy.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_memhammer_lower_auto_copy.py @@ -421,12 +421,14 @@ def main(a: T.handle, b: T.handle): B[bx * 128 + ax0, by * 128 + ax1] = A_shared_dyn[ax0, ax1] +s0 = T.dynamic("s0", "int32") +s1 = T.dynamic("s1", "int32") + + @tvm.script.ir_module class TransformedSharedToWmma: @Ts.prim_func def main() -> None: - s0 = T.int32() - s1 = T.int32() # body with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) @@ -502,12 +504,14 @@ def main() -> None: ) +s0 = T.dynamic("s0", "int32") +s1 = T.dynamic("s1", "int32") + + @tvm.script.ir_module class TransformedWmmaToShared: @Ts.prim_func def main() -> None: - s0 = T.int32() - s1 = T.int32() # body with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) @@ -583,6 +587,10 @@ def main() -> None: ) +s1 = T.dynamic("s1", "int32") +s0 = T.dynamic("s0", "int32") + + @tvm.script.ir_module class TransformedWmmaToGlobal: @Ts.prim_func @@ -622,8 +630,6 @@ def main(C: T.Buffer((1024, 1024), "float32")): scope="wmma.accumulator", offset_factor=16, ) - s1 = T.int32() - s0 = T.int32() tgt = T.match_buffer( C_accum_shared_dyn[ty, ax1_0, 0:16, 0:16], (16, 16), @@ -780,12 +786,16 @@ def main(C: T.Buffer((1024, 1024), "float32")): ] +s0_0 = T.dynamic("s0_0", "int32") +s1_0 = T.dynamic("s1_0", "int32") +s1_1 = T.dynamic("s1_1", "int32") +s0_1 = T.dynamic("s0_1", "int32") + + @tvm.script.ir_module class TransformedWmmaToGlobalWithFusion: @Ts.prim_func def main(A: T.Buffer((1024,), "float32"), C: T.Buffer((1024, 1024), "float32")) -> None: - s0 = T.int32() - s1 = T.int32() # body with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) @@ -824,12 +834,10 @@ def main(A: T.Buffer((1024,), "float32"), C: T.Buffer((1024, 1024), "float32")) scope="wmma.accumulator", offset_factor=16, ) - s1 = T.int32() - s0 = T.int32() tgt = T.match_buffer( C_accum_shared_dyn[ty, ax1_0, 0:16, 0:16], (16, 16), - strides=(s1, s0), + strides=(s1_1, s0_1), scope="shared.dyn", offset_factor=16, ) @@ -844,10 +852,10 @@ def main(A: T.Buffer((1024,), "float32"), C: T.Buffer((1024, 1024), "float32")) T.type_annotation("float32"), tgt.data, tgt.elem_offset, - s1 * 16, + s1_1 * 16, 2, ), - s1, + s1_1, "row_major", ) for ( @@ -1005,6 +1013,10 @@ def main(A: T.Buffer((1024,), "float32"), C: T.Buffer((1024, 1024), "float32")) ) +s1 = T.dynamic("s1", "int32") +s0 = T.dynamic("s0", "int32") + + @tvm.script.ir_module class TransformedMmaToGlobal: @Ts.prim_func @@ -1044,7 +1056,6 @@ def main(C: T.Buffer((1024, 1024), "float32")): scope="m16n8k8.matrixC", offset_factor=8, ) - s1, s0 = T.int32(), T.int32() tgt = T.match_buffer( C_accum_shared_dyn[ty, ax1_0, 0:8, 0:8], (8, 8), diff --git a/tests/python/te/test_te_create_primfunc.py b/tests/python/te/test_te_create_primfunc.py index daf21578ffb3..c88d0dbfc64b 100644 --- a/tests/python/te/test_te_create_primfunc.py +++ b/tests/python/te/test_te_create_primfunc.py @@ -203,11 +203,13 @@ def te_multi_output(): return [A0, A1, B0, B1] +m = T.dynamic("m", "int32") +n = T.dynamic("n", "int32") + + @Ts.prim_func def tir_multi_output(a0: T.handle, a1: T.handle, b0: T.handle, b1: T.handle) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - m = T.int32() - n = T.int32() A0 = T.match_buffer(a0, (m, n)) A1 = T.match_buffer(a1, (m, n)) B0 = T.match_buffer(b0, (m, n)) @@ -240,12 +242,14 @@ def te_extern(): return [A, B, C] +off1 = T.dynamic("off1", "int32") +off2 = T.dynamic("off2", "int32") +off3 = T.dynamic("off3", "int32") + + @Ts.prim_func def tir_extern(a: T.handle, b: T.handle, c: T.handle) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - off1 = T.int32() - off2 = T.int32() - off3 = T.int32() A = T.match_buffer(a, (128, 128), elem_offset=off1) B = T.match_buffer(b, (128, 128), elem_offset=off2) C = T.match_buffer(c, (128, 128), elem_offset=off3) @@ -553,13 +557,15 @@ def f_identity(dtype0: tvm.DataType, dtype1: tvm.DataType): return [idx, val, max_idx, max_val] +m = T.dynamic("m", "int32") +n = T.dynamic("n", "int32") + + @Ts.prim_func def tir_argmax_idx_val( var_idx: T.handle, var_val: T.handle, var_argmax_v0: T.handle, var_argmax_v1: T.handle ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - m = T.int32() - n = T.int32() idx = T.match_buffer(var_idx, [m, n], dtype="int32") val = T.match_buffer(var_val, [m, n], dtype="float32") argmax_v0 = T.match_buffer(var_argmax_v0, [m], dtype="int32") @@ -604,13 +610,15 @@ def f_identity(dtype0: tvm.DataType, dtype1: tvm.DataType): return [val, idx, max_val, max_idx] +m = T.dynamic("m", "int32") +n = T.dynamic("n", "int32") + + @Ts.prim_func def tir_argmax_val_idx( var_val: T.handle, var_idx: T.handle, var_argmax_v0: T.handle, var_argmax_v1: T.handle ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - m = T.int32() - n = T.int32() val = T.match_buffer(var_val, [m, n], dtype="float32") idx = T.match_buffer(var_idx, [m, n], dtype="int32") argmax_v0 = T.match_buffer(var_argmax_v0, [m], dtype="float32") @@ -727,14 +735,16 @@ def te_resize2d_symbolic(): return [A, B] +oh = T.dynamic("oh") +ow = T.dynamic("ow") + + @Ts.prim_func def tir_resize2d_symbolic( A: T.Buffer((T.int64(2), T.int64(3), T.int64(128), T.int64(128)), "float32"), var_resize: T.handle, ): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - oh = T.int64() - ow = T.int64() resize = T.match_buffer(var_resize, [T.int64(2), T.int64(3), oh, ow], dtype="float32") for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(3), oh, ow): with Ts.sblock("resize"): @@ -816,10 +826,13 @@ def te_slice_with_var_input(): return [tensor, idx, slice0] +m = T.dynamic("m") +n = T.dynamic("n") + + @Ts.prim_func def tir_slice_with_var_input(var_tensor: T.handle, idx: T.int64, var_slice: T.handle): T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) - m, n = T.int64(), T.int64() tensor = T.match_buffer(var_tensor, (m, n)) slice = T.match_buffer(var_slice, (idx, n)) # with Ts.sblock("root"): diff --git a/tests/python/tirx-analysis/test_tir_analysis_verify_well_formed.py b/tests/python/tirx-analysis/test_tir_analysis_verify_well_formed.py index d64d1cd4b598..1d7ee37667f8 100644 --- a/tests/python/tirx-analysis/test_tir_analysis_verify_well_formed.py +++ b/tests/python/tirx-analysis/test_tir_analysis_verify_well_formed.py @@ -74,9 +74,10 @@ def test_error_for_out_of_scope_usage(): def test_error_for_nested_rebind_usage(): """A variable may not be re-defined within the initial scope""" + i = T.dynamic("i", "int32") + @T.prim_func(check_well_formed=False) def func(): - i = T.int32() T.bind(42, var=i) T.bind(42, var=i) T.evaluate(i) @@ -96,9 +97,10 @@ def test_error_for_repeated_binding(): scope extends to all subsequent siblings). """ + i = T.dynamic("i", "int32") + @T.prim_func(check_well_formed=False) def func(): - i = T.int32() T.bind(42, var=i) T.evaluate(i) T.bind(17, var=i) @@ -113,7 +115,7 @@ def func(): def test_error_for_cross_function_reuse(): """A variable may not be re-defined in another function""" - i = tvm.tirx.Var("i", "int32") + i = T.dynamic("i", "int32") @I.ir_module(check_well_formed=False) class mod: @@ -181,7 +183,7 @@ def test_reuse_of_env_thread_across_functions_is_ill_formed(): PrimFuncs. """ - threadIdx_x = tvm.tirx.Var("threadIdx_x", "int32") + threadIdx_x = T.dynamic("threadIdx_x", "int32") @I.ir_module(check_well_formed=False) class mod: @@ -241,10 +243,10 @@ def test_error_message_without_previous_definition_location(): IS known, so the message includes location info. """ + x = T.dynamic("x", "int32") + @T.prim_func(check_well_formed=False) def func(): - x = T.int32() - T.bind(42, var=x) T.evaluate(x) @@ -268,10 +270,10 @@ def test_error_message_with_previous_definition_location(): contain 'It was first defined at' with the location information. """ + x = T.dynamic("x", "int32") + @T.prim_func(check_well_formed=False) def func(): - x = T.int32() - T.bind(42, var=x) T.bind(99, var=x) # This should trigger the error T.evaluate(x) @@ -297,10 +299,10 @@ def test_sequential_redefinition_with_location(): are treated as nested definitions with location info. """ + x = T.dynamic("x", "int32") + @T.prim_func(check_well_formed=False) def func(): - x = T.int32() - T.bind(1, var=x) T.evaluate(x) diff --git a/tests/python/tirx-base/test_tir_intrin.py b/tests/python/tirx-base/test_tir_intrin.py index e412c1ba56cb..59572d25acfa 100644 --- a/tests/python/tirx-base/test_tir_intrin.py +++ b/tests/python/tirx-base/test_tir_intrin.py @@ -308,17 +308,19 @@ def run_and_check(): tvm.testing.run_with_gpu_lock(run_and_check) +n = T.dynamic("n", "int32") +stride = T.dynamic("stride", "int32") +stride_1 = T.dynamic("stride_1", "int32") +stride_2 = T.dynamic("stride_2", "int32") +stride_3 = T.dynamic("stride_3", "int32") + + @tvm.script.ir_module class Module: @T.prim_func def test_tir_fma(A: T.handle, B: T.handle, C: T.handle, d: T.handle) -> None: # function attr dict T.func_attr({"global_symbol": "test_fma", "tirx.noalias": True}) - n = T.int32() - stride = T.int32() - stride_1 = T.int32() - stride_2 = T.int32() - stride_3 = T.int32() A_1 = T.match_buffer( A, [n], diff --git a/tests/python/tirx-base/test_tir_specialize.py b/tests/python/tirx-base/test_tir_specialize.py index 732473fcf6af..7c623d502622 100644 --- a/tests/python/tirx-base/test_tir_specialize.py +++ b/tests/python/tirx-base/test_tir_specialize.py @@ -22,6 +22,8 @@ import tvm from tvm.script import tirx as T +m = T.dynamic("m", "int32") + def assert_structural_equal_ignore_global_symbol(lhs, rhs): tvm.ir.assert_structural_equal( @@ -31,7 +33,6 @@ def assert_structural_equal_ignore_global_symbol(lhs, rhs): @T.prim_func def matmul(a: T.handle, b: T.handle, c: T.handle, n: T.int32) -> None: - m = T.int32() A = T.match_buffer(a, [m, n]) B = T.match_buffer(b, [m, n]) C = T.match_buffer(c, [m, m]) @@ -54,9 +55,11 @@ def matmul_128(a: T.handle, b: T.handle, c: T.handle) -> None: C[i, j] = C[i, j] + A[i, k] * B[j, k] +m = T.dynamic("m", "int32") + + @T.prim_func def matmul_m_128(a: T.handle, b: T.handle, c: T.handle) -> None: - m = T.int32() A = T.match_buffer(a, [m, 128]) B = T.match_buffer(b, [m, 128]) C = T.match_buffer(c, [m, m]) @@ -69,10 +72,12 @@ def matmul_m_128(a: T.handle, b: T.handle, c: T.handle) -> None: # x is considered undefined because it appears as part of x*8, # but not on its own +x = T.dynamic("x", "int32") +m = T.dynamic("m", "int32") + + @T.prim_func(check_well_formed=False) def matmul_m_8x(a: T.handle, b: T.handle, c: T.handle) -> None: - x = T.int32() - m = T.int32() A = T.match_buffer(a, [m, x * 8]) B = T.match_buffer(b, [m, x * 8]) C = T.match_buffer(c, [m, m]) @@ -83,10 +88,12 @@ def matmul_m_8x(a: T.handle, b: T.handle, c: T.handle) -> None: C[i, j] = C[i, j] + A[i, k] * B[j, k] +m = T.dynamic("m", "int32") +n = T.dynamic("n", "int32") + + @T.prim_func def element_wise(a: T.handle, c: T.handle) -> None: - m = T.int32() - n = T.int32() A = T.match_buffer(a, (m, n), "float32") C = T.match_buffer(c, (m, n), "float32") @@ -112,9 +119,11 @@ def element_wise_128_64(a: T.handle, c: T.handle) -> None: C[i, j] = B[i, j] + 1.0 +n = T.dynamic("n", "int32") + + @T.prim_func def element_wise_128_n(a: T.handle, c: T.handle) -> None: - n = T.int32() A = T.match_buffer(a, (128, n), "float32") C = T.match_buffer(c, (128, n), "float32") B = T.alloc_buffer((128, n), "float32") @@ -200,9 +209,10 @@ def test_specialize_recursive_load(): def test_specialize_with_const_folding(): + n = T.dynamic("n", "int32") + @T.prim_func def before(a: T.handle, b: T.handle): - n = T.int32() A = T.match_buffer(a, [n // 8, 8], "int32") B = T.match_buffer(b, [n], "int32") for i in range(n - 1): diff --git a/tests/python/tirx-transform/test_tir_transform_convert_ssa.py b/tests/python/tirx-transform/test_tir_transform_convert_ssa.py index d2d77580a184..9032fab0d8d7 100644 --- a/tests/python/tirx-transform/test_tir_transform_convert_ssa.py +++ b/tests/python/tirx-transform/test_tir_transform_convert_ssa.py @@ -314,7 +314,7 @@ def test_de_duplicate_thread_idx_across_multiple_functions(): Var/IterVar usage across the two PrimFuncs. """ - threadIdx_x = tvm.tirx.Var("threadIdx_x", "int32") + threadIdx_x = T.dynamic("threadIdx_x", "int32") # threadIdx_x is defined outside @I.ir_module(check_well_formed=False) @@ -337,27 +337,28 @@ def kernel_2(A: T.Buffer([256], "float32")): ) A[threadIdx_x] = A[threadIdx_x] + T.float32(1) + kernel_1_threadIdx_x = T.dynamic("threadIdx_x", "int32") + kernel_2_threadIdx_x = T.dynamic("threadIdx_x", "int32") + @I.ir_module class expected: @T.prim_func def kernel_1(A: T.Buffer([256], "float32")): - threadIdx_x = T.int32() T.attr( - T.iter_var(threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), + T.iter_var(kernel_1_threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", 256, ) - A[threadIdx_x] = A[threadIdx_x] + T.float32(1) + A[kernel_1_threadIdx_x] = A[kernel_1_threadIdx_x] + T.float32(1) @T.prim_func def kernel_2(A: T.Buffer([256], "float32")): - threadIdx_x = T.int32() T.attr( - T.iter_var(threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), + T.iter_var(kernel_2_threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", 256, ) - A[threadIdx_x] = A[threadIdx_x] + T.float32(1) + A[kernel_2_threadIdx_x] = A[kernel_2_threadIdx_x] + T.float32(1) after = tvm.tirx.transform.ConvertSSA()(before) tvm.ir.assert_structural_equal(after, expected) @@ -371,7 +372,7 @@ def test_de_duplicate_thread_idx_iter_var_across_multiple_functions(): PrimFuncs, not just the `tirx.Var` inside the `IterVar`. """ - threadIdx_x = tvm.tirx.Var("threadIdx_x", "int32") + threadIdx_x = T.dynamic("threadIdx_x", "int32") iter_var = tvm.tirx.IterVar( tvm.ir.Range(0, 256), threadIdx_x, tvm.tirx.IterVar.ThreadIndex, "threadIdx.x" ) @@ -389,27 +390,28 @@ def kernel_2(A: T.Buffer([256], "float32")): T.attr(iter_var, "thread_extent", 256) A[threadIdx_x] = A[threadIdx_x] + T.float32(1) + kernel_1_threadIdx_x = T.dynamic("threadIdx_x", "int32") + kernel_2_threadIdx_x = T.dynamic("threadIdx_x", "int32") + @I.ir_module(check_well_formed=False) class expected: @T.prim_func def kernel_1(A: T.Buffer([256], "float32")): - threadIdx_x = T.int32() T.attr( - T.iter_var(threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), + T.iter_var(kernel_1_threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", 256, ) - A[threadIdx_x] = A[threadIdx_x] + T.float32(1) + A[kernel_1_threadIdx_x] = A[kernel_1_threadIdx_x] + T.float32(1) @T.prim_func def kernel_2(A: T.Buffer([256], "float32")): - threadIdx_x = T.int32() T.attr( - T.iter_var(threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), + T.iter_var(kernel_2_threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", 256, ) - A[threadIdx_x] = A[threadIdx_x] + T.float32(1) + A[kernel_2_threadIdx_x] = A[kernel_2_threadIdx_x] + T.float32(1) after = tvm.tirx.transform.ConvertSSA()(before) tvm.ir.assert_structural_equal(after, expected) @@ -425,7 +427,7 @@ def test_thread_idx_reused_within_and_across_functions(): de-duplicated. """ - threadIdx_x = tvm.tirx.Var("threadIdx_x", "int32") + threadIdx_x = T.dynamic("threadIdx_x", "int32") iter_var = tvm.tirx.IterVar( tvm.ir.Range(0, 256), threadIdx_x, tvm.tirx.IterVar.ThreadIndex, "threadIdx.x" ) diff --git a/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py b/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py index 865fcd80e466..ddd6d53d250f 100644 --- a/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py +++ b/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py @@ -478,21 +478,23 @@ def main( def test_let_binding(): + n = T.dynamic("n") + @tvm.script.ir_module class Before: @T.prim_func def main(buf: T.handle): - n = T.int64() Buf = T.match_buffer(buf, [n], "int32") ceil_log2: T.int64 = T.Cast("int64", T.ceil(T.log2(T.Cast("float32", n)))) for i in T.serial(ceil_log2): T.evaluate(0) + n = T.dynamic("n", "int32") + @tvm.script.ir_module class Expected: @T.prim_func def main(buf: T.handle): - n = T.int32() Buf = T.match_buffer(buf, [n], "int32") # The pass narrows indexing variables (n, the For extent) but leaves # an explicitly-typed `T.Cast("int64", ...)` storage alone; a Cast to diff --git a/tests/python/tirx-transform/test_tir_transform_make_packed_api.py b/tests/python/tirx-transform/test_tir_transform_make_packed_api.py index 3531579950ab..5565b5b4184d 100644 --- a/tests/python/tirx-transform/test_tir_transform_make_packed_api.py +++ b/tests/python/tirx-transform/test_tir_transform_make_packed_api.py @@ -447,12 +447,13 @@ def test_forward_reference_symbolic_variable(): ensures all variable definitions precede all assertions. """ + batch_size = T.dynamic("batch_size") + @I.ir_module class Before: @T.prim_func def main(a: T.handle, b: T.handle): T.func_attr({"target": T.target("llvm", host="llvm")}) - batch_size = T.int64() A = T.match_buffer(a, (batch_size + 1,), "int32") B = T.match_buffer(b, (batch_size,), "int32") for i in range(batch_size): diff --git a/tests/python/tirx-transform/test_tir_transform_simplify.py b/tests/python/tirx-transform/test_tir_transform_simplify.py index 8452ec9f756d..e6fcee5e9c5a 100644 --- a/tests/python/tirx-transform/test_tir_transform_simplify.py +++ b/tests/python/tirx-transform/test_tir_transform_simplify.py @@ -1257,19 +1257,21 @@ def before(A_ptr: T.handle("float32"), B_ptr: T.handle("float32"), n: T.int32): def test_buffer_shape_constraint(): + n = T.dynamic("n") + @I.ir_module(check_well_formed=False) class Before: @T.prim_func def main(a: T.handle): - n = T.int64() A = T.match_buffer(a, (n * 32,), "float32") A[T.min(T.int64(0), n)] = T.float32(0) + n = T.dynamic("n") + @I.ir_module(check_well_formed=False) class Expected: @T.prim_func def main(a: T.handle): - n = T.int64() A = T.match_buffer(a, (n * 32,), "float32") A[T.int64(0)] = T.float32(0) @@ -1278,19 +1280,21 @@ def main(a: T.handle): def test_buffer_shape_constraint_with_offset(): + n = T.dynamic("n") + @I.ir_module(check_well_formed=False) class Before: @T.prim_func def main(a: T.handle): - n = T.int64() A = T.match_buffer(a, (n * 32 + 1 - 2,), "float32") A[T.min(T.int64(1), n)] = T.float32(0) + n = T.dynamic("n") + @I.ir_module(check_well_formed=False) class Expected: @T.prim_func def main(a: T.handle): - n = T.int64() A = T.match_buffer(a, (n * 32 + 1 - 2,), "float32") A[T.int64(1)] = T.float32(0) diff --git a/tests/python/tirx-transform/test_tir_transform_split_host_device.py b/tests/python/tirx-transform/test_tir_transform_split_host_device.py index 7e141d57d325..06d002ad0d96 100644 --- a/tests/python/tirx-transform/test_tir_transform_split_host_device.py +++ b/tests/python/tirx-transform/test_tir_transform_split_host_device.py @@ -327,12 +327,13 @@ def default_function_kernel( def test_symbolic_var_parameter(): + m = T.dynamic("m") + @I.ir_module class Module: @T.prim_func def main(var_A: T.handle, var_B: T.handle): T.func_attr({"target": T.target("cuda")}) - m = T.int64() A = T.match_buffer(var_A, (m,)) B = T.match_buffer(var_B, (m,)) T.attr(T.target("cuda"), "target", 0) diff --git a/tests/python/tirx-transform/test_tir_transform_vectorize.py b/tests/python/tirx-transform/test_tir_transform_vectorize.py index feba13559f00..200c58ee06b7 100644 --- a/tests/python/tirx-transform/test_tir_transform_vectorize.py +++ b/tests/python/tirx-transform/test_tir_transform_vectorize.py @@ -498,11 +498,12 @@ def main(A: T.Buffer((25,), "int32"), B: T.Buffer((25,), "float32")): def test_illegal_extent(): + n = T.dynamic("n", "int32") + @I.ir_module(check_well_formed=False) class Mod: @T.prim_func def main(A: T.Buffer((25,), "int32")): - n = T.int32() for j in T.vectorized(n): A[j] = 3 diff --git a/tests/python/tirx/test_parser_printer.py b/tests/python/tirx/test_parser_printer.py index 2054faede1dd..97b71c163151 100644 --- a/tests/python/tirx/test_parser_printer.py +++ b/tests/python/tirx/test_parser_printer.py @@ -988,8 +988,8 @@ def test_buffer_shape_repeated_var_prints_out_of_line(): buffer = tvm.tirx.decl_buffer((n + n,), name="A") func = tvm.tirx.PrimFunc([buffer], tvm.tirx.Evaluate(0)) - code = func.script() - assert "n = T.int32()" in code + code = func.script(extra_config={"script.use_pep695": False}) + assert 'n = I.dynamic("n", dtype="int32")' in code assert_structural_equal(func, from_source(code)) diff --git a/tests/python/tirx/transform/test_transform_lower_tirx.py b/tests/python/tirx/transform/test_transform_lower_tirx.py index 0b05e215cbfa..946759488932 100644 --- a/tests/python/tirx/transform/test_transform_lower_tirx.py +++ b/tests/python/tirx/transform/test_transform_lower_tirx.py @@ -659,6 +659,9 @@ def before(A: T.Buffer((128, 32), "float16")) -> None: for vec in T.vectorized(8): A_smem[row, col + vec] = A[bx * 128 + row, col + vec] + compose_m = T.dynamic("compose_m", "int32") + compose_q = T.dynamic("compose_q", "int32") + @T.prim_func(private=True) def after(A_handle: T.handle) -> None: A = T.match_buffer(A_handle, (128, 32), "float16", layout=None) @@ -687,8 +690,6 @@ def after(A_handle: T.handle) -> None: # folded closed form: compose_m is the flat element index, so # compose_m // 8 is the row and compose_m % 8 the lane, which # substituted back gives the same address. - compose_m = T.int32() - compose_q = T.int32() A_smem[ T.Let( T.Let( From eeb8a9a9a379415e3146ebfed7ac76adb79b5002 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 21:23:37 +0000 Subject: [PATCH 04/16] [Script] Align symbol registration with consolidated builder APIs --- docs/reference/api/python/script/parser.rst | 2 +- python/tvm/tirx/script/ir_builder/__init__.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/reference/api/python/script/parser.rst b/docs/reference/api/python/script/parser.rst index 02e5304d606c..2b2f7c4b26be 100644 --- a/docs/reference/api/python/script/parser.rst +++ b/docs/reference/api/python/script/parser.rst @@ -32,7 +32,7 @@ registered through :func:`tvm.script.parser.register_namespace` and :func:`tvm.script.parser.register_namespace_initializer`. .. automodule:: tvm.script.parser.protocol_registry - :members: constexpr, args_policy, register_type_var_decl, mutable_cell_decl, result_span, module_decorator, declaration_kind + :members: constexpr, args_policy, register_scalar_annotation, mutable_cell_decl, result_span, module_decorator, declaration_kind The language variant aliases below share the public construction namespaces documented in :doc:`script`. Parser entry points above use the canonical frontend. diff --git a/python/tvm/tirx/script/ir_builder/__init__.py b/python/tvm/tirx/script/ir_builder/__init__.py index cec51e5d2e33..8eb021065726 100644 --- a/python/tvm/tirx/script/ir_builder/__init__.py +++ b/python/tvm/tirx/script/ir_builder/__init__.py @@ -23,10 +23,10 @@ from tvm import ir as _ir from tvm import tirx as _tir +from tvm.script.ir_builder import dynamic as dynamic from tvm.script.ir_builder.base import annotation_constructor as _annotation_constructor from tvm.script.ir_builder.base import at as _at from tvm.script.ir_builder.base import source_span as _source_span -from tvm.script.ir_builder import dynamic as dynamic from tvm.script.parser.protocol_registry import constexpr as constexpr from tvm.script.parser.protocol_registry import ( mutable_cell_decl as _mutable_cell_decl, From 2ad856e72f243173b93475092b10a3b0866753c2 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 22:17:05 +0000 Subject: [PATCH 05/16] [Script] Use explicit symbols in runtime library generators --- jvm/core/src/test/scripts/prepare_test_libs.py | 4 +++- python/tvm/relax/block_builder.py | 2 +- tests/nightly/python/test_nnapi/test_ops.py | 3 --- web/tests/python/prepare_test_libs.py | 4 +++- 4 files changed, 7 insertions(+), 6 deletions(-) diff --git a/jvm/core/src/test/scripts/prepare_test_libs.py b/jvm/core/src/test/scripts/prepare_test_libs.py index 591715137260..202efe4e6baf 100644 --- a/jvm/core/src/test/scripts/prepare_test_libs.py +++ b/jvm/core/src/test/scripts/prepare_test_libs.py @@ -21,16 +21,18 @@ import tvm from tvm import relax, te +from tvm.script import ir as I from tvm.script import relax as R def prepare_relax_lib(base_path): pipeline = relax.get_pipeline() + n = I.dynamic("n") @tvm.script.ir_module class Mod: @R.function - def main(x: R.Tensor(["n"], "float32"), y: R.Tensor(["n"], "float32")): + def main(x: R.Tensor([n], "float32"), y: R.Tensor([n], "float32")): lv0 = R.add(x, y) return lv0 diff --git a/python/tvm/relax/block_builder.py b/python/tvm/relax/block_builder.py index 07e055408348..e62d71568afe 100644 --- a/python/tvm/relax/block_builder.py +++ b/python/tvm/relax/block_builder.py @@ -646,7 +646,7 @@ def emit_func_output( # `bb.function()`, then any variables provided from the params # are not in scope. Otherwise, TIR variables used in dynamic # inputs are removed as undefined (e.g. Replacing - # `R.Tensor(["batch_size"])` with `R.Tensor(ndims=1)`). + # `R.Tensor([batch_size])` with `R.Tensor(ndims=1)`). self.begin_scope(self._func._params) try: seqe = self.normalize(seqe) diff --git a/tests/nightly/python/test_nnapi/test_ops.py b/tests/nightly/python/test_nnapi/test_ops.py index a4a866e4677e..6181add25bb2 100644 --- a/tests/nightly/python/test_nnapi/test_ops.py +++ b/tests/nightly/python/test_nnapi/test_ops.py @@ -14,7 +14,6 @@ # KIND, either express or implied. See the License for the # specific language governing permissions and limitations # under the License. -# ruff: noqa: F841 """NNAPI integration operator tests.""" import numpy as np @@ -23,7 +22,6 @@ import tvm import tvm.script import tvm.script.relax as R -import tvm.script.tirx as T from test_nnapi.conftest import remote from test_nnapi.infrastructure import build_and_run @@ -269,7 +267,6 @@ class Module: def main( i0: R.Tensor((1, 10, 15), "float32"), ) -> R.Tensor((1, 10, 1), "float32"): - n = T.int64() with R.dataflow(): t0: R.Tensor((1, 10, 1), "float32") = R.mean(i0, axis=[-1], keepdims=True) R.output(t0) diff --git a/web/tests/python/prepare_test_libs.py b/web/tests/python/prepare_test_libs.py index 239453317017..bbc6c1b18e0c 100644 --- a/web/tests/python/prepare_test_libs.py +++ b/web/tests/python/prepare_test_libs.py @@ -21,16 +21,18 @@ import tvm from tvm import relax, te from tvm.contrib import tvmjs +from tvm.script import ir as I from tvm.script import relax as R def prepare_relax_lib(base_path): pipeline = relax.get_pipeline() + n = I.dynamic("n") @tvm.script.ir_module class Mod: @R.function - def main(x: R.Tensor(["n"], "float32"), y: R.Tensor(["n"], "float32")): + def main(x: R.Tensor([n], "float32"), y: R.Tensor([n], "float32")): lv0 = R.add(x, y) return lv0 From f6f0428bb9d5c892a89e081ee1cbd63a34514e09 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 23:04:03 +0000 Subject: [PATCH 06/16] [Script] Document explicit symbolic dimensions in Relax builders --- python/tvm/relax/script/ir_builder/__init__.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/python/tvm/relax/script/ir_builder/__init__.py b/python/tvm/relax/script/ir_builder/__init__.py index f296028e5624..bfcda8c88961 100644 --- a/python/tvm/relax/script/ir_builder/__init__.py +++ b/python/tvm/relax/script/ir_builder/__init__.py @@ -92,8 +92,8 @@ def Tensor(shape=None, dtype=None, vdevice=None, ndim=-1, *, span=None): ---------- shape : Expr or sequence of Expr, optional Tensor shape, or None when unknown. A string supplied without dtype - is shorthand for the dtype. Script dimension strings follow the - registered shape-expression argument policy. + is shorthand for the dtype. Symbolic dimensions are expressions over + explicit variables, such as those created with I.dynamic. dtype : str or PrimType, optional Element type; None leaves the element type unknown. vdevice : VDevice or str, optional @@ -126,8 +126,8 @@ def DTensor(shape=None, dtype=None, device_mesh=None, placement="", *, ndim=-1, Parameters ---------- shape : Expr or sequence of Expr, optional - Global tensor shape, or None when unknown; dimension strings in script - follow the registered shape-expression argument policy. + Global tensor shape, or None when unknown. Symbolic dimensions are + expressions over explicit variables, such as those created with I.dynamic. dtype : str or PrimType, optional Element type; None leaves the element type unknown. device_mesh : DeviceMesh or str, optional @@ -172,8 +172,8 @@ def Shape(values=None, ndim=-1, *, span=None): Parameters ---------- values : sequence of Expr, optional - Known dimensions, or None for an unknown shape value. Script dimension - strings use the registered int64 shape-expression policy. + Known dimensions, or None for an unknown shape value. Symbolic dimensions + are expressions over explicit variables, such as those created with I.dynamic. ndim : int, optional Number of dimensions when values is None; -1 leaves it unknown. Do not supply an explicit count together with known values. From ac9fb4b67fe6e11ee95c88be1c88fc0c570fe7ec Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 22:22:21 +0000 Subject: [PATCH 07/16] [REFACTOR][TVMScript] Read explicit annotation symbols directly --- .../tvm/relax/script/ir_builder/__init__.py | 15 +- .../script/ir_builder/parser_protocol.py | 2 +- python/tvm/script/ir_builder/__init__.py | 4 - python/tvm/script/ir_builder/base.py | 148 ++--------- python/tvm/script/parser/__init__.py | 3 +- python/tvm/script/parser/entry.py | 1 - python/tvm/script/parser/inspect_source.py | 20 ++ python/tvm/script/parser/transpile.py | 238 ++++++++---------- python/tvm/tirx/script/ir_builder/__init__.py | 4 +- .../tirx/script/ir_builder/parser_protocol.py | 1 - tests/python/script/minilang.py | 7 - 11 files changed, 144 insertions(+), 299 deletions(-) diff --git a/python/tvm/relax/script/ir_builder/__init__.py b/python/tvm/relax/script/ir_builder/__init__.py index bfcda8c88961..a9919cf4d264 100644 --- a/python/tvm/relax/script/ir_builder/__init__.py +++ b/python/tvm/relax/script/ir_builder/__init__.py @@ -28,8 +28,6 @@ from tvm.relax.distributed import Placement as _Placement from tvm.relax.distributed import device_mesh as device_mesh from tvm.script.ir_builder import resolve_global_info_args as _resolve_global_info_args -from tvm.script.ir_builder import IRBuilder as _IRBuilder -from tvm.script.ir_builder.base import annotation_constructor as _annotation_constructor from tvm.script.ir_builder.base import at as _at from tvm.script.ir_builder.base import source_span as _source_span from tvm.script.parser.protocol_registry import constexpr as constexpr @@ -84,7 +82,6 @@ @_resolve_global_info_args("vdevice", resolver=resolve_global_info_) -@_annotation_constructor("shape") def Tensor(shape=None, dtype=None, vdevice=None, ndim=-1, *, span=None): """Construct a Relax tensor type. @@ -109,8 +106,8 @@ def Tensor(shape=None, dtype=None, vdevice=None, ndim=-1, *, span=None): Returns ------- - result : TensorType or Type - The tensor type, or a missing type for an unresolved eager shape annotation. + result : TensorType + The constructed tensor type. String selectors outside an active module always raise ValueError. """ if isinstance(shape, _python.str) and dtype is None: @@ -119,7 +116,6 @@ def Tensor(shape=None, dtype=None, vdevice=None, ndim=-1, *, span=None): @_resolve_global_info_args("device_mesh", resolver=resolve_global_info_) -@_annotation_constructor("shape") def DTensor(shape=None, dtype=None, device_mesh=None, placement="", *, ndim=-1, span=None): """Construct a Relax distributed tensor type. @@ -144,8 +140,8 @@ def DTensor(shape=None, dtype=None, device_mesh=None, placement="", *, ndim=-1, Returns ------- - result : DTensorType or Type - The distributed type, or a missing type for an unresolved eager shape annotation. + result : DTensorType + The constructed distributed type. String selectors outside an active module always raise ValueError. """ if device_mesh is None: @@ -155,7 +151,7 @@ def DTensor(shape=None, dtype=None, device_mesh=None, placement="", *, ndim=-1, return _DTensorType(Tensor(shape, dtype, ndim=ndim), device_mesh, placement, _source_span(span)) -# The distributed source spelling shares concrete constructors and argument policy. +# The distributed source spelling shares the decorated concrete constructor. dist.DTensor = DTensor dist.device_mesh = device_mesh @@ -165,7 +161,6 @@ def DTensor(shape=None, dtype=None, device_mesh=None, placement="", *, ndim=-1, __tvm_value_if__ = True -@_annotation_constructor("values") def Shape(values=None, ndim=-1, *, span=None): """Construct a Relax shape type. diff --git a/python/tvm/relax/script/ir_builder/parser_protocol.py b/python/tvm/relax/script/ir_builder/parser_protocol.py index f7ba61e7c769..d3bcedd936bb 100644 --- a/python/tvm/relax/script/ir_builder/parser_protocol.py +++ b/python/tvm/relax/script/ir_builder/parser_protocol.py @@ -374,7 +374,7 @@ def func_name(name: str) -> None: def func_ret_type(annotation: Any, *, span: _Span = None) -> None: """Implements :func:`tvm.script.ir_builder.parser_protocol.func_ret_type`.""" - return _native.func_ret_type(_builder._type(_base._return_annotation(annotation))) + return _native.func_ret_type(_builder._type(annotation)) def check_well_formed_(function: _relax.Function) -> None: diff --git a/python/tvm/script/ir_builder/__init__.py b/python/tvm/script/ir_builder/__init__.py index 96ab99308f26..34c272207e8b 100644 --- a/python/tvm/script/ir_builder/__init__.py +++ b/python/tvm/script/ir_builder/__init__.py @@ -26,9 +26,7 @@ MISSING, AlreadyEmitted, IRBuilder, - annotation_value_, at_, - require_defined, resolve_global_info_args, with_at_group_, ) @@ -57,7 +55,6 @@ "Range", "StringImm", "StringType", - "annotation_value_", "at_", "check_well_formed_", "constexpr", @@ -72,7 +69,6 @@ "module_global_infos", "module_member_", "module_set_attr", - "require_defined", "resolve_global_info_args", "with_at_group_", ] diff --git a/python/tvm/script/ir_builder/base.py b/python/tvm/script/ir_builder/base.py index bcba3d2efacb..87d459f4ccd3 100644 --- a/python/tvm/script/ir_builder/base.py +++ b/python/tvm/script/ir_builder/base.py @@ -18,12 +18,11 @@ from collections.abc import Callable from contextlib import contextmanager, nullcontext -from functools import wraps from inspect import signature from typing import Any, Generic, TypeVar from tvm_ffi import register_object as _register_object -from tvm_ffi.dataclasses import MISSING +from tvm_ffi.dataclasses import MISSING as MISSING from tvm import ir from tvm.runtime import Object as _Object @@ -377,32 +376,6 @@ def at( return value -def require_defined(value, name): - """Report a source name whose designated region output was not produced. - - Parameters - ---------- - value : Any - Candidate binding. Only the canonical ``MISSING`` singleton denotes an - absent value; explicit None is a defined value. - name : str - Source identifier to include in the missing-name diagnostic. - - Returns - ------- - Any - The exact ``value`` when it is defined. - - Raises - ------ - NameError - If ``value`` is ``MISSING``. - """ - if value is MISSING: - raise NameError(f"name {name!r} is not defined") - return value - - def with_at_group_( location: SpanEntry | ir.Span | tuple[str | ir.SourceName, int, int, int, int] | None, thunk: Callable[[], _T], @@ -465,7 +438,7 @@ def _resolve_type_var(frame, ffi_resolver, name, dtype=None, *, value=None, span def _current_function_frame(): - """Find the nearest function for eager shared annotation constructors.""" + """Find the nearest native function frame for explicit symbol declarations.""" if IRBuilder.is_in_scope(): for frame in reversed(IRBuilder.current().frames): if callable(getattr(frame, "resolve_type_var", None)): @@ -473,108 +446,19 @@ def _current_function_frame(): raise ValueError("Symbol resolution requires an active function frame") -def annotation_constructor(*fields: str, as_type: bool = False): - """Adapt eager Python type parameters on concrete annotation constructors. +def annotation_constructor(constructor): + """Expose a constructor as a real annotation class supporting Python unions. - Unresolved ``typing.TypeVar`` values defer annotations outside a builder; - an active native function owns their resolution. Ordinary values, including - strings, are passed unchanged to the concrete API. + Calls construct ordinary native values directly. The class preserves the + constructor's signature and documentation without adapting its arguments. """ - - def decorate(constructor): - call_signature = signature(constructor) - - def unresolved(value): - if isinstance(value, TypeVar): - return True - if isinstance(value, tuple | list): - return any(unresolved(item) for item in value) - return False - - def resolve(value): - if isinstance(value, TypeVar): - if value.__bound__ is not None or value.__constraints__: - raise TypeError("A symbolic TypeVar cannot have constraints or a bound") - return _current_function_frame().resolve_type_var(value.__name__) - if isinstance(value, tuple): - return tuple(resolve(item) for item in value) - if isinstance(value, list): - return [resolve(item) for item in value] - return value - - @wraps(constructor) - def invoke(*args, **kwargs): - bound = call_signature.bind(*args, **kwargs) - for field in fields: - if field not in bound.arguments: - continue - value = bound.arguments[field] - if IRBuilder.is_in_scope(): - bound.arguments[field] = resolve(value) - elif unresolved(value): - return ir.Type.missing() - return constructor(*bound.args, **bound.kwargs) - - if not as_type: - return invoke - # A real annotation class supports Python unions while constructing - # ordinary native types, with no proxy values or parser policy state. - return type( - constructor.__name__, - (), - { - "__new__": lambda cls, *args, **kwargs: invoke(*args, **kwargs), - "__signature__": call_signature, - "__doc__": constructor.__doc__, - "__module__": constructor.__module__, - }, - ) - - return decorate - - -def _return_annotation(annotation): - """Evaluate a return annotation without introducing return-only symbols.""" - if not callable(annotation) or isinstance(annotation, ir.Expr | ir.Type): - return annotation - frame = _current_function_frame() - declared = set(frame.type_var_map) - annotation = annotation() - introduced = set(frame.type_var_map) - declared - if introduced: - raise ValueError(f"Return annotation introduces unbound symbol {sorted(introduced)[0]!r}") - return annotation - - -def annotation_value_(name, value): - """Adapt a real definition-context symbol using the native function map. - - Parameters - ---------- - name : str - Source spelling used to resolve the symbol in the nearest active - function frame. - value : TypeVar, Var or Any - Captured definition-context value. An unconstrained ``typing.TypeVar`` - resolves a symbolic variable; a primitive IR Var supplies its existing - value to the resolver. Other values pass through unchanged. - - Returns - ------- - Any - The function frame's resolved symbol, or the unchanged nonsymbolic value. - - Raises - ------ - TypeError - If a TypeVar has a bound or constraints. - ValueError - If a symbolic value requires resolution without an active function frame. - """ - if isinstance(value, TypeVar): - if value.__bound__ is not None or value.__constraints__: - raise TypeError("A symbolic TypeVar cannot have constraints or a bound") - return _current_function_frame().resolve_type_var(name) - if ir.is_prim_var(value): - return _current_function_frame().resolve_type_var(name, value=value) - return value + return type( + constructor.__name__, + (), + { + "__new__": lambda cls, *args, **kwargs: constructor(*args, **kwargs), + "__signature__": signature(constructor), + "__doc__": constructor.__doc__, + "__module__": constructor.__module__, + }, + ) diff --git a/python/tvm/script/parser/__init__.py b/python/tvm/script/parser/__init__.py index 46c85f168510..944d4a3d2ce0 100644 --- a/python/tvm/script/parser/__init__.py +++ b/python/tvm/script/parser/__init__.py @@ -20,7 +20,7 @@ import importlib from collections.abc import Callable -from typing import Any, TypeVar +from typing import Any _NAMESPACES: dict[str, object] = {} _NAMESPACE_INITIALIZERS: list[Callable[[], None]] = [] @@ -104,7 +104,6 @@ def _initialize() -> None: _initializing = True try: importlib.import_module(__name__ + ".ir") - register_namespace("TypeVar", TypeVar) for initialize in _NAMESPACE_INITIALIZERS: initialize() _initialized = True diff --git a/python/tvm/script/parser/entry.py b/python/tvm/script/parser/entry.py index cd883cf8e344..f3d45a4befb5 100644 --- a/python/tvm/script/parser/entry.py +++ b/python/tvm/script/parser/entry.py @@ -470,7 +470,6 @@ def _prepare_transpiler( No builder frame or expression is created here. """ namespace = { - "TypeVar": TypeVar, "tvm": sys.modules.get("tvm"), **_NAMESPACES, **environment, diff --git a/python/tvm/script/parser/inspect_source.py b/python/tvm/script/parser/inspect_source.py index 26b5603c9486..08de2ef0ce24 100644 --- a/python/tvm/script/parser/inspect_source.py +++ b/python/tvm/script/parser/inspect_source.py @@ -34,12 +34,32 @@ from types import CodeType, FrameType, FunctionType from typing import Any +from tvm_ffi.dataclasses import MISSING + from tvm.ir import SourceName, Span from .annotation import parse_annotation from .prescan import collect_annotation_free_names +class _AnnotationScope(dict): + """Snapshot selected bindings; absent names fail only when an annotation reads them.""" + + def __init__(self, names, *scopes): + super().__init__() + for name in names: + for scope in scopes: + if name in scope: + value = scope[name] + if value is not MISSING: + self[name] = value + # An explicit absence still shadows the remaining scopes. + break + + def __missing__(self, name): + raise NameError(f"name {name!r} is not defined") + + class Source: """Source code class for TVMScript. diff --git a/python/tvm/script/parser/transpile.py b/python/tvm/script/parser/transpile.py index 2e4f15b33103..400b7830a944 100644 --- a/python/tvm/script/parser/transpile.py +++ b/python/tvm/script/parser/transpile.py @@ -38,6 +38,7 @@ import ast import builtins +import copy from collections.abc import Callable, Iterator, Mapping from contextlib import contextmanager from types import FunctionType @@ -45,6 +46,7 @@ from . import protocol_registry as protocol from .annotation import parse_annotation +from .inspect_source import _AnnotationScope from .prescan import Binding, PrescanContext, resolve_namespace_key, resolve_namespace_value _Node = TypeVar("_Node", bound=ast.AST) @@ -130,10 +132,10 @@ def __init__(self, current_scope: ast.AST | None, dialect_prefix: str) -> None: self.current_scope = current_scope # Fixed generated namespace name selects this function's language variant operations. self.dialect_prefix = dialect_prefix - # Source-to-generated names start with definition captures and grow in signature + # Annotation reads start with definition captures and grow in signature # order. Annotation rewriting reads this one map; lexical masks temporarily # replace it and restore it on exit. Ordinary body lookup does not use it. - self.annotation_aliases: dict[str, str] = {} + self.annotation_bindings: dict[str, ast.expr] = {} class IRBuilderTranspiler(ast.NodeTransformer): @@ -287,16 +289,16 @@ def _bypass_rewrite(self) -> Iterator[None]: self.bypass_ast_rewrite = old @contextmanager - def _use_aliases(self, mapping: dict[str, str]) -> Iterator[None]: + def _use_aliases(self, mapping: dict[str, ast.expr]) -> Iterator[None]: """Restore lexical annotation substitutions even when a visitor fails.""" # Mask the active context's one map, retaining the enclosing map by identity. - old = self.function.annotation_aliases - self.function.annotation_aliases = mapping + old = self.function.annotation_bindings + self.function.annotation_bindings = mapping try: yield finally: # Nested annotation scopes cannot leak substitutions into their caller. - self.function.annotation_aliases = old + self.function.annotation_bindings = old @contextmanager def _rewrite_annotation(self) -> Iterator[None]: @@ -362,7 +364,7 @@ def visit_Name(self, node: ast.Name) -> ast.expr: # value: X.Tensor((n,)) # # Builder: - # value_type = X.Tensor((I.require_defined(I.annotation_value_("n", captured_n), "n"),)) + # value_type = X.Tensor((_definition["n"],)) # ------------------------------------------------- # Definition substitutions apply only within annotation syntax. # A preceding body target is already a Python binding (including symbols). @@ -378,28 +380,10 @@ def visit_Name(self, node: ast.Name) -> ast.expr: if ( self.annotation_expression and isinstance(node.ctx, ast.Load) - and node.id in self.function.annotation_aliases + and node.id in self.function.annotation_bindings ): - if node.id in self.module.prescan.namespaces: - return ast.copy_location( - ast.Name(self.function.annotation_aliases[node.id], ast.Load()), node - ) - return self._call( - self.module.infrastructure_name, - "require_defined", - [ - self._call( - self.module.infrastructure_name, - "annotation_value_", - [ - ast.Constant(node.id), - ast.Name(self.function.annotation_aliases[node.id], ast.Load()), - ], - node, - ), - ast.Constant(node.id), - ], - node, + return ast.copy_location( + copy.deepcopy(self.function.annotation_bindings[node.id]), node ) return node @@ -444,7 +428,7 @@ def visit_Lambda(self, node: ast.Lambda) -> ast.Lambda: with self._use_aliases( { name: alias - for name, alias in self.function.annotation_aliases.items() + for name, alias in self.function.annotation_bindings.items() if name not in local_names } ): @@ -462,12 +446,12 @@ def visit_ListComp( # [_S[i].ctx(lambda: f(n)) for n in values] # ------------------------------------------------- # Comprehension binders mask annotation substitutions in Python evaluation order. - with self._use_aliases(dict(self.function.annotation_aliases)): + with self._use_aliases(dict(self.function.annotation_bindings)): for generator in node.generators: generator.iter = self.visit(generator.iter) for target in ast.walk(generator.target): if isinstance(target, ast.Name): - self.function.annotation_aliases.pop(target.id, None) + self.function.annotation_bindings.pop(target.id, None) generator.ifs = [self.visit(value) for value in generator.ifs] if isinstance(node, ast.DictComp): node.key, node.value = self.visit(node.key), self.visit(node.value) @@ -1598,11 +1582,15 @@ def _function_namespace(self, node: ast.FunctionDef, namespace: object) -> str: return self._inject(namespace, "_X").id def _read_function_annotations( - self, node: ast.FunctionDef, parameters: list[ast.arg], facts: list[Binding] - ) -> tuple[list[ast.expr | None], ast.expr | None, dict[str, str]]: + self, + node: ast.FunctionDef, + parameters: list[ast.arg], + facts: list[Binding], + *, + captures: str, + ) -> tuple[list[ast.expr | None], ast.expr | None, dict[str, ast.expr]]: """Find definition-scope names needed by signatures and body annotations.""" declared_names = {item.name for item in getattr(node, "type_params", ())} - # Quoted expression names are created later by argument normalization. annotations = [ parse_annotation(parameter.annotation, self.module.filename) if parameter.annotation @@ -1630,42 +1618,31 @@ def _read_function_annotations( if isinstance(item, ast.Name) and isinstance(item.ctx, ast.Load) } annotation_names.update(body_annotation_names) - # Unconflicted captures keep their spelling in the declaration scope. - # A different execution binding or source-local binding needs a distinct - # compiler name, including preceding signature parameters. - conflicts = {item.name for item in facts} - # Body annotations execute alongside ordinary body globals. A definition- - # only name needs an alias there to keep those two lookup meanings separate. - conflicts.update(body_annotation_names - self.module.bindings.keys()) - if self.module.module_name is not None: - conflicts.update(self.module.module_functions) - conflicts.add(self.module.module_name) - aliases = { + # Fixed namespaces cannot be rebound. Other definition values are + # captured at execution, including names assigned by source prefix code. + bindings = { name: ( - name - if name not in conflicts - and ( - name not in self.module.bindings - or self.module.environment.get(name) is self.module.bindings[name] - ) - else self.module.fresh("_annotation") + ast.Name(name, ast.Load()) + if name in self.module.prescan.namespaces + and name in self.module.bindings + and self.module.environment.get(name) is self.module.bindings[name] + else ast.Subscript(ast.Name(captures, ast.Load()), ast.Constant(name), ast.Load()) ) for name in sorted(annotation_names - declared_names) } - return annotations, returns, aliases + return annotations, returns, bindings def _create_definition_bindings( - self, node: ast.FunctionDef, aliases: dict[str, str], *, captures: str, local_function: bool + self, + node: ast.FunctionDef, + bindings: dict[str, ast.expr], + *, + captures: str, + local_function: bool, ) -> list[ast.stmt]: - """Capture lexical annotation values without entering a construction frame.""" - # Inject builtin objects under fresh names: a source binding named globals, - # locals, iter or next must not replace these generated operations. - # Each function retains only annotation/constexpr names. Refer directly to - # the single root scope; never copy globals or enclosing locals wholesale. + """Snapshot only lexical annotation/constexpr names at their definition site.""" names = { - name - for name, alias in aliases.items() - if name != alias or name not in self.module.bindings + name for name, binding in bindings.items() if isinstance(binding, ast.Subscript) } | { parameter.arg for parameter in [*node.args.posonlyargs, *node.args.args, *node.args.kwonlyargs] @@ -1673,42 +1650,25 @@ def _create_definition_bindings( } if not names: return [] - values: list[ast.expr] = [] - for name in sorted(names): - value: ast.expr = ( - self._inject(getattr(builtins, name)) - if name in aliases and hasattr(builtins, name) - else ast.Attribute( - ast.Name(self.module.infrastructure_name, ast.Load()), "MISSING", ast.Load() - ) - ) - scopes: list[ast.expr] = [ast.Call(self._inject(globals), [], [])] - if self.module.definition_scope_name is not None and not local_function: - scopes.append(ast.Name(self.module.definition_scope_name, ast.Load())) - scopes.append(ast.Call(self._inject(locals), [], [])) - for scope in scopes: - # locals().get("n", definition_scope.get("n", globals().get("n", MISSING))) - value = ast.Call( - ast.Attribute(scope, "get", ast.Load()), [ast.Constant(name), value], [] - ) - values.append(value) - captures_expr = ast.Dict([ast.Constant(name) for name in sorted(names)], values) - # _definition = {"n": resolved_definition_value, ...} - statements: list[ast.stmt] = [self._assign(captures, captures_expr, node)] - for name, alias in aliases.items(): - if name == alias and name in self.module.bindings: - continue - fallback = ( - self._inject(getattr(builtins, name)) - if hasattr(builtins, name) - else ast.Attribute( - ast.Name(self.module.infrastructure_name, ast.Load()), "MISSING", ast.Load() - ) - ) - # _annotation = _definition.get("n", MISSING) - value = self._call(captures, "get", [ast.Constant(name), fallback], node) - statements.append(self._assign(alias, value, node)) - return statements + # Inject builtin operations so same-named source bindings cannot replace them. + # Acquire each scope once; the snapshot retains selected values, never frames. + scopes: list[ast.expr] = [ast.Call(self._inject(locals), [], [])] + if self.module.definition_scope_name is not None and not local_function: + scopes.append(ast.Name(self.module.definition_scope_name, ast.Load())) + scopes.append(ast.Call(self._inject(globals), [], [])) + defaults = { + name: getattr(builtins, name) + for name in names + if name in bindings and hasattr(builtins, name) + } + if defaults: + scopes.append(self._inject(defaults)) + captured = ast.Call( + self._inject(_AnnotationScope), + [ast.Tuple([ast.Constant(name) for name in sorted(names)], ast.Load()), *scopes], + [], + ) + return [self._assign(captures, captured, node)] def _create_specialization_bindings( self, node: ast.FunctionDef, *, special: str @@ -1726,10 +1686,10 @@ def _create_specialization_bindings( def _create_symbol_declarations( self, node: ast.FunctionDef - ) -> tuple[list[ast.stmt], dict[str, str]]: + ) -> tuple[list[ast.stmt], dict[str, ast.expr]]: """Bind explicit signature type parameters with their declared dtypes.""" declaration: list[ast.stmt] = [] - symbol_aliases: dict[str, str] = {} + symbol_aliases: dict[str, ast.expr] = {} # -------------------- Pattern -------------------- # Python source: # def f[n](): @@ -1752,7 +1712,7 @@ def _create_symbol_declarations( if getattr(parameter, "default_value", None) is not None: self._raise_error(parameter, "A symbolic type parameter cannot have a default") alias = self.module.fresh("_symbol") - symbol_aliases[parameter.name] = alias + symbol_aliases[parameter.name] = ast.Name(alias, ast.Load()) declaration.append( self._assign( alias, @@ -1814,7 +1774,15 @@ def _rewrite_parameters( is_constexpr = self._is_constexpr_annotation(annotation) group = constexpr_params if is_constexpr else other_params group.append((parameter, annotation, is_constexpr)) - for parameter, annotation, is_constexpr in [*constexpr_params, *other_params]: + constexpr_reads = { + parameter.arg + for parameter, _, _ in constexpr_params + if parameter.arg in self.function.annotation_bindings + } + constexpr_scope = self.module.fresh("_constexpr") if constexpr_reads else None + for index, (parameter, annotation, is_constexpr) in enumerate( + [*constexpr_params, *other_params] + ): if annotation is None: self._raise_error(parameter, f"Parameter {parameter.arg!r} requires an annotation") name = parameter.arg @@ -1864,7 +1832,30 @@ def _rewrite_parameters( else self._select_specialized_value(name, fallback, parameter, special=special) ) declaration.append(self._assign(alias, value, parameter)) - self.function.annotation_aliases[name] = alias + self.function.annotation_bindings[name] = ( + ast.Subscript(ast.Name(constexpr_scope, ast.Load()), ast.Constant(name), ast.Load()) + if is_constexpr and name in constexpr_reads + else ast.Name(alias, ast.Load()) + ) + if constexpr_scope is not None and index + 1 == len(constexpr_params): + # Selected host values may be MISSING. Preserve lazy missing-name + # errors without changing definition captures or Python body locals. + names = sorted(constexpr_reads) + values = ast.Dict( + [ast.Constant(name) for name in names], + [ast.Name(constexpr_aliases[name], ast.Load()) for name in names], + ) + declaration.append( + self._assign( + constexpr_scope, + ast.Call( + self._inject(_AnnotationScope), + [ast.Tuple([ast.Constant(name) for name in names], ast.Load()), values], + [], + ), + parameter, + ) + ) return declaration, constexpr_aliases def _create_function_frame( @@ -2073,9 +2064,9 @@ def create_function_builder_fragments( ) -> tuple[list[ast.stmt], str, ast.With]: """Declare a native frame and emit a lexical body helper inside its scope. - Definition aliases retain outer annotation values. Signature aliases add + Definition captures retain outer annotation values. Signature bindings add declared symbols and each preceding parameter; constexpr aliases retain - compile-time values for the body. One alias map serves annotation syntax; + compile-time values for the body. One substitution map serves annotation syntax; ordinary body reads retain Python globals/closures. None owns native IR. """ kind, options = self.read_function_metadata(node) @@ -2091,11 +2082,11 @@ def create_function_builder_fragments( if node.args.vararg or node.args.kwarg: self._raise_error(node, "IR signatures require ordinary named parameters") facts = self.module.prescan.bindings.get(node, []) + captures = self.module.fresh("_definition") annotations, returns, definition_aliases = self._read_function_annotations( - node, parameters, facts + node, parameters, facts, captures=captures ) - self.function.annotation_aliases = definition_aliases - captures = self.module.fresh("_definition") + self.function.annotation_bindings = definition_aliases statements = self._create_definition_bindings( node, definition_aliases, captures=captures, local_function=local_function ) @@ -2133,7 +2124,7 @@ def create_function_builder_fragments( returns, ) ) - definition_aliases.update(self.function.annotation_aliases) + definition_aliases.update(self.function.annotation_bindings) with self._bypass_rewrite(): options = self.visit(options) frame_declaration, body_entry = self._create_function_frame( @@ -2148,35 +2139,6 @@ def create_function_builder_fragments( node, parameters, constexpr_aliases, frame=frame, special=special ) body.extend(self.transform_statements(node.body)) - if not local_function and node.name not in self.module.source_functions: - # Source-text bodies use the supplied environment, not - # definition-only locals introduced for their annotations. - referenced = { - item.id - for statement in body - for item in ast.walk(statement) - if isinstance(item, ast.Name) - } - # Prefix statements in source text execute in the outer - # builder callable. Keep their values as Python closures; - # only externally supplied bindings belong to its globals. - prefix_names = { - item.name - for scope, bindings in self.module.prescan.bindings.items() - if isinstance(scope, ast.Module) - for item in bindings - } - global_names = sorted( - name - for name, alias in definition_aliases.items() - if name == alias - and name not in prefix_names - and name in referenced - and name not in self.module.prescan.namespaces - and name not in {parameter.arg for parameter in parameters} - ) - if global_names: - body.insert(0, ast.copy_location(ast.Global(global_names), node)) definition = self._create_definition(body_name, body, node) if not local_function: protected = ( diff --git a/python/tvm/tirx/script/ir_builder/__init__.py b/python/tvm/tirx/script/ir_builder/__init__.py index 8eb021065726..6f2e2ed55012 100644 --- a/python/tvm/tirx/script/ir_builder/__init__.py +++ b/python/tvm/tirx/script/ir_builder/__init__.py @@ -98,9 +98,7 @@ @_result_span("T.Buffer") @_mutable_cell_decl("T.Buffer", syntax="parameter") -@_annotation_constructor( - "shape", "strides", "elem_offset", "byte_offset", "allocated_addr", as_type=True -) +@_annotation_constructor def Buffer( shape, dtype="float32", diff --git a/python/tvm/tirx/script/ir_builder/parser_protocol.py b/python/tvm/tirx/script/ir_builder/parser_protocol.py index 0d170dbc548e..9046fbca9ed8 100644 --- a/python/tvm/tirx/script/ir_builder/parser_protocol.py +++ b/python/tvm/tirx/script/ir_builder/parser_protocol.py @@ -273,7 +273,6 @@ def func_name(name: str) -> None: def func_ret_type(annotation: Any, *, span: _Span = None) -> None: """Implements :func:`tvm.script.ir_builder.parser_protocol.func_ret_type`.""" - annotation = _base._return_annotation(annotation) if callable(annotation) and not isinstance(annotation, _ir.Expr | _ir.Type): annotation = annotation() if isinstance(annotation, _ir.Expr): diff --git a/tests/python/script/minilang.py b/tests/python/script/minilang.py index 25905fdf6c8f..ae44b255bc8f 100644 --- a/tests/python/script/minilang.py +++ b/tests/python/script/minilang.py @@ -160,8 +160,6 @@ def __init__(self): at_=self.at, with_at_group_=self.with_at_group, module_member_=lambda name, value: value, - require_defined=self.require_defined, - annotation_value_=lambda name, value: value, MISSING=self.missing, check_well_formed_=lambda result: None, constexpr=registry.constexpr, @@ -243,11 +241,6 @@ def get(self): def frame(self): return next(frame for frame in reversed(self.stack) if frame.kind == "function") - def require_defined(self, value, name): - if value is self.missing: - raise NameError(name) - return value - def func_name(self, name): self.frame().function.name = name self.frame().global_var = self.references.setdefault(name, Value("global", (name,))) From fb70668062ea723dea2d25252b60060b9dbec971 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 22:35:19 +0000 Subject: [PATCH 08/16] [REFACTOR][TVMScript] Preserve deferred return constructor evaluation --- python/tvm/relax/script/ir_builder/parser_protocol.py | 2 +- python/tvm/script/ir_builder/base.py | 7 +++++++ python/tvm/tirx/script/ir_builder/parser_protocol.py | 1 + 3 files changed, 9 insertions(+), 1 deletion(-) diff --git a/python/tvm/relax/script/ir_builder/parser_protocol.py b/python/tvm/relax/script/ir_builder/parser_protocol.py index d3bcedd936bb..f7ba61e7c769 100644 --- a/python/tvm/relax/script/ir_builder/parser_protocol.py +++ b/python/tvm/relax/script/ir_builder/parser_protocol.py @@ -374,7 +374,7 @@ def func_name(name: str) -> None: def func_ret_type(annotation: Any, *, span: _Span = None) -> None: """Implements :func:`tvm.script.ir_builder.parser_protocol.func_ret_type`.""" - return _native.func_ret_type(_builder._type(annotation)) + return _native.func_ret_type(_builder._type(_base._return_annotation(annotation))) def check_well_formed_(function: _relax.Function) -> None: diff --git a/python/tvm/script/ir_builder/base.py b/python/tvm/script/ir_builder/base.py index 87d459f4ccd3..d3577fe0720b 100644 --- a/python/tvm/script/ir_builder/base.py +++ b/python/tvm/script/ir_builder/base.py @@ -446,6 +446,13 @@ def _current_function_frame(): raise ValueError("Symbol resolution requires an active function frame") +def _return_annotation(annotation): + """Evaluate the deferred return expression before normalizing its annotation value.""" + if callable(annotation) and not isinstance(annotation, ir.Expr | ir.Type): + return annotation() + return annotation + + def annotation_constructor(constructor): """Expose a constructor as a real annotation class supporting Python unions. diff --git a/python/tvm/tirx/script/ir_builder/parser_protocol.py b/python/tvm/tirx/script/ir_builder/parser_protocol.py index 9046fbca9ed8..0d170dbc548e 100644 --- a/python/tvm/tirx/script/ir_builder/parser_protocol.py +++ b/python/tvm/tirx/script/ir_builder/parser_protocol.py @@ -273,6 +273,7 @@ def func_name(name: str) -> None: def func_ret_type(annotation: Any, *, span: _Span = None) -> None: """Implements :func:`tvm.script.ir_builder.parser_protocol.func_ret_type`.""" + annotation = _base._return_annotation(annotation) if callable(annotation) and not isinstance(annotation, _ir.Expr | _ir.Type): annotation = annotation() if isinstance(annotation, _ir.Expr): From 79bfa0267bfdae935d8586412e0c06ad3122d95e Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 22:42:17 +0000 Subject: [PATCH 09/16] [REFACTOR][TVMScript] Bind captured scalar parameter symbols explicitly --- tests/python/tvmscript/test_tvmscript_parser_tir.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/python/tvmscript/test_tvmscript_parser_tir.py b/tests/python/tvmscript/test_tvmscript_parser_tir.py index 44c8c24c7f35..474ddc526d30 100644 --- a/tests/python/tvmscript/test_tvmscript_parser_tir.py +++ b/tests/python/tvmscript/test_tvmscript_parser_tir.py @@ -82,7 +82,7 @@ def test_tir_external_symbol_adopted_by_later_prim_param(): """ n = T.dynamic("n", "int32") @T.prim_func -def main(A: T.Buffer((n,), "float32"), n: T.int32): +def main(A: T.Buffer((n,), "float32"), n: n): T.evaluate(n) """ ) @@ -97,7 +97,7 @@ def main(A: T.Buffer((n,), "float32"), n: T.int32): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), "float32"), n: T.int32): + def main(A: T.Buffer((n,), "float32"), n: n): T.evaluate(n) """ ) @@ -112,7 +112,7 @@ def test_tir_external_symbol_preserves_later_prim_param_dtype(): """ n = T.dynamic("n", "int64") @T.prim_func -def main(A: T.Buffer((n,), "float32"), n: T.int64): +def main(A: T.Buffer((n,), "float32"), n: n): T.evaluate(n) """ ) From 757047a9388690101ef22a9707128cb9eb0c7d95 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 23:13:25 +0000 Subject: [PATCH 10/16] [FIX][TVMScript] Allow body locals to shadow annotation captures --- python/tvm/script/parser/prescan.py | 26 +++------------- tests/python/script/test_symbolic_shape.py | 36 +++++++++------------- 2 files changed, 18 insertions(+), 44 deletions(-) diff --git a/python/tvm/script/parser/prescan.py b/python/tvm/script/parser/prescan.py index 59de07454999..dd5153b3aa2c 100644 --- a/python/tvm/script/parser/prescan.py +++ b/python/tvm/script/parser/prescan.py @@ -446,37 +446,19 @@ def visit_FunctionDef(self, node: ast.FunctionDef | ast.AsyncFunctionDef) -> Non visit_AsyncFunctionDef = visit_FunctionDef def _validate_symbols(self, node: ast.FunctionDef | ast.AsyncFunctionDef) -> None: - """Reject ordinary writes after symbolic introduction in this function.""" + """Reject ordinary writes to explicit function-local symbolic declarations.""" facts = self.bindings[node] - # Mutable storage and loop/parameter names retain their existing assignment - # rules. Annotation lambda/comprehension binders are handled by free-name lookup. + # Captured annotation values belong to the definition scope and may be + # shadowed by ordinary body locals. Only header declarations introduce symbols. + # Mutable storage and loop/parameter names retain their assignment rules. assignable = { item.name for item in facts if item.kind in ("parameter", "mutable_parameter", "mutable", "loop") } - annotations = [item.annotation for item in facts if item.annotation is not None] - if node.returns is not None: - annotations.append(node.returns) # This temporary diagnostic index is derived from existing source facts; # it is not retained in the prescan result or used as value state. origins: dict[str, ast.AST] = {} - for annotation in annotations: - for name, annotation_origin in collect_annotation_free_names(annotation).items(): - # A prior body target makes this a local annotation operand, - # not a free symbolic introduction. Explicit symbols below - # still establish their own origin and reject ordinary writes. - bound_in_body = any( - item.name == name - and isinstance(item.node, ast.Name) - and (item.node.lineno, item.node.col_offset) - < (annotation_origin.lineno, annotation_origin.col_offset) - for item in facts - ) - if name not in assignable and name not in self.namespaces and not bound_in_body: - previous = origins.get(name) - if previous is None or annotation_origin.lineno < previous.lineno: - origins[name] = annotation_origin for item in facts: if item.kind == "symbol": origin = origins.get(item.name) diff --git a/tests/python/script/test_symbolic_shape.py b/tests/python/script/test_symbolic_shape.py index 4e6a5ae6b584..1d0c760da496 100644 --- a/tests/python/script/test_symbolic_shape.py +++ b/tests/python/script/test_symbolic_shape.py @@ -21,8 +21,7 @@ from __future__ import annotations # Invalid script examples deliberately contain unresolved or unused bindings. -# ruff: noqa: F821, F841 -import inspect +# ruff: noqa: F821 import sys import pytest @@ -97,11 +96,6 @@ def main(x: T.Buffer(shape, "float32")): build(("n", 16)) -def _line_of(function, statement): - lines, first = inspect.getsourcelines(function) - return first + next(index for index, line in enumerate(lines) if line.strip() == statement) - - def test_argument_policies_reuse_symbols_and_resolve_only_marked_literals(language): # Shape expressions reuse the external symbol while marked device strings resolve. M = language.M @@ -133,28 +127,26 @@ def main(x: M.Tensor((symbol, symbol))): M.record(n) +@pytest.mark.skipif(sys.version_info < (3, 12), reason="PEP 695 requires Python 3.12") def test_symbol_reassignment_reports_introduction_and_exact_write(language): - # Ordinary writes to a symbolic dimension must point to its original introduction and - # exact target. - M = language.M - n = M.dynamic("n") + # Ordinary writes to a header symbol must point to its declaration and exact target. + source = """ +@M.function +def main[n](x: M.Tensor((n,))): + n = 2 +""" + filename = "symbol_reassignment.py" with pytest.raises(SyntaxError) as caught: - - @M.function - def main(x: M.Tensor((n,))): - n = 2 + entry.parse(source, extra_vars={"M": language.M}, filename=filename) message = str(caught.value) - introduction = _line_of( - test_symbol_reassignment_reports_introduction_and_exact_write, - "def main(x: M.Tensor((n,))):", - ) - offending = _line_of(test_symbol_reassignment_reports_introduction_and_exact_write, "n = 2") + introduction = 3 + offending = 4 assert "Symbolic variable 'n' cannot be reassigned" in message assert f"introduced at line {introduction}" in message error = caught.value - assert (error.filename, error.lineno, error.end_lineno) == (__file__, offending, offending) - assert (error.offset, error.end_offset) == (13, 14) + assert (error.filename, error.lineno, error.end_lineno) == (filename, offending, offending) + assert (error.offset, error.end_offset) == (5, 6) def test_external_dynamic_symbols_reuse_identity(language): From f1a14be1044af9236243daef8e0c3d334aef4f5f Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 23:49:44 +0000 Subject: [PATCH 11/16] [FIX][Relax] Preserve frontend batch dimensions in sampling outputs --- python/tvm/relax/frontend/nn/op.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/python/tvm/relax/frontend/nn/op.py b/python/tvm/relax/frontend/nn/op.py index f7e008895ec5..b45188db5034 100644 --- a/python/tvm/relax/frontend/nn/op.py +++ b/python/tvm/relax/frontend/nn/op.py @@ -2778,7 +2778,7 @@ def sample_top_p_top_k_from_sorted_prob( prob_dtype = sorted_prob.dtype index_dtype = sorted_index.dtype prob_batch = sorted_prob.shape[0] - out_batch = uniform_sample.shape[0] + sample_batch = uniform_sample.shape[0] if sample_indices is not None: assert sample_indices.shape == uniform_sample.shape, ( @@ -2789,7 +2789,7 @@ def sample_top_p_top_k_from_sorted_prob( "Number of samples must match the number of probability distributions." ) sample_indices = Tensor.from_const( - np.arange(out_batch).reshape(out_batch, 1).astype(np.int64) + np.arange(sample_batch).reshape(sample_batch, 1).astype(np.int64) ) print("sample_indices: ", sample_indices) sample_indices_dtype = sample_indices.dtype @@ -2867,7 +2867,7 @@ def _get_index_from_sorted( _get_index_from_sorted, "get_index_from_sorted", args=[cumsum_sorted, sorted_index, renorm_prob, uniform_sample, sample_indices], - out=Tensor.placeholder([out_batch, 1], index_dtype), + out=Tensor.placeholder([sample_batch, 1], index_dtype), ) return out_index_in_sorted @@ -2902,7 +2902,7 @@ def renormalize_top_p_top_k_prob(prob, sorted_prob, top_p, top_k): """ prob_dtype = prob.dtype top_k_dtype = top_k.dtype - batch = sorted_prob.shape[0] + prob_batch = sorted_prob.shape[0] def _cumsum_mask(cumsum_sorted, top_p, top_k, i, j): return _tir.all(cumsum_sorted[i, j] < top_p[i, 0], j + 1 < top_k[i, 0]) @@ -2935,7 +2935,7 @@ def _get_renorm_cutoff(A: T.handle, B: T.handle, C: T.handle, D: T.handle, E: T. "get_renorm_cutoff", args=[sorted_prob, cumsum_sorted, top_p, top_k], out=Tensor.placeholder( - [batch, 1], + [prob_batch, 1], prob_dtype, ), ) From c8a3a46f1b32258c9d9a5d2966889fc0cf71e10e Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Thu, 24 Sep 2026 23:40:15 +0000 Subject: [PATCH 12/16] [Script] Keep dynamic shape captures distinct from local bindings --- tests/python/relax/test_dataflow_inplace.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/tests/python/relax/test_dataflow_inplace.py b/tests/python/relax/test_dataflow_inplace.py index b78efa188f19..27ba3f09a04b 100644 --- a/tests/python/relax/test_dataflow_inplace.py +++ b/tests/python/relax/test_dataflow_inplace.py @@ -546,15 +546,15 @@ def main( def test_dynamic(): - a = T.dynamic("a") + a_dim = T.dynamic("a") b = T.dynamic("b") @I.ir_module class DynamicTestCase: @R.function def main( - x: R.Tensor((a, b), dtype="float32"), y: R.Tensor((a, b), dtype="float32") - ) -> R.Tensor((a, b), dtype="float32"): + x: R.Tensor((a_dim, b), dtype="float32"), y: R.Tensor((a_dim, b), dtype="float32") + ) -> R.Tensor((a_dim, b), dtype="float32"): with R.dataflow(): z = R.add(x, y) # Cannot be done in-place because x and y are arguments @@ -637,7 +637,7 @@ def main( def test_dynamic_mismatch(): # cannot statically prove the shapes to be equal so the module should be unchanged - a = T.dynamic("a") + a_dim = T.dynamic("a") b = T.dynamic("b") c = T.dynamic("c") d = T.dynamic("d") @@ -645,7 +645,7 @@ def test_dynamic_mismatch(): @I.ir_module class DynamicMistmatchTestCase: @R.function - def main(x: R.Tensor((a, b), dtype="float32"), y: R.Tensor((c, d), dtype="float32")): + def main(x: R.Tensor((a_dim, b), dtype="float32"), y: R.Tensor((c, d), dtype="float32")): with R.dataflow(): z = R.add(x, y) # Cannot be done in-place because x and y are arguments From 72e236c353998974fdfc41cfd0001309d05bfbba Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Fri, 25 Sep 2026 00:32:58 +0000 Subject: [PATCH 13/16] [REFACTOR][TVMScript] Align retained builder protocols with direct arguments --- docs/reference/api/python/script/parser.rst | 2 +- python/tvm/script/ir_builder/base.py | 1 + python/tvm/script/parser/transpile.py | 1 - tests/python/script/test_special_parser_protocol.py | 1 - 4 files changed, 2 insertions(+), 3 deletions(-) diff --git a/docs/reference/api/python/script/parser.rst b/docs/reference/api/python/script/parser.rst index 2b2f7c4b26be..453ceaa0a203 100644 --- a/docs/reference/api/python/script/parser.rst +++ b/docs/reference/api/python/script/parser.rst @@ -32,7 +32,7 @@ registered through :func:`tvm.script.parser.register_namespace` and :func:`tvm.script.parser.register_namespace_initializer`. .. automodule:: tvm.script.parser.protocol_registry - :members: constexpr, args_policy, register_scalar_annotation, mutable_cell_decl, result_span, module_decorator, declaration_kind + :members: constexpr, register_scalar_annotation, mutable_cell_decl, result_span, module_decorator, declaration_kind The language variant aliases below share the public construction namespaces documented in :doc:`script`. Parser entry points above use the canonical frontend. diff --git a/python/tvm/script/ir_builder/base.py b/python/tvm/script/ir_builder/base.py index d3577fe0720b..bba7ea947121 100644 --- a/python/tvm/script/ir_builder/base.py +++ b/python/tvm/script/ir_builder/base.py @@ -18,6 +18,7 @@ from collections.abc import Callable from contextlib import contextmanager, nullcontext +from functools import wraps from inspect import signature from typing import Any, Generic, TypeVar diff --git a/python/tvm/script/parser/transpile.py b/python/tvm/script/parser/transpile.py index 400b7830a944..f8b8af648daa 100644 --- a/python/tvm/script/parser/transpile.py +++ b/python/tvm/script/parser/transpile.py @@ -531,7 +531,6 @@ def _is_module_owner(self, node: ast.expr) -> bool: ] return bool(records) and all(item.kind == "module_alias" for item in records) - def _visit_direct_operand(self, node: ast.expr) -> ast.expr: """Preserve an existing payload's span without bypassing child operations.""" if isinstance(node, ast.Name): diff --git a/tests/python/script/test_special_parser_protocol.py b/tests/python/script/test_special_parser_protocol.py index d85d11df81ea..6cb7c63bf838 100644 --- a/tests/python/script/test_special_parser_protocol.py +++ b/tests/python/script/test_special_parser_protocol.py @@ -25,7 +25,6 @@ from tvm.script import tirx as T from tvm.script.ir_builder import resolve_global_info_args -from tvm.script.parser import protocol_registry as registry def test_constructor_policy_survives_a_failed_definition(language): From 11539f8aa75b4bccb629ae5f06cde418ef8cb582 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Fri, 25 Sep 2026 01:12:24 +0000 Subject: [PATCH 14/16] [TEST][TVMScript] Keep dynamic name hints literal in address matching --- tests/python/tvmscript/test_tvmscript_printer_tir.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/python/tvmscript/test_tvmscript_printer_tir.py b/tests/python/tvmscript/test_tvmscript_printer_tir.py index aba2cee8fd13..15c704a5d7d2 100644 --- a/tests/python/tvmscript/test_tvmscript_printer_tir.py +++ b/tests/python/tvmscript/test_tvmscript_printer_tir.py @@ -931,8 +931,8 @@ def func(a_name: T.handle): expected_regex = re.escape(without_address) for name in ["a_name", "A_name", "N_name", "i_name"]: - # Replace all occurrences with a backref to an earlier match - expected_regex = expected_regex.replace(name, rf"(?P={name})") + # Identifiers carry addresses; dynamic name-hint strings stay literal. + expected_regex = re.sub(rf"(?{name}_0x[A-Fa-f0-9]+)", 1 From d3fe3111b55765ae5004509d576eb91949b619cf Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Fri, 25 Sep 2026 01:19:14 +0000 Subject: [PATCH 15/16] [REFACTOR][TVMScript] Share explicit symbol annotations with S-TIR --- docs/arch/tvmscript.rst | 18 ++++++++++++++++-- python/tvm/s_tir/script/__init__.py | 3 +-- .../tvmscript/test_tvmscript_printer_tir.py | 2 +- .../test_tvmscript_s_tir_namespace.py | 4 +++- 4 files changed, 21 insertions(+), 6 deletions(-) diff --git a/docs/arch/tvmscript.rst b/docs/arch/tvmscript.rst index 7ee7e4556ee1..ef95206b13f1 100644 --- a/docs/arch/tvmscript.rst +++ b/docs/arch/tvmscript.rst @@ -45,6 +45,20 @@ function decorators retain definitions until module construction, which can decl signatures before building bodies. Captured Python values belong to the source context; symbolic IR values are created and resolved by builders. +Annotations read concrete values from their definition scope, preserving missing-name +errors when a value is used. Create external symbols with ``n = I.dynamic("n")`` +or the identical ``T.dynamic`` and ``Ts.dynamic`` constructors. Each call creates a +fresh native variable, defaulting to int64. On Python 3.12+, explicit headers such as +``def f[n, k: T.int32](...)`` declare local symbols; ``n: int`` retains the int64 +default. Quote the whole annotation or use ``from __future__ import annotations`` +to defer eager Python evaluation of header symbols. Captured runtime ``typing.TypeVar`` +objects are not script symbols; ordinary Python typing uses remain unaffected. + +An explicit scalar annotation ``n: n`` preserves a captured native symbol's identity. +An independently typed parameter such as ``n: T.int32`` and ordinary body locals +shadow definition captures normally. Annotation classes, Python unions and deferred +return-constructor evaluation retain their builder behavior. + Syntax and construction protocol -------------------------------- @@ -55,8 +69,8 @@ builder frames. The namespace owns the meaning of these operations and the suppo IR constructs. ``tvm.script.parser.protocol_registry`` records syntax policies under registered -namespace paths. These policies identify declarations and arguments that need special -handling, such as symbolic expression strings. Source aliases resolve to those paths; +namespace paths. These policies identify scalar annotations, mutable declarations and +result span handling. Symbolic shapes use concrete expressions. Source aliases resolve to those paths; ordinary Python calls remain calls in the generated program. Explicit ``constexpr`` markers select host control flow during construction. diff --git a/python/tvm/s_tir/script/__init__.py b/python/tvm/s_tir/script/__init__.py index b6536aab80c4..7f9c18194130 100644 --- a/python/tvm/s_tir/script/__init__.py +++ b/python/tvm/s_tir/script/__init__.py @@ -58,8 +58,7 @@ def _initialize(): globals()[name] = getattr(builder, name) # The operations are shared, but syntax is keyed by the registered namespace. for table in ( - protocol_registry.ARGS_POLICIES, - protocol_registry.TYPE_VAR_DECL, + protocol_registry.SCALAR_ANNOTATION_DTYPE, protocol_registry.MUTABLE_CELL_DECL, protocol_registry.RESULT_SPAN, ): diff --git a/tests/python/tvmscript/test_tvmscript_printer_tir.py b/tests/python/tvmscript/test_tvmscript_printer_tir.py index 15c704a5d7d2..b2bdb76b00f5 100644 --- a/tests/python/tvmscript/test_tvmscript_printer_tir.py +++ b/tests/python/tvmscript/test_tvmscript_printer_tir.py @@ -25,8 +25,8 @@ from tvm import ir, s_tir, tirx from tvm.ir import Range from tvm.s_tir.script.ir_builder import prim_func as build_prim_func -from tvm.script import s_tir as Ts from tvm.script import ir as I +from tvm.script import s_tir as Ts from tvm.script.ir_builder import IRBuilder from tvm.tirx.script import ir_builder as T diff --git a/tests/python/tvmscript/test_tvmscript_s_tir_namespace.py b/tests/python/tvmscript/test_tvmscript_s_tir_namespace.py index 16395621acb2..57f6aadfc736 100644 --- a/tests/python/tvmscript/test_tvmscript_s_tir_namespace.py +++ b/tests/python/tvmscript/test_tvmscript_s_tir_namespace.py @@ -33,8 +33,10 @@ def test_shared_operations_and_aliases(): from tvm.script import s_tir as S + n = S.dynamic("n") + @S.prim_func - def shared(A: S.Buffer(("n",), "float32")): + def shared(A: S.Buffer((n,), "float32")): for i in T.serial(A.shape[0]): with S.sblock("copy"): v = Ts.axis.spatial(A.shape[0], i) From 9e40422706d28f93c37d32872e58357ef1d337c0 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Fri, 25 Sep 2026 01:22:49 +0000 Subject: [PATCH 16/16] [REFACTOR][TVMScript] Drop unused stride captures from lowered kernel --- .../codegen/test_target_codegen_vulkan.py | 21 ------------------- 1 file changed, 21 deletions(-) diff --git a/tests/python/codegen/test_target_codegen_vulkan.py b/tests/python/codegen/test_target_codegen_vulkan.py index f8664e1f7fad..3b371fef8752 100644 --- a/tests/python/codegen/test_target_codegen_vulkan.py +++ b/tests/python/codegen/test_target_codegen_vulkan.py @@ -459,27 +459,6 @@ def test_cooperative_matrix(out_dtype): M, N, K = 16, 16, 32 # fmt: off - C_0_s0 = T.dynamic("C_0_s0") - C_0_s1 = T.dynamic("C_0_s1") - A_3_s0 = T.dynamic("A_3_s0") - A_3_s1 = T.dynamic("A_3_s1") - C_4_s0 = T.dynamic("C_4_s0") - C_4_s1 = T.dynamic("C_4_s1") - A_2_s0 = T.dynamic("A_2_s0") - A_2_s1 = T.dynamic("A_2_s1") - B_0_s0 = T.dynamic("B_0_s0") - B_0_s1 = T.dynamic("B_0_s1") - C_3_s0 = T.dynamic("C_3_s0") - C_3_s1 = T.dynamic("C_3_s1") - A_0_s0 = T.dynamic("A_0_s0") - A_0_s1 = T.dynamic("A_0_s1") - C_1_s0 = T.dynamic("C_1_s0") - C_1_s1 = T.dynamic("C_1_s1") - A_1_s0 = T.dynamic("A_1_s0") - A_1_s1 = T.dynamic("A_1_s1") - C_2_s0 = T.dynamic("C_2_s0") - C_2_s1 = T.dynamic("C_2_s1") - @I.ir_module class Module: @T.prim_func