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/docs/deep_dive/relax/tutorials/relax_creation.py b/docs/deep_dive/relax/tutorials/relax_creation.py index f1437b66e008..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) @@ -70,12 +72,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 +89,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 +168,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 +232,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..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 @@ -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): @@ -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") @@ -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/docs/reference/api/python/script/parser.rst b/docs/reference/api/python/script/parser.rst index 02e5304d606c..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_type_var_decl, 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/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/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..e62d71568afe 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 @@ -642,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/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..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 @@ -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) @@ -2863,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 @@ -2898,14 +2902,16 @@ 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]) + 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) @@ -2929,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, ), ) diff --git a/python/tvm/relax/script/ir_builder/__init__.py b/python/tvm/relax/script/ir_builder/__init__.py index 29987e713f99..a9919cf4d264 100644 --- a/python/tvm/relax/script/ir_builder/__init__.py +++ b/python/tvm/relax/script/ir_builder/__init__.py @@ -30,8 +30,6 @@ from tvm.script.ir_builder import resolve_global_info_args as _resolve_global_info_args 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 +82,6 @@ @_resolve_global_info_args("vdevice", resolver=resolve_global_info_) -@_args_policy("R.Tensor", {"shape": "expr_str"}, scalar_strings=False) def Tensor(shape=None, dtype=None, vdevice=None, ndim=-1, *, span=None): """Construct a Relax tensor type. @@ -92,8 +89,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 @@ -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,15 +116,14 @@ 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) def DTensor(shape=None, dtype=None, device_mesh=None, placement="", *, ndim=-1, span=None): """Construct a Relax distributed tensor type. 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 @@ -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,9 +151,8 @@ 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 -_ARGS_POLICIES["R.dist.DTensor"] = _ARGS_POLICIES["R.DTensor"] dist.device_mesh = device_mesh Range = _ir.Range @@ -166,15 +161,14 @@ 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") def Shape(values=None, ndim=-1, *, span=None): """Construct a Relax shape type. 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. 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/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/python/tvm/script/ir_builder/__init__.py b/python/tvm/script/ir_builder/__init__.py index 93bcec4b8d64..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_, ) @@ -36,6 +34,7 @@ from .ir import ( decl_function, def_function, + dynamic, ir_module, lookup_name, meta_var, @@ -56,12 +55,12 @@ "Range", "StringImm", "StringType", - "annotation_value_", "at_", "check_well_formed_", "constexpr", "decl_function", "def_function", + "dynamic", "ir_module", "lookup_name", "meta_var", @@ -70,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 14a960e9ca4e..bba7ea947121 100644 --- a/python/tvm/script/ir_builder/base.py +++ b/python/tvm/script/ir_builder/base.py @@ -23,7 +23,7 @@ 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 +377,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 +439,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,106 +447,26 @@ 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 - - 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( - constructor.__name__, - (), - { - "__new__": lambda cls, *args, **kwargs: invoke(*args, **kwargs), - "__signature__": call_signature, - "__doc__": constructor.__doc__, - "__module__": constructor.__module__, - }, - ) - return result - - 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}") + """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_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. +def annotation_constructor(constructor): + """Expose a constructor as a real annotation class supporting Python unions. - Raises - ------ - TypeError - If a TypeVar has a bound or constraints. - ValueError - If a symbolic value requires resolution without an active function frame. + Calls construct ordinary native values directly. The class preserves the + constructor's signature and documentation without adapting its arguments. """ - 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/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..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. @@ -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 @@ -872,8 +871,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/__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/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/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 fb9da9177292..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 .expr_str_handling import parse_annotation +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/prescan.py b/python/tvm/script/parser/prescan.py index 3de503356391..dd5153b3aa2c 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( @@ -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,6 @@ 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) self._record_binding( arg.arg, arg, @@ -422,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) @@ -440,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) @@ -501,10 +489,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..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,44 +33,13 @@ 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] = {} -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] = {} @@ -111,110 +79,13 @@ def constexpr(value: object) -> NoReturn: raise TypeError("constexpr is a parser syntax marker, not a runtime operation") -def args_policy( - namespace_path: str, - fields: Mapping[str, str], - *, - 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. - - 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. - 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. - - Returns - ------- - Callable - Decorator retaining the existing builder-owned annotation adapter where - needed. Concrete arguments still invoke the original constructor. - - 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. - - .. 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_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 ---------- @@ -225,8 +96,8 @@ def register_type_var_decl( 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. + Static scalar dtype for an explicit PEP 695 symbol bound. None (the + default) leaves the annotation without a supported scalar bound dtype. Returns ------- @@ -235,18 +106,16 @@ def register_type_var_decl( 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. + This metadata does not give constructor calls special assignment semantics. + Scalar runtime parameters retain their ordinary annotation construction. .. code:: python - register_type_var_decl("T.int32", T.int32, dtype="int32") - # Source: n = T.int32() + register_scalar_annotation("T.int32", T.int32, dtype="int32") + # Source: def f[n: T.int32](...): # Builder: n = X.resolve_type_var_("n", dtype="int32") """ - 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..f8b8af648daa 100644 --- a/python/tvm/script/parser/transpile.py +++ b/python/tvm/script/parser/transpile.py @@ -38,13 +38,15 @@ import ast import builtins +import copy from collections.abc import Callable, Iterator, Mapping from contextlib import contextmanager from types import FunctionType 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 .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): @@ -157,8 +159,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 @@ -289,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]: @@ -361,32 +361,10 @@ 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: - # 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). @@ -402,37 +380,14 @@ 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 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 @@ -473,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 } ): @@ -491,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) @@ -576,50 +531,10 @@ 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 +549,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 +570,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 -------------------- @@ -927,7 +819,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 +854,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 +1011,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 -------------------- @@ -1709,11 +1581,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 @@ -1741,42 +1617,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] @@ -1784,42 +1649,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 @@ -1836,11 +1684,11 @@ def _create_specialization_bindings( return [self._assign(special, special_expr, node)] def _create_symbol_declarations( - self, node: ast.FunctionDef, facts: list[Binding] - ) -> tuple[list[ast.stmt], dict[str, str]]: - """Predeclare symbol types and bind explicit signature type parameters.""" + self, node: ast.FunctionDef + ) -> 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](): @@ -1849,38 +1697,33 @@ 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") 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") - symbol_aliases[parameter.name] = alias + symbol_aliases[parameter.name] = ast.Name(alias, ast.Load()) declaration.append( 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"): - 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 @@ -1930,7 +1773,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 @@ -1980,7 +1831,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( @@ -2189,9 +2063,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) @@ -2207,11 +2081,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 ) @@ -2225,7 +2099,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}), @@ -2249,7 +2123,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( @@ -2264,25 +2138,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) - } - global_names = sorted( - name - for name, alias in definition_aliases.items() - if name == alias - 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 babeaf151d32..6f2e2ed55012 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.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, @@ -96,44 +96,9 @@ 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( - "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 def Buffer( shape, dtype="float32", @@ -218,7 +183,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): @@ -376,15 +340,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. @@ -451,8 +406,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 5544ca32147c..53efba58f30e 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 @@ -323,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 = [] @@ -535,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 = [] @@ -1525,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 @@ -2038,7 +2034,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 +2223,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..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->func_vars->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 3063a7a740e7..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; { @@ -130,7 +133,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/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 dcf059ed8d99..f697bd9aaf0f 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; @@ -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)); } @@ -123,10 +123,10 @@ 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 (!declared_type_vars.count(name.value())) { - declared_type_vars.insert(name.value()); - type_var_decls.push_back(stmt); + if (ffi::Optional name = GetDynamicDeclarationName(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/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..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 = {}, - std::unordered_set stringify_compound_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,23 +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. - // Bare TypeVars are real annotation bindings; compound expressions involving - // them, or any expression referring to a later parameter, must 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)); - 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); @@ -113,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)); } @@ -154,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)); } @@ -175,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` { @@ -232,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)); } @@ -324,11 +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, - std::unordered_set stringify_compound_shape_vars) { + const IRDocsifier& d) { ffi::Map attrs = - BufferAttrs(buffer, p, frame, d, BufferVarDefinition::MatchBuffer, std::nullopt, true, - std::move(stringify_shape_vars), std::move(stringify_compound_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/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..12a64e5f2e99 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; @@ -80,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)); @@ -99,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(); @@ -109,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); @@ -119,32 +131,8 @@ TVM_FFI_STATIC_INIT_BLOCK() { (*f)->stmts.push_back(AssignDoc(lhs, rhs, std::nullopt)); 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) { - 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), - std::move(stringify_compound_shape_vars)); + ExprDoc annotation = BufferAttn(buffer, var_p->Attr("ty"), *f, d); args.push_back(AssignDoc(lhs, std::nullopt, annotation)); continue; } @@ -152,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 fca09c45b511..d6ae5a222ef9 100644 --- a/src/tirx/script/printer/utils.h +++ b/src/tirx/script/printer/utils.h @@ -322,15 +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. - * \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); /*! * \brief Print the creation of a Var 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/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..3b371fef8752 100644 --- a/tests/python/codegen/test_target_codegen_vulkan.py +++ b/tests/python/codegen/test_target_codegen_vulkan.py @@ -581,11 +581,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..27ba3f09a04b 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_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 @@ -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_dim = 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_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 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_tvmscript_parser.py b/tests/python/relax/test_tvmscript_parser.py index 5a0d39cfd68b..215608a2ce3f 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, m), "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,17 @@ 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: # 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 @@ -1521,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 @@ -1545,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 @@ -1564,13 +1592,14 @@ 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" ): - m = T.int64() z = R.call_dps_packed("test_intrin", (x, y), R.Tensor((T.max(m, 20) + 1,), dtype="float32")) return z @@ -1596,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) @@ -1619,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 """ @@ -1633,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 """ ) @@ -1649,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 """ ) @@ -1686,9 +1715,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 @@ -1703,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 @@ -1772,10 +1797,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 +1814,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) @@ -1802,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 @@ -2371,17 +2401,21 @@ 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") + 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 @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 +2433,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,10 +2451,12 @@ 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.""" + + 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 @@ -2484,9 +2522,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 +2543,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 +2595,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 +2610,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..a86682250474 100644 --- a/tests/python/relax/test_tvmscript_printer_relax.py +++ b/tests/python/relax/test_tvmscript_printer_relax.py @@ -56,12 +56,12 @@ 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")) 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") @@ -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") @@ -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..00499acb192b 100644 --- a/tests/python/relax/test_tvmscript_type_vars.py +++ b/tests/python/relax/test_tvmscript_type_vars.py @@ -16,21 +16,21 @@ # 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(): @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() @@ -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/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/script/minilang.py b/tests/python/script/minilang.py index de9c99b30e77..ae44b255bc8f 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] @@ -158,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, @@ -216,18 +216,18 @@ 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)) - 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 @@ -241,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,))) 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 20d6e2e0eeb7..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,25 +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")): - n = M.symbol() # noqa: F841 + def first(x: M.Tensor((first_n,), "float32")): return x @M.function - def second(x: M.Tensor(("n",), "float32")): - n = M.symbol() # noqa: F841 + 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..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): @@ -67,19 +66,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 fe281434b569..1d0c760da496 100644 --- a/tests/python/script/test_symbolic_shape.py +++ b/tests/python/script/test_symbolic_shape.py @@ -14,34 +14,34 @@ # 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. -Quoted dimensions must not create or replace ordinary Python bindings. +"""Explicit symbolic dimensions retain identity in signatures and function bodies. +Dimensions use ordinary Python expressions over explicit symbols. """ 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 -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 @@ -53,22 +53,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 @@ -80,8 +81,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")): @@ -89,29 +89,22 @@ 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( - 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): - 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): - # 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] @@ -123,48 +116,48 @@ 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) +@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 + # 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(): - n = M.symbol() - 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, "n = M.symbol()" - ) - 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_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 +169,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 +182,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/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/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/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( 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 372f25b105db..474ddc526d30 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: n): 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: n): 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: n): T.evaluate(n) """ ) @@ -119,12 +122,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) """ ) @@ -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,12 +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"): - n = T.int32() +def main() -> T.Buffer((n,), "float32"): A = T.alloc_buffer((n,), "float32") return A """ @@ -631,10 +633,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..b2bdb76b00f5 100644 --- a/tests/python/tvmscript/test_tvmscript_printer_tir.py +++ b/tests/python/tvmscript/test_tvmscript_printer_tir.py @@ -25,6 +25,7 @@ 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 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 @@ -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 @@ -929,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 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_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) 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) 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