diff --git a/docs/arch/codegen.rst b/docs/arch/codegen.rst index 799568b42039..168bb6d30592 100644 --- a/docs/arch/codegen.rst +++ b/docs/arch/codegen.rst @@ -152,7 +152,7 @@ backend (x86, ARM, NVPTX, AMDGPU, etc.). ``Cast`` → LLVM type conversions, ``Call`` → intrinsic or extern function calls. - **Statements** (``VisitStmt_``) emit LLVM IR side effects: ``BufferStore`` → store instructions, ``For`` → loop basic blocks with branches, - ``IfThenElse`` → conditional branches, ``AllocBuffer`` → stack or heap allocation. + ``IfThenElse`` → conditional branches, ``tirx.alloc_tensor`` calls → stack or heap allocation. The key methods on ``CodeGenLLVM`` are: diff --git a/docs/arch/tvmscript.rst b/docs/arch/tvmscript.rst index fc29376a8d2d..c7441abbc6ae 100644 --- a/docs/arch/tvmscript.rst +++ b/docs/arch/tvmscript.rst @@ -174,7 +174,7 @@ For example, a small function can be authored, printed and parsed again: from tvm.script import tirx as T @T.prim_func - def increment(A: T.Buffer((4,), "float32")): + def increment(A: T.Tensor((4,), "float32")): for i in T.serial(4): A[i] = A[i] + 1.0 diff --git a/docs/deep_dive/relax/learning.rst b/docs/deep_dive/relax/learning.rst index 3748aad9efe7..428ddff62f12 100644 --- a/docs/deep_dive/relax/learning.rst +++ b/docs/deep_dive/relax/learning.rst @@ -133,13 +133,13 @@ for the end-to-end model execution. The code block below shows a TVMScript imple class Module: M, N, K = T.int64(), T.int64(), T.int64() @Ts.prim_func(private=True) - def linear(X: T.Buffer((M, K), 'float32'), W: T.Buffer((K, N), 'float32'), B: T.Buffer((N,), 'float32'), Z: T.Buffer((M, N), 'float32')): + def linear(X: T.Tensor((M, K), 'float32'), W: T.Tensor((K, N), 'float32'), B: T.Tensor((N,), 'float32'), Z: T.Tensor((M, N), 'float32')): - Y = T.alloc_buffer((M, N), "float32") + Y = T.alloc_tensor((M, N), "float32") for i, j, k in T.grid(M, N, K): with Ts.sblock("Y"): v_i, v_j, v_k = Ts.axis.remap("SSR", [i, j, k]) @@ -153,7 +153,7 @@ for the end-to-end model execution. The code block below shows a TVMScript imple M, N = T.int64(), T.int64() @Ts.prim_func(private=True) - def relu(X: T.Buffer((M, N), 'float32'), Y: T.Buffer((M, N), 'float32')): + def relu(X: T.Tensor((M, N), 'float32'), Y: T.Tensor((M, N), 'float32')): diff --git a/docs/deep_dive/relax/tutorials/relax_creation.py b/docs/deep_dive/relax/tutorials/relax_creation.py index 917091f53599..0886fdfe77c8 100644 --- a/docs/deep_dive/relax/tutorials/relax_creation.py +++ b/docs/deep_dive/relax/tutorials/relax_creation.py @@ -77,7 +77,7 @@ def forward( @I.ir_module class RelaxModuleWithTIR: @Ts.prim_func - def relu(X: T.Buffer((n, m), "float32"), Y: T.Buffer((n, m), "float32")): + def relu(X: T.Tensor((n, m), "float32"), Y: T.Tensor((n, m), "float32")): for i, j in T.grid(n, m): with Ts.sblock("relu"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -170,10 +170,10 @@ def forward(self, x): @Ts.prim_func def tir_linear( - X: T.Buffer((M, K), "float32"), - W: T.Buffer((N, K), "float32"), - B: T.Buffer((N,), "float32"), - Z: T.Buffer((M, N), "float32"), + X: T.Tensor((M, K), "float32"), + W: T.Tensor((N, K), "float32"), + B: T.Tensor((N,), "float32"), + Z: T.Tensor((M, N), "float32"), ): for i, j, k in T.grid(M, N, K): with Ts.sblock("linear"): diff --git a/docs/deep_dive/tensor_ir/abstraction.rst b/docs/deep_dive/tensor_ir/abstraction.rst index bfd790150ff0..49d96264fcea 100644 --- a/docs/deep_dive/tensor_ir/abstraction.rst +++ b/docs/deep_dive/tensor_ir/abstraction.rst @@ -34,9 +34,9 @@ the compute statements themselves. @Ts.prim_func def main( - A: T.Buffer((128,), "float32"), - B: T.Buffer((128,), "float32"), - C: T.Buffer((128,), "float32"), + A: T.Tensor((128,), "float32"), + B: T.Tensor((128,), "float32"), + C: T.Tensor((128,), "float32"), ) -> None: for i in range(128): with Ts.sblock("C"): diff --git a/docs/deep_dive/tensor_ir/learning.rst b/docs/deep_dive/tensor_ir/learning.rst index 7d72d597c18d..821b427a99a0 100644 --- a/docs/deep_dive/tensor_ir/learning.rst +++ b/docs/deep_dive/tensor_ir/learning.rst @@ -65,10 +65,10 @@ language called TVMScript, which is a domain-specific dialect embedded in python @tvm.script.ir_module class MyModule: @Ts.prim_func - def mm_relu(A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32")): - Y = T.alloc_buffer((128, 128), dtype="float32") + def mm_relu(A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32")): + Y = T.alloc_tensor((128, 128), dtype="float32") for i, j, k in T.grid(128, 128, 128): with Ts.sblock("Y"): vi = Ts.axis.spatial(128, i) @@ -93,15 +93,15 @@ Function Parameters and Buffers .. code:: python # TensorIR - def mm_relu(A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32")): + def mm_relu(A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32")): ... # NumPy def lnumpy_mm_relu(A: np.ndarray, B: np.ndarray, C: np.ndarray): ... -Here ``A``, ``B``, and ``C`` takes a type named ``T.Buffer``, which with shape +Here ``A``, ``B``, and ``C`` takes a type named ``T.Tensor``, which with shape argument ``(128, 128)`` and data type ``float32``. This additional information helps possible MLC process to generate code that specializes in the shape and data type. @@ -111,7 +111,7 @@ type. .. code:: python # TensorIR - Y = T.alloc_buffer((128, 128), dtype="float32") + Y = T.alloc_tensor((128, 128), dtype="float32") # NumPy Y = np.empty((128, 128), dtype="float32") @@ -240,10 +240,10 @@ So we can also write the programs as follows. @tvm.script.ir_module class MyModuleWithAxisRemapSugar: @Ts.prim_func - def mm_relu(A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32")): - Y = T.alloc_buffer((128, 128), dtype="float32") + def mm_relu(A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32")): + Y = T.alloc_tensor((128, 128), dtype="float32") for i, j, k in T.grid(128, 128, 128): with Ts.sblock("Y"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) diff --git a/docs/deep_dive/tensor_ir/tutorials/tir_creation.py b/docs/deep_dive/tensor_ir/tutorials/tir_creation.py index d50f7531c2f9..f5059ec95071 100644 --- a/docs/deep_dive/tensor_ir/tutorials/tir_creation.py +++ b/docs/deep_dive/tensor_ir/tutorials/tir_creation.py @@ -65,11 +65,11 @@ class MyModule: @Ts.prim_func def mm_relu( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ): - Y = T.alloc_buffer((128, 128), dtype="float32") + Y = T.alloc_tensor((128, 128), dtype="float32") for i in range(128): for j in range(128): for k in range(128): @@ -108,11 +108,11 @@ def mm_relu( class ConciseModule: @Ts.prim_func def mm_relu( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ): - Y = T.alloc_buffer((128, 128), dtype="float32") + Y = T.alloc_tensor((128, 128), dtype="float32") for i, j, k in T.grid(128, 128, 128): with Ts.sblock("Y"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) @@ -147,11 +147,11 @@ def mm_relu( class ConciseModuleFromPython: @Ts.prim_func def mm_relu( - A: T.Buffer((M, K), dtype), - B: T.Buffer((K, N), dtype), - C: T.Buffer((M, N), dtype), + A: T.Tensor((M, K), dtype), + B: T.Tensor((K, N), dtype), + C: T.Tensor((M, N), dtype), ): - Y = T.alloc_buffer((M, N), dtype) + Y = T.alloc_tensor((M, N), dtype) for i, j, k in T.grid(M, N, K): with Ts.sblock("Y"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) @@ -185,10 +185,10 @@ def mm_relu( @I.ir_module class DynamicShapeModule: @Ts.prim_func - def mm_relu(A: T.Buffer([M, K], dtype), B: T.Buffer([K, N], dtype), C: T.Buffer([M, N], dtype)): + def mm_relu(A: T.Tensor([M, K], dtype), B: T.Tensor([K, N], dtype), C: T.Tensor([M, N], dtype)): # Bind the input buffers with the dynamic shapes - Y = T.alloc_buffer((M, N), dtype) + Y = T.alloc_tensor((M, N), dtype) for i, j, k in T.grid(M, N, K): with Ts.sblock("Y"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) diff --git a/docs/deep_dive/tensor_ir/tutorials/tir_transformation.py b/docs/deep_dive/tensor_ir/tutorials/tir_transformation.py index a4ff02669dcf..0ca72d51605a 100644 --- a/docs/deep_dive/tensor_ir/tutorials/tir_transformation.py +++ b/docs/deep_dive/tensor_ir/tutorials/tir_transformation.py @@ -46,9 +46,9 @@ class MyModule: @Ts.prim_func def main( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ): T.func_attr({"tirx.noalias": True}) with Ts.sblock("root"): 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 c95afd994c59..4fc4421a9a66 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 @@ -87,9 +87,9 @@ class MyFirstModule(BasePyModule): @Ts.prim_func def add_tir( - A: T.Buffer((4,), "float32"), - B: T.Buffer((4,), "float32"), - C: T.Buffer((4,), "float32"), + A: T.Tensor((4,), "float32"), + B: T.Tensor((4,), "float32"), + C: T.Tensor((4,), "float32"), ): for i in range(4): C[i] = A[i] + B[i] @@ -133,9 +133,9 @@ def forward(self, x, y): class DebugModule(BasePyModule): @Ts.prim_func def matmul_tir( - A: T.Buffer((n, 4), "float32"), - B: T.Buffer((4, 3), "float32"), - C: T.Buffer((n, 3), "float32"), + A: T.Tensor((n, 4), "float32"), + B: T.Tensor((4, 3), "float32"), + C: T.Tensor((n, 3), "float32"), ): for i, j, k in T.grid(n, 3, 4): with Ts.sblock("matmul"): @@ -210,9 +210,9 @@ def my_bias_add(x, bias, out): class PipelineModule(BasePyModule): @Ts.prim_func def matmul_tir( - A: T.Buffer((2, 4), "float32"), - B: T.Buffer((4, 3), "float32"), - C: T.Buffer((2, 3), "float32"), + A: T.Tensor((2, 4), "float32"), + B: T.Tensor((4, 3), "float32"), + C: T.Tensor((2, 3), "float32"), ): for i, j, k in T.grid(2, 3, 4): with Ts.sblock("matmul"): @@ -274,9 +274,9 @@ def forward(self, x, weights, bias): class DenseLayer: @Ts.prim_func def bias_add_tir( - x: T.Buffer((2, 4), "float32"), - b: T.Buffer((4,), "float32"), - out: T.Buffer((2, 4), "float32"), + x: T.Tensor((2, 4), "float32"), + b: T.Tensor((4,), "float32"), + out: T.Tensor((2, 4), "float32"), ): for i, j in T.grid(2, 4): out[i, j] = x[i, j] + b[j] @@ -401,7 +401,7 @@ def main( @R.py_module class DynamicModule(BasePyModule): @Ts.prim_func - def scale_tir(x: T.Buffer((n,), "float32"), out: T.Buffer((n,), "float32")): + def scale_tir(x: T.Tensor((n,), "float32"), out: T.Tensor((n,), "float32")): for i in T.serial(n): out[i] = x[i] * T.float32(2.0) diff --git a/docs/tirx/api/script.rst b/docs/tirx/api/script.rst index 5f0d50c76f3e..1589e0da32aa 100644 --- a/docs/tirx/api/script.rst +++ b/docs/tirx/api/script.rst @@ -22,7 +22,7 @@ TIRx kernels use ``tvm.script.tirx`` for the parser and core IR builders:: from tvm.script import tirx as Tx - Tx.alloc_buffer(...) + Tx.alloc_tensor(...) Tile primitives and backend-specific namespaces are documented separately in :doc:`tile`, :doc:`cuda`, and :doc:`ptx`. For the relationship between these diff --git a/docs/tirx/api/tirx.rst b/docs/tirx/api/tirx.rst index 51c9c1744b95..443b7d027f22 100644 --- a/docs/tirx/api/tirx.rst +++ b/docs/tirx/api/tirx.rst @@ -25,7 +25,7 @@ here so the same objects are not expanded twice. For C++ construction, include ``tvm/tirx/expr.h`` for ``BufferVar``, buffer loads, and buffer-region constructors. Include ``tvm/tirx/type.h`` for -``BufferType``, ``BufferRegionType``, and ``TensorMapType``. Buffer regions +``TensorType``, ``BufferRegionType``, and ``TensorMapType``. Buffer regions use the shared ``TensorRegion`` expression from ``tvm/ir/expr.h``. .. automodule:: tvm.tirx diff --git a/docs/tirx/arch/lowering_pipeline.rst b/docs/tirx/arch/lowering_pipeline.rst index b55a5d1b07b4..869d22d680ed 100644 --- a/docs/tirx/arch/lowering_pipeline.rst +++ b/docs/tirx/arch/lowering_pipeline.rst @@ -157,7 +157,7 @@ Take a one-line scale kernel: .. code-block:: python @Tx.prim_func - def scale(A: Tx.Buffer((256,), "float32"), B: Tx.Buffer((256,), "float32")): + def scale(A: Tx.Tensor((256,), "float32"), B: Tx.Tensor((256,), "float32")): Tx.device_entry() bx = Tx.cta_id([1]) diff --git a/docs/tirx/native_basics.rst b/docs/tirx/native_basics.rst index e30e3e7f89f5..ef349438ba91 100644 --- a/docs/tirx/native_basics.rst +++ b/docs/tirx/native_basics.rst @@ -43,7 +43,7 @@ The authoring model - ``@Tx.prim_func`` (or ``@Tx.jit`` for compile-time-specialized) kernels, written with ``from tvm.script import tirx as Tx``; - ``Tx.device_entry()`` plus *scope-id* intrinsics for thread binding; -- ``Tx.Buffer`` parameter annotations and ``Tx.alloc_*`` scratch buffers; +- ``Tx.Tensor`` parameter annotations and ``Tx.alloc_*`` scratch buffers; - ordinary loops, branches, and scalar math; - ``tvm.compile(mod, target=..., tir_pipeline="tirx")`` to build, then call the result directly. diff --git a/docs/tirx/native_basics/cuda/buffers.rst b/docs/tirx/native_basics/cuda/buffers.rst index 6b923e9316b1..9cabed75aff6 100644 --- a/docs/tirx/native_basics/cuda/buffers.rst +++ b/docs/tirx/native_basics/cuda/buffers.rst @@ -18,29 +18,31 @@ Buffers and memory ================== -Parameter buffers use ``Tx.Buffer`` signature annotations; scratch buffers are -created in the body with one of two declaration APIs (below). Index a buffer with +``Tx.Tensor`` constructs a ``tirx.TensorType`` for a tensor parameter. Its +shape, dtype, strides, offsets, layout and storage scope describe the same +low-level storage contract used by the allocation and declaration helpers. +Scratch tensors are created in the body with the APIs below. Index a tensor with ``A[i, j]``, slice it with ``A[m0:m0+BM, 0:BK]`` (a ``BufferRegion``), and take a pointer with ``A.ptr_to([i, j])`` or the raw data pointer ``A.data``. Declaring buffers ----------------- -Two fundamental APIs create a buffer: +Two fundamental APIs create a tensor variable: -- ``Tx.alloc_buffer(shape, dtype, scope=..., ...)`` — **allocates new storage** - (emits an ``AllocBuffer`` node) and returns the ``Buffer``. ``Tx.alloc_shared`` / - ``Tx.alloc_local`` are just ``alloc_buffer`` with ``scope="shared"`` / +- ``Tx.alloc_tensor(shape, dtype, scope=..., ...)`` — **allocates new storage** + (binds a ``tirx.alloc_tensor`` call) and returns an ``ir.Var``. ``Tx.alloc_shared`` / + ``Tx.alloc_local`` are just ``alloc_tensor`` with ``scope="shared"`` / ``scope="local"``. -- ``Tx.decl_buffer(shape, dtype, data=..., ...)`` — **declares a view** over an +- ``Tx.decl_tensor(shape, dtype, data=..., ...)`` — **declares a view** over an existing pointer ``data`` (no allocation); use it to alias or reinterpret storage — a sub-region of a pool, or a tensor-memory address. With ``data=None`` - it allocates, like ``alloc_buffer``. + it allocates, like ``alloc_tensor``, except in ``tmem`` scope, where + ``allocated_addr`` identifies externally allocated tensor memory. -A buffer's ``data`` pointer is an immutable ``Var`` (``alloc_buffer`` defines it; -``decl_buffer`` takes one). To back a buffer with a pointer *expression*, assign -it to a name first; the parser creates an immutable pointer binding. See -:doc:`data_types`. +``A.data`` projects the physical pointer from the tensor variable. +``alloc_tensor`` supplies new storage; ``decl_tensor`` with ``data`` binds an +existing pointer expression. See :doc:`data_types`. Both share one descriptor; the parameters that matter most: @@ -91,10 +93,10 @@ The ``scope`` argument selects the memory space: .. code-block:: python @Tx.prim_func - def kernel(A: Tx.Buffer((M, K), "float16", align=16)): + def kernel(A: Tx.Tensor((M, K), "float16", align=16)): As = Tx.alloc_shared((BM, BK), "float16") # new shared tile acc = Tx.alloc_local((4,), "float32") # per-thread accumulator - view = Tx.decl_buffer((BM, BK), "float16", data=As.data) # a view over As + view = Tx.decl_tensor((BM, BK), "float16", data=As.data) # a view over As **A ptr-based buffer is just metadata over a pointer.** For any non-tmem buffer, the declaration is a pointer plus a layout, and indexing resolves to an address:: @@ -111,10 +113,10 @@ the function signature: from tvm.tirx.layout import TileLayout, S - B: Tx.Buffer((4, 8), "float32") # row-major - B: Tx.Buffer((4, 8), "float32", layout=TileLayout(S[(4, 8) : (1, 4)])) # column-major - B: Tx.Buffer((4, 8), "float32", elem_offset=64) # shifted view - B: Tx.Buffer((4, 8), "float32", layout=TileLayout(S[(4, 8) : (16, 1)])) # row stride 16 + B: Tx.Tensor((4, 8), "float32") # row-major + B: Tx.Tensor((4, 8), "float32", layout=TileLayout(S[(4, 8) : (1, 4)])) # column-major + B: Tx.Tensor((4, 8), "float32", elem_offset=64) # shifted view + B: Tx.Tensor((4, 8), "float32", layout=TileLayout(S[(4, 8) : (16, 1)])) # row stride 16 each makes ``B[i, j]`` lower to a different index in the generated CUDA (the ``A[i, j]`` load stays ``i*8 + j`` — only ``B``'s metadata changed): @@ -142,7 +144,7 @@ whole block sees the writes, then read it back: .. code-block:: python @Tx.prim_func - def smem_demo(A: Tx.Buffer((128,), "float32"), B: Tx.Buffer((128,), "float32")): + def smem_demo(A: Tx.Tensor((128,), "float32"), B: Tx.Tensor((128,), "float32")): Tx.device_entry() bx = Tx.cta_id([1]) @@ -171,14 +173,14 @@ Dynamic **Dynamic** shared memory (``scope="shared.dyn"``) is sized per launch (the ``sharedMemBytes`` launch parameter), not at compile time. A kernel may have **only one** dynamic-shared allocation — the *arena*. So you allocate it once and ``decl`` -each buffer as a view into it: ``Tx.decl_buffer`` with ``data=`` the arena pointer +each buffer as a view into it: ``Tx.decl_tensor`` with ``data=`` the arena pointer and an ``elem_offset``: .. code-block:: python - arena = Tx.alloc_buffer((128,), "float32", scope="shared.dyn") # the one arena - As = Tx.decl_buffer((64,), "float32", data=arena.data, scope="shared.dyn") # offset 0 - Bs = Tx.decl_buffer((64,), "float32", data=arena.data, elem_offset=64, scope="shared.dyn") # offset 64 + arena = Tx.alloc_tensor((128,), "float32", scope="shared.dyn") # the one arena + As = Tx.decl_tensor((64,), "float32", data=arena.data, scope="shared.dyn") # offset 0 + Bs = Tx.decl_tensor((64,), "float32", data=arena.data, elem_offset=64, scope="shared.dyn") # offset 64 As[tx] = A[tx] Bs[tx] = B[tx] Tx.cuda.cta_sync() @@ -195,7 +197,7 @@ boilerplate elided; arena named ``smem`` for clarity): __syncthreads(); C_ptr[tx] = smem[tx] + smem[tx + 64]; -(Two separate ``alloc_buffer(scope="shared.dyn")`` is an error — *only one dynamic +(Two separate ``alloc_tensor(scope="shared.dyn")`` is an error — *only one dynamic shared memory allocation is allowed*.) So static shared memory is sized at compile time (``__shared__ T x[N];``); dynamic shared memory is this one launch-sized arena with views decl'd at offsets inside it. @@ -356,10 +358,10 @@ Tensor memory Blackwell *tensor memory* is not a plain scratch scope: it must be explicitly reserved and freed with the warp-uniform ``Tx.ptx.tcgen05.alloc`` / ``tcgen05.dealloc`` intrinsics, and each tensor is a view into it declared with -``Tx.decl_buffer(..., scope="tmem", allocated_addr=
, layout=)``. +``Tx.decl_tensor(..., scope="tmem", allocated_addr=
, layout=)``. The ``allocated_addr`` is the allocated tensor-memory base address plus any desired column offset. It is mandatory — the tensor-core dispatch asserts it — so -``Tx.alloc_buffer(scope="tmem")`` (which does **not** set it) will not work. Unlike +``Tx.alloc_tensor(scope="tmem")`` (which does **not** set it) will not work. Unlike shared memory, tensor memory is not directly addressable: it is read and written only through ``tcgen05`` ``mma`` / ``ld`` / ``st`` / ``cp``. @@ -372,7 +374,7 @@ tensor as a view at a column offset, and one warp frees it at the end: if warp_id == alloc_warp: # tcgen05.alloc is warp-uniform Tx.ptx[f"tcgen05.alloc.cta_group::{cta_group}.sync.aligned.shared::cta.b32"]( Tx.address_of(addr), Tx.uint32(512)) - acc = Tx.decl_buffer((CTA_M, 512), "float32", scope="tmem", + acc = Tx.decl_tensor((CTA_M, 512), "float32", scope="tmem", allocated_addr=addr[0], layout=tmem_layout) # allocated base # ... use acc as a gemm_async / copy_async operand ... if warp_id == alloc_warp: @@ -491,7 +493,7 @@ no shape infers a one-dimensional shape from the logical storage size: .. code-block:: python - R = Tx.alloc_buffer((32, 8), "float32", scope="local", layout=TileLayout(S[(32, 8) : (1 @ laneid, 1)])) + R = Tx.alloc_tensor((32, 8), "float32", scope="local", layout=TileLayout(S[(32, 8) : (1 @ laneid, 1)])) R_flat = R.local() # this lane's 8 local elements, physical order R_2d = R.local(2, 4) # the same elements, row-major 2x4 reshape diff --git a/docs/tirx/native_basics/cuda/compiling.rst b/docs/tirx/native_basics/cuda/compiling.rst index bbd11c4b363c..03773635224f 100644 --- a/docs/tirx/native_basics/cuda/compiling.rst +++ b/docs/tirx/native_basics/cuda/compiling.rst @@ -76,7 +76,7 @@ Rung 2 in full — a 256-element block sum via a shared-memory tree reduction .. code-block:: python @Tx.prim_func - def block_sum(A: Tx.Buffer((256,), "float32"), out: Tx.Buffer((1,), "float32")): + def block_sum(A: Tx.Tensor((256,), "float32"), out: Tx.Tensor((1,), "float32")): Tx.device_entry() bx = Tx.cta_id([1]) diff --git a/docs/tirx/native_basics/cuda/data_types.rst b/docs/tirx/native_basics/cuda/data_types.rst index d8ad1ef3409a..fadb1ecd399c 100644 --- a/docs/tirx/native_basics/cuda/data_types.rst +++ b/docs/tirx/native_basics/cuda/data_types.rst @@ -32,7 +32,7 @@ shared buffers across several dtypes, plus a vectorized ``float32x4`` load/store .. code-block:: python @Tx.prim_func - def dtypes(A: Tx.Buffer((256,), "float32"), O: Tx.Buffer((256,), "float32")): + def dtypes(A: Tx.Tensor((256,), "float32"), O: Tx.Tensor((256,), "float32")): Tx.device_entry() bx = Tx.cta_id([1]) @@ -100,8 +100,8 @@ Pointers (``handle``) A buffer's ``data`` — its pointer — is a ``Var`` of pointer type, and it is **immutable** (a pointer is never reassigned). That shapes how you obtain one: -- ``Tx.alloc_buffer(...)`` allocates storage **and** defines its ``data`` pointer. -- ``Tx.decl_buffer(..., data=ptr)`` declares a buffer over an existing pointer +- ``Tx.alloc_tensor(...)`` allocates storage **and** defines its ``data`` pointer. +- ``Tx.decl_tensor(..., data=ptr)`` declares a buffer over an existing pointer ``Var`` ``ptr``. - To back a buffer with a pointer **expression** — e.g. ``Tx.ptx.mapa`` giving another cluster CTA's shared address — convert the ``uint64`` address the @@ -118,7 +118,7 @@ A buffer's ``data`` — its pointer — is a ``Var`` of pointer type, and it is Tx.ptx.mapa.u64(mapped[0], mbar.ptr_to([0]), Tx.uint32(0)) ptr_ty = PointerType(PrimType("uint64"), "shared") ptr = Tx.reinterpret(ptr_ty, mapped[0]) - remote_mbar = Tx.decl_buffer([1], "uint64", data=ptr, scope="shared") + remote_mbar = Tx.decl_tensor([1], "uint64", data=ptr, scope="shared") Pointer bindings cannot be reassigned; use a new name for a different pointer value. diff --git a/docs/tirx/native_basics/cuda/first_kernel.rst b/docs/tirx/native_basics/cuda/first_kernel.rst index c0550f08e13e..7af02da34ab5 100644 --- a/docs/tirx/native_basics/cuda/first_kernel.rst +++ b/docs/tirx/native_basics/cuda/first_kernel.rst @@ -29,7 +29,7 @@ with one block of 256 threads. @Tx.prim_func - def scale(A: Tx.Buffer((256,), "float32"), B: Tx.Buffer((256,), "float32")): + def scale(A: Tx.Tensor((256,), "float32"), B: Tx.Tensor((256,), "float32")): Tx.device_entry() # everything below runs on the device bx = Tx.cta_id([1]) # 1 block (blockIdx) diff --git a/docs/tirx/native_basics/cuda/functions.rst b/docs/tirx/native_basics/cuda/functions.rst index d73ae924ea7f..661d5877c4f9 100644 --- a/docs/tirx/native_basics/cuda/functions.rst +++ b/docs/tirx/native_basics/cuda/functions.rst @@ -26,13 +26,13 @@ pass, symbolic shapes, and the ``prim_func`` / ``jit`` distinction. Declaring buffer parameters --------------------------- -Declare tensor parameters with ``Tx.Buffer`` annotations. The annotation accepts +Declare tensor parameters with ``Tx.Tensor`` annotations. The annotation accepts shape, dtype, layout, offset, scope, and alignment metadata: .. code-block:: python @Tx.prim_func - def f(A: Tx.Buffer((256,), "float32", align=16), B: Tx.Buffer((256,), "float32")): ... + def f(A: Tx.Tensor((256,), "float32", align=16), B: Tx.Tensor((256,), "float32")): ... The parameters are buffers that you index with ``A[i]`` or ``A[i, j]``. Annotations also support :ref:`symbolic shapes `. @@ -50,7 +50,7 @@ pass on the Python side when you call the compiled ``Executable``: * - Annotation - Is - Pass at call time - * - ``Tx.Buffer((d0, d1), dtype)`` + * - ``Tx.Tensor((d0, d1), dtype)`` - a tensor parameter (shape + dtype fixed) - a tensor on the right device * - ``Tx.handle`` @@ -69,7 +69,7 @@ interop) or order. For example, a kernel with a scalar parameter:: @Tx.prim_func - def scal(A: Tx.Buffer((256,), 'float32'), B: Tx.Buffer((256,), 'float32'), s: Tx.float32): + def scal(A: Tx.Tensor((256,), 'float32'), B: Tx.Tensor((256,), 'float32'), s: Tx.float32): Tx.device_entry(); bx = Tx.cta_id([1]); tx = Tx.thread_id([256]) @@ -92,7 +92,7 @@ passed tensor** at run time, so a *single compiled kernel* handles any size: @Tx.prim_func - def scale_dyn(A: Tx.Buffer((n,), "float32"), B: Tx.Buffer((n,), "float32")): + def scale_dyn(A: Tx.Tensor((n,), "float32"), B: Tx.Tensor((n,), "float32")): Tx.device_entry() bx = Tx.cta_id([1]) tx = Tx.thread_id([1]) @@ -144,7 +144,7 @@ merged function (trimmed): @Tx.prim_func - def main(A: Tx.Buffer((n,)), B: Tx.Buffer((n,))): + def main(A: Tx.Tensor((n,)), B: Tx.Tensor((n,))): with Tx.launch_thread("blockIdx.x", 1), Tx.launch_thread("threadIdx.x", 1): for i in range(n): @@ -167,7 +167,7 @@ trailing ``1, 1`` are the grid/block launch dims): @Tx.prim_func # host - def main(A: Tx.Buffer((n,)), B: Tx.Buffer((n,))): + def main(A: Tx.Tensor((n,)), B: Tx.Tensor((n,))): Tx.call_packed("scale_dyn_kernel", A.data, B.data, n, 1, 1) # n forwarded @@ -186,7 +186,7 @@ device checks (e.g. asserting ``B.shape[0] == n``):: parameters annotated ``Tx.constexpr`` are baked in as compile-time constants and the result is an ordinary ``PrimFunc``. Use it when you want sizes/flags fixed at compile time (so the compiler can unroll, statically size shared memory, etc.). - Referencing a constexpr inside an annotation (e.g. ``Tx.Buffer((N,), ...)``) + Referencing a constexpr inside an annotation (e.g. ``Tx.Tensor((N,), ...)``) requires ``from __future__ import annotations`` at the top of the file. .. code-block:: python @@ -196,9 +196,9 @@ device checks (e.g. asserting ``B.shape[0] == n``):: @Tx.jit def add( - A: Tx.Buffer((N,), "float32"), - B: Tx.Buffer((N,), "float32"), - C: Tx.Buffer((N,), "float32"), + A: Tx.Tensor((N,), "float32"), + B: Tx.Tensor((N,), "float32"), + C: Tx.Tensor((N,), "float32"), *, N: Tx.constexpr, ): diff --git a/docs/tirx/native_basics/cuda/parser_utils.rst b/docs/tirx/native_basics/cuda/parser_utils.rst index 61ddbe53217b..9ecad32404b7 100644 --- a/docs/tirx/native_basics/cuda/parser_utils.rst +++ b/docs/tirx/native_basics/cuda/parser_utils.rst @@ -64,7 +64,7 @@ allocations and state into one object and use it in the kernel body. class State: def __init__(self, smem): self.acc = Tx.alloc_local([1], "float32") - self.buf = Tx.decl_buffer([64], "float16", smem, scope="shared.dyn") + self.buf = Tx.decl_tensor([64], "float16", smem, scope="shared.dyn") s = State(smem.data) s.acc[0] = Tx.float32(0.0) # use its fields like ordinary buffers diff --git a/docs/tirx/native_basics/cuda/profiling.rst b/docs/tirx/native_basics/cuda/profiling.rst index 2553b00268cf..377d008aba77 100644 --- a/docs/tirx/native_basics/cuda/profiling.rst +++ b/docs/tirx/native_basics/cuda/profiling.rst @@ -68,9 +68,9 @@ are a plain ``enum.Enum`` whose integer values start at 0 and index a names list @Tx.prim_func def profiled_kernel( - out: Tx.Buffer((N,), "float32"), - inp: Tx.Buffer((N,), "float32"), - prof: Tx.Buffer((PROF_SIZE,), "uint64"), + out: Tx.Tensor((N,), "float32"), + inp: Tx.Tensor((N,), "float32"), + prof: Tx.Tensor((PROF_SIZE,), "uint64"), ): Tx.device_entry() diff --git a/docs/tirx/native_basics/cuda/threads_sync.rst b/docs/tirx/native_basics/cuda/threads_sync.rst index 477de34025f4..be9b4122933f 100644 --- a/docs/tirx/native_basics/cuda/threads_sync.rst +++ b/docs/tirx/native_basics/cuda/threads_sync.rst @@ -46,7 +46,7 @@ A complete, runnable example — a warp all-reduce via ``Tx.tvm_warp_shuffle_xor .. code-block:: python @Tx.prim_func - def warp_reduce(A: Tx.Buffer((32,), "float32", align=16)): + def warp_reduce(A: Tx.Tensor((32,), "float32", align=16)): Tx.device_entry() cta_id = Tx.cta_id([1]) @@ -86,7 +86,7 @@ source string with ``Tx.cuda.func_call(name, *args, source_code=..., return_type @Tx.prim_func - def k(A: Tx.Buffer((256,), "float32"), B: Tx.Buffer((256,), "float32")): + def k(A: Tx.Tensor((256,), "float32"), B: Tx.Tensor((256,), "float32")): Tx.device_entry() bx = Tx.cta_id([1]) diff --git a/docs/tirx/tile_primitives/copy/fallback.rst b/docs/tirx/tile_primitives/copy/fallback.rst index 0659aeabefdc..eb325f5a8a56 100644 --- a/docs/tirx/tile_primitives/copy/fallback.rst +++ b/docs/tirx/tile_primitives/copy/fallback.rst @@ -77,13 +77,13 @@ divisible by ``32``, so this falls through to ``fallback`` (from @Tx.prim_func - def kernel(A: Tx.Buffer(shape, dtype), B: Tx.Buffer(shape, dtype)): + def kernel(A: Tx.Tensor(shape, dtype), B: Tx.Tensor(shape, dtype)): Tx.device_entry() Tx.cta_id([1]) Tx.lane_id([32]) Tx.thread_id([32]) - A_smem = Tx.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = Tx.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) Tx.tile.warp.copy(A_smem[full], A[full]) # fallback Tx.cuda.cta_sync() Tx.tile.warp.copy(B[full], A_smem[full]) # fallback diff --git a/docs/tirx/tile_primitives/copy/gmem_smem.rst b/docs/tirx/tile_primitives/copy/gmem_smem.rst index ae7a13e1d4bc..99e363355206 100644 --- a/docs/tirx/tile_primitives/copy/gmem_smem.rst +++ b/docs/tirx/tile_primitives/copy/gmem_smem.rst @@ -92,13 +92,13 @@ A warp (32 threads) copies a ``32×32`` ``float32`` tile global → shared and b @Tx.prim_func - def kernel(A: Tx.Buffer(shape, dtype), B: Tx.Buffer(shape, dtype)): + def kernel(A: Tx.Tensor(shape, dtype), B: Tx.Tensor(shape, dtype)): Tx.device_entry() Tx.cta_id([1]) Tx.lane_id([32]) Tx.thread_id([32]) - A_smem = Tx.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = Tx.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) Tx.tile.warp.copy(A_smem[fs], A[fs]) # global -> shared (this dispatch) Tx.cuda.cta_sync() Tx.tile.warp.copy(B[fs], A_smem[fs]) # shared -> global (this dispatch) diff --git a/docs/tirx/tile_primitives/copy/ldstmatrix.rst b/docs/tirx/tile_primitives/copy/ldstmatrix.rst index 0fec611bf50f..f4eeb54b904c 100644 --- a/docs/tirx/tile_primitives/copy/ldstmatrix.rst +++ b/docs/tirx/tile_primitives/copy/ldstmatrix.rst @@ -113,16 +113,16 @@ register, from ``test_ld_stmatrix.py`` (register layout = the m8n8 fragment, @Tx.prim_func - def kernel(A: Tx.Buffer((M, N), "float16"), B: Tx.Buffer((M, N), "float16")): + def kernel(A: Tx.Tensor((M, N), "float16"), B: Tx.Tensor((M, N), "float16")): Tx.device_entry() Tx.cta_id([1]) Tx.lane_id([32]) tid = Tx.thread_id([32]) - A_smem = Tx.alloc_buffer((8, 4, num, 2), "float16", scope="shared", layout=s_layout) + A_smem = Tx.alloc_tensor((8, 4, num, 2), "float16", scope="shared", layout=s_layout) # ... stage A into A_smem (row = tid//4, cp = tid%4) ... Tx.cuda.cta_sync() - R = Tx.alloc_buffer((8, 4, num, 2), "float16", scope="local", layout=r_layout) + R = Tx.alloc_tensor((8, 4, num, 2), "float16", scope="local", layout=r_layout) Tx.tile.warp.copy(R[full], A_smem[full]) # shared -> register (ldmatrix) # ... write R back out to B ... diff --git a/docs/tirx/tile_primitives/copy/reg.rst b/docs/tirx/tile_primitives/copy/reg.rst index 5e0e76d334eb..0480b458bfdc 100644 --- a/docs/tirx/tile_primitives/copy/reg.rst +++ b/docs/tirx/tile_primitives/copy/reg.rst @@ -90,17 +90,17 @@ contiguous elements). From ``test_reg.py``: @Tx.prim_func - def kernel(B: Tx.Buffer(shape, dtype)): + def kernel(B: Tx.Tensor(shape, dtype)): Tx.device_entry() Tx.cta_id([1]) Tx.lane_id([32]) tid = Tx.thread_id([32]) - A_smem = Tx.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = Tx.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) for kk in range(8): A_smem[tid, kk] = Tx.cast(tid * 100 + kk + 1, dtype) Tx.cuda.cta_sync() - R = Tx.alloc_buffer(shape, dtype, scope="local", layout=r_layout) + R = Tx.alloc_tensor(shape, dtype, scope="local", layout=r_layout) Tx.tile.warp.copy(R[fs], A_smem[fs]) # shared -> register (this dispatch) # ... clear A_smem, cta_sync ... Tx.tile.warp.copy(A_smem[fs], R[fs]) # register -> shared (this dispatch) @@ -153,7 +153,7 @@ Generated TIRx IR .. code-block:: python - r_local = Tx.decl_buffer((8,), data=R.data, scope="local") # 8 fp32 elements / lane + r_local = Tx.decl_tensor((8,), data=R.data, scope="local") # 8 fp32 elements / lane r_words = r_local.view("uint32") for f in range(2): # outer = 8 / vec 4 s_ptr = pointer_offset(A_smem, ...) # this lane's row diff --git a/docs/tirx/tile_primitives/copy_async/dsmem.rst b/docs/tirx/tile_primitives/copy_async/dsmem.rst index 70003290633c..091691d6be6b 100644 --- a/docs/tirx/tile_primitives/copy_async/dsmem.rst +++ b/docs/tirx/tile_primitives/copy_async/dsmem.rst @@ -84,7 +84,7 @@ mbarrier and writes the result out (from ``test_dsmem.py``): @Tx.prim_func - def dsmem_copy(A: Tx.Buffer(shape, dtype), B: Tx.Buffer(shape, dtype)): + def dsmem_copy(A: Tx.Tensor(shape, dtype), B: Tx.Tensor(shape, dtype)): Tx.device_entry() cbx = Tx.cta_id_in_cluster([CLUSTER_N]) @@ -92,7 +92,7 @@ mbarrier and writes the result out (from ``test_dsmem.py``): tid = Tx.thread_id([1]) pool = Tx.SMEMPool() src_raw = pool.alloc([8192], dtype, align=128) - src_smem = Tx.decl_buffer( + src_smem = Tx.decl_tensor( list(shape), dtype, src_raw.data, @@ -101,7 +101,7 @@ mbarrier and writes the result out (from ``test_dsmem.py``): layout=src_layout, ) dst_raw = pool.alloc([8192], dtype, align=128) - dst_smem = Tx.decl_buffer( + dst_smem = Tx.decl_tensor( list(shape), dtype, dst_raw.data, diff --git a/docs/tirx/tile_primitives/copy_async/ldgsts.rst b/docs/tirx/tile_primitives/copy_async/ldgsts.rst index 8c511350a322..47d21155ac7c 100644 --- a/docs/tirx/tile_primitives/copy_async/ldgsts.rst +++ b/docs/tirx/tile_primitives/copy_async/ldgsts.rst @@ -97,14 +97,14 @@ shared, then commits and waits before reading it back (from ``test_ldgsts.py``): @Tx.prim_func - def copy_async(A: Tx.Buffer(shape, dtype), B: Tx.Buffer(shape, dtype)): + def copy_async(A: Tx.Tensor(shape, dtype), B: Tx.Tensor(shape, dtype)): Tx.device_entry() Tx.cta_id([1]) Tx.warp_id([4]) Tx.lane_id([32]) tid = Tx.thread_id([128]) - A_smem = Tx.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = Tx.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) Tx.tile.cta.copy_async(A_smem[full], A[full], dispatch="ldgsts") # async global -> shared Tx.ptx.cp.async_.commit_group() # caller commits ... Tx.ptx.cp.async_.wait_group(0) # ... and waits diff --git a/docs/tirx/tile_primitives/copy_async/tcgen05_cp.rst b/docs/tirx/tile_primitives/copy_async/tcgen05_cp.rst index fc67c8d7ba40..bdcf11843079 100644 --- a/docs/tirx/tile_primitives/copy_async/tcgen05_cp.rst +++ b/docs/tirx/tile_primitives/copy_async/tcgen05_cp.rst @@ -124,7 +124,7 @@ dealloc tail elided): from tvm.tirx.layout import R, S, TCol, TileLayout, TLane - A_smem = Tx.alloc_buffer([32, 16], "uint8", scope="shared", + A_smem = Tx.alloc_tensor([32, 16], "uint8", scope="shared", layout=TileLayout(S[(32, 16) : (16, 1)]), align=1024) tmem_addr = Tx.alloc_shared([1], "uint32") cp_mbar = Tx.alloc_shared([1], "uint64") @@ -132,7 +132,7 @@ dealloc tail elided): Tx.ptx["tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32"]( Tx.address_of(tmem_addr), Tx.uint32(16)) # ... mbarrier.init, fence, cta_sync, fill A_smem from global ... - tmem = Tx.decl_buffer([32, 16], "uint8", scope="tmem", allocated_addr=tmem_addr[0], + tmem = Tx.decl_tensor([32, 16], "uint8", scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(32, 16) : (1 @ TLane, 1 @ TCol)] + R[4 : 32 @ TLane])) if tid_in_wg == 0: Tx.tile.copy_async(tmem[0:32, 0:16], A_smem[0:32, 0:16], cta_group=1) # smem -> tmem diff --git a/docs/tirx/tile_primitives/copy_async/tcgen05_ldst.rst b/docs/tirx/tile_primitives/copy_async/tcgen05_ldst.rst index 3188cccddcbe..363ecf47db90 100644 --- a/docs/tirx/tile_primitives/copy_async/tcgen05_ldst.rst +++ b/docs/tirx/tile_primitives/copy_async/tcgen05_ldst.rst @@ -86,7 +86,7 @@ fp16): @Tx.prim_func def copy_async_test( - A: Tx.Buffer((128, WIDTH), "float16"), B: Tx.Buffer((128, WIDTH), "float16") + A: Tx.Tensor((128, WIDTH), "float16"), B: Tx.Tensor((128, WIDTH), "float16") ): Tx.device_entry() @@ -100,7 +100,7 @@ fp16): Tx.address_of(tmem_addr), Tx.uint32(32) ) Tx.tvm_storage_sync("shared") - tmem = Tx.decl_buffer( + tmem = Tx.decl_tensor( (128, WIDTH), "float16", scope="tmem", @@ -169,14 +169,14 @@ Selecting the upper F sub-slab from tvm.tirx.layout import tmem_datapath_layout - lower = Tx.decl_buffer( + lower = Tx.decl_tensor( (64, cols), "float32", scope="tmem", allocated_addr=tmem_addr[0], layout=tmem_datapath_layout("F", 64, cols, sub_slab=0), ) - upper = Tx.decl_buffer( + upper = Tx.decl_tensor( (64, cols), "float32", scope="tmem", diff --git a/docs/tirx/tile_primitives/elementwise/reg.rst b/docs/tirx/tile_primitives/elementwise/reg.rst index f106ca20fcba..7f746c748cd5 100644 --- a/docs/tirx/tile_primitives/elementwise/reg.rst +++ b/docs/tirx/tile_primitives/elementwise/reg.rst @@ -86,16 +86,16 @@ A warp takes the elementwise ``sqrt`` of a ``32×8`` ``float32`` local tile @Tx.prim_func - def k(A: Tx.Buffer((32, 8), "float32"), B: Tx.Buffer((32, 8), "float32")): + def k(A: Tx.Tensor((32, 8), "float32"), B: Tx.Tensor((32, 8), "float32")): Tx.device_entry() Tx.cta_id([1]) Tx.lane_id([32]) tid = Tx.thread_id([32]) - A_smem = Tx.alloc_buffer((32, 8), "float32", scope="shared", layout=TileLayout(S[(32, 8)])) + A_smem = Tx.alloc_tensor((32, 8), "float32", scope="shared", layout=TileLayout(S[(32, 8)])) Tx.tile.warp.copy(A_smem[fs], A[fs]) Tx.cuda.cta_sync() - R = Tx.alloc_buffer((32, 8), "float32", scope="local", layout=r_layout) + R = Tx.alloc_tensor((32, 8), "float32", scope="local", layout=r_layout) Tx.tile.warp.copy(R[fs], A_smem[fs]) Tx.tile.warp.sqrt(R[fs], R[fs]) # elementwise reg dispatch Tx.tile.warp.copy(A_smem[fs], R[fs]) diff --git a/docs/tirx/tile_primitives/elementwise/smem.rst b/docs/tirx/tile_primitives/elementwise/smem.rst index cdb43e5dd9ec..646769cdbd61 100644 --- a/docs/tirx/tile_primitives/elementwise/smem.rst +++ b/docs/tirx/tile_primitives/elementwise/smem.rst @@ -80,14 +80,14 @@ round): @Tx.prim_func - def unary_op(A: Tx.Buffer((32, 32), "float32", layout=s_layout)): + def unary_op(A: Tx.Tensor((32, 32), "float32", layout=s_layout)): Tx.device_entry() Tx.cta_id([1]) Tx.warp_id([8]) Tx.lane_id([32]) Tx.thread_id([256]) - A_smem = Tx.alloc_buffer((32, 32), "float32", scope="shared", layout=s_layout) + A_smem = Tx.alloc_tensor((32, 32), "float32", scope="shared", layout=s_layout) Tx.tile.cta.copy(A_smem[full], A[full]) Tx.tile.cta.sqrt(A_smem[full], A_smem[full]) # elementwise smem dispatch Tx.tile.cta.copy(A[full], A_smem[full]) diff --git a/docs/tirx/tile_primitives/gemm.rst b/docs/tirx/tile_primitives/gemm.rst index ff8df19d2435..c04383ecaaae 100644 --- a/docs/tirx/tile_primitives/gemm.rst +++ b/docs/tirx/tile_primitives/gemm.rst @@ -91,18 +91,18 @@ accumulate) — one ``m16n8k16`` atom (from ``test_gemm_mma_m16n8k_.py``): @Tx.prim_func def gemm( - A_g: Tx.Buffer((16, 16), "float16"), - B_g: Tx.Buffer((16, 8), "float16"), - D_g: Tx.Buffer((16, 8), "float32"), + A_g: Tx.Tensor((16, 16), "float16"), + B_g: Tx.Tensor((16, 8), "float16"), + D_g: Tx.Tensor((16, 8), "float32"), ): Tx.device_entry() Tx.cta_id([1]) Tx.warp_id([1]) lane = Tx.lane_id([32]) - A_f = Tx.alloc_buffer((16, 16), "float16", scope="local", layout=A_FRAG) - B_f = Tx.alloc_buffer((16, 8), "float16", scope="local", layout=B_FRAG) - D_f = Tx.alloc_buffer((16, 8), "float32", scope="local", layout=D_FRAG) + A_f = Tx.alloc_tensor((16, 16), "float16", scope="local", layout=A_FRAG) + B_f = Tx.alloc_tensor((16, 8), "float16", scope="local", layout=B_FRAG) + D_f = Tx.alloc_tensor((16, 8), "float32", scope="local", layout=D_FRAG) A_reg = A_f.local(8) # stage A into the lane's 8 regs for s in Tx.unroll(8): kp, rM, kHi = s % 2, (s // 2) % 2, s // 4 diff --git a/docs/tirx/tile_primitives/gemm_async.rst b/docs/tirx/tile_primitives/gemm_async.rst index f3ea96f538d6..b601fb988d81 100644 --- a/docs/tirx/tile_primitives/gemm_async.rst +++ b/docs/tirx/tile_primitives/gemm_async.rst @@ -103,15 +103,15 @@ into shared (from from tvm.tirx.layout import S, TCol, TLane, TileLayout, tid_in_wg as axis_tid_in_wg from tvm.backend.cuda.tile_primitive.tma_utils import mma_shared_layout - A_smem = Tx.alloc_buffer((3,128,64), "float16", scope="shared", layout=mma_shared_layout("float16", 3, (3,128,64))) - B_smem = Tx.alloc_buffer((3,128,64), "float16", scope="shared", layout=mma_shared_layout("float16", 3, (3,128,64))) + A_smem = Tx.alloc_tensor((3,128,64), "float16", scope="shared", layout=mma_shared_layout("float16", 3, (3,128,64))) + B_smem = Tx.alloc_tensor((3,128,64), "float16", scope="shared", layout=mma_shared_layout("float16", 3, (3,128,64))) tmem_addr = Tx.alloc_shared([1], "uint32"); mma_mbar = Tx.alloc_shared([1], "uint64") # ... mbarrier.init, cta_sync ... if warp_id == 0: Tx.ptx["tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32"]( Tx.address_of(tmem_addr), Tx.uint32(512)) Tx.cuda.cta_sync() - tmem = Tx.decl_buffer((128, 512), "float32", scope="tmem", allocated_addr=tmem_addr[0], + tmem = Tx.decl_tensor((128, 512), "float32", scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, 512) : (1 @ TLane, 1 @ TCol)])) # ... TMA-load A_smem, B_smem from global, wait ... if tid_in_wg == 0: diff --git a/docs/tirx/tile_primitives/permute_layout.rst b/docs/tirx/tile_primitives/permute_layout.rst index fdaecd699272..4ebbd32d133b 100644 --- a/docs/tirx/tile_primitives/permute_layout.rst +++ b/docs/tirx/tile_primitives/permute_layout.rst @@ -87,7 +87,7 @@ canonical SF-transpose, from ``test_permute_layout.py``): @Tx.prim_func - def f(A_buf: Tx.Buffer(shape, dtype, layout=pre), B_buf: Tx.Buffer(shape, dtype, layout=post)): + def f(A_buf: Tx.Tensor(shape, dtype, layout=pre), B_buf: Tx.Tensor(shape, dtype, layout=post)): Tx.device_entry() Tx.cta_id([1]) @@ -117,7 +117,7 @@ destination layout: .. code-block:: python - regs = Tx.alloc_buffer((P,), dtype, scope="local") + regs = Tx.alloc_tensor((P,), dtype, scope="local") for r in Tx.unroll(0, P): # read via src layout j = r ^ ((lane_id >> shift) & mask) idx = decompose(lane_id + j * 32, extent) diff --git a/docs/tirx/tile_primitives/reduction/local.rst b/docs/tirx/tile_primitives/reduction/local.rst index 41a3fc2c628b..0451452f7c27 100644 --- a/docs/tirx/tile_primitives/reduction/local.rst +++ b/docs/tirx/tile_primitives/reduction/local.rst @@ -68,15 +68,15 @@ A single thread reduces a 4-element ``float32`` local vector to a scalar @Tx.prim_func def test_func( - A: Tx.Buffer([4], "float32", layout=TileLayout(S[4,])), - B: Tx.Buffer([1], "float32", layout=TileLayout(S[1,])), + A: Tx.Tensor([4], "float32", layout=TileLayout(S[4,])), + B: Tx.Tensor([1], "float32", layout=TileLayout(S[1,])), ): Tx.device_entry() Tx.cta_id([1]) Tx.thread_id([1]) - A_local = Tx.alloc_buffer([4], "float32", scope="local") - B_local = Tx.alloc_buffer([1], "float32", scope="local") + A_local = Tx.alloc_tensor([4], "float32", scope="local") + B_local = Tx.alloc_tensor([1], "float32", scope="local") for i in Tx.serial(4): A_local[i] = A[i] Tx.tile.sum(B_local, A_local, accum=False) # reduction local dispatch diff --git a/docs/tirx/tile_primitives/reduction/shared.rst b/docs/tirx/tile_primitives/reduction/shared.rst index 0e669e6fd1b0..edfda1b83a54 100644 --- a/docs/tirx/tile_primitives/reduction/shared.rst +++ b/docs/tirx/tile_primitives/reduction/shared.rst @@ -64,15 +64,15 @@ axis ``-1``) to a ``4``-vector (from ``test_reduction.py``): @Tx.prim_func def test_reduction( - A: Tx.Buffer((4, 8), "float32", layout=TileLayout(S[4, 8])), - B: Tx.Buffer((4,), "float32", layout=TileLayout(S[4,])), + A: Tx.Tensor((4, 8), "float32", layout=TileLayout(S[4, 8])), + B: Tx.Tensor((4,), "float32", layout=TileLayout(S[4,])), ): Tx.device_entry() Tx.cta_id([1]) Tx.thread_id([32]) - A_smem = Tx.alloc_buffer((4, 8), "float32", scope="shared", layout=TileLayout(S[(4, 8)])) - B_smem = Tx.alloc_buffer((4,), "float32", scope="shared", layout=TileLayout(S[(4,)])) + A_smem = Tx.alloc_tensor((4, 8), "float32", scope="shared", layout=TileLayout(S[(4, 8)])) + B_smem = Tx.alloc_tensor((4,), "float32", scope="shared", layout=TileLayout(S[(4,)])) Tx.tile.cta.copy(A_smem, A) Tx.cuda.cta_sync() Tx.tile.cta.sum(B_smem, A_smem, axes=(-1,), accum=False) # reduction shared dispatch diff --git a/docs/tirx/tile_primitives/reduction/sm100_packed.rst b/docs/tirx/tile_primitives/reduction/sm100_packed.rst index ad41892e285b..ecc1cf8111fd 100644 --- a/docs/tirx/tile_primitives/reduction/sm100_packed.rst +++ b/docs/tirx/tile_primitives/reduction/sm100_packed.rst @@ -73,15 +73,15 @@ A single thread sums a 32-element ``float32`` local vector on ``sm_100a`` (from @Tx.prim_func def test_func( - A: Tx.Buffer([32], "float32", layout=TileLayout(S[32,])), - B: Tx.Buffer([1], "float32", layout=TileLayout(S[1,])), + A: Tx.Tensor([32], "float32", layout=TileLayout(S[32,])), + B: Tx.Tensor([1], "float32", layout=TileLayout(S[1,])), ): Tx.device_entry() Tx.cta_id([1]) Tx.thread_id([1]) - A_local = Tx.alloc_buffer([32], "float32", scope="local") - B_local = Tx.alloc_buffer([1], "float32", scope="local") + A_local = Tx.alloc_tensor([32], "float32", scope="local") + B_local = Tx.alloc_tensor([1], "float32", scope="local") for i in Tx.serial(32): A_local[i] = A[i] Tx.tile.sum(B_local, A_local, accum=False) # -> packed_add_sum diff --git a/include/tvm/relax/distributed/axis_group_graph.h b/include/tvm/relax/distributed/axis_group_graph.h index 6c944ff275fb..22cc6216977d 100644 --- a/include/tvm/relax/distributed/axis_group_graph.h +++ b/include/tvm/relax/distributed/axis_group_graph.h @@ -73,14 +73,14 @@ class BufferAxisGraphExtractor : public s_tir::StmtExprVisitor { extractor->Visit(prim_func->body); ffi::Map inverse_buffer_map; for (const Var& param : prim_func->params) { - if (param->ty.as()) { + if (param->ty.as()) { inverse_buffer_map.Set(param.as_or_throw(), param); } } std::vector> tir_var_axis_group_list; std::unordered_set visited; for (const Var& param : prim_func->params) { - if (!param->ty.as()) { + if (!param->ty.as()) { continue; } BufferVar buffer = param.as_or_throw(); diff --git a/include/tvm/s_tir/schedule/schedule.h b/include/tvm/s_tir/schedule/schedule.h index 7af08a1630f3..e69ae2b53e31 100644 --- a/include/tvm/s_tir/schedule/schedule.h +++ b/include/tvm/s_tir/schedule/schedule.h @@ -834,7 +834,7 @@ class ScheduleNode : public ffi::Object { * appears in the block's ancestor loops as `rolling axis`, fold and circularize the buffer along * the rolling dimension, append block predicate to avoid recomputing overlapping elements. * It requires: - * 1) The buffer to be an intermediate buffer defined via `alloc_buffer`. + * 1) The buffer to be an intermediate buffer defined via `alloc_tensor`. * 2) The LCA of the producer and consumer of the buffer is a for loop, typically, * the producer and consumer of the buffer are cascaded through compute_at. * 3) The access region of the buffer has at least one dimension that contains diff --git a/include/tvm/s_tir/stmt.h b/include/tvm/s_tir/stmt.h index d79674ff3e5e..3f35087765ee 100644 --- a/include/tvm/s_tir/stmt.h +++ b/include/tvm/s_tir/stmt.h @@ -84,7 +84,7 @@ class MatchBufferRegion : public ffi::ObjectRef { * T.reads([buffer0[start:end, ...], ...]) * T.writes([buffer1[start:end, ...], ...]) * T.where(predicate) - * buffer2 = T.alloc_buffer(shape, dtype) + * buffer2 = T.alloc_tensor(shape, dtype) * buffer3 = Ts.match_buffer(source_buffer[start:end, ...]) * T.attr({attr_key: attr_value, ...}) * with T.init(): diff --git a/include/tvm/s_tir/transform.h b/include/tvm/s_tir/transform.h index 54493013f13f..0f2c2c4b1fe3 100644 --- a/include/tvm/s_tir/transform.h +++ b/include/tvm/s_tir/transform.h @@ -108,7 +108,7 @@ TVM_DLL Pass LiftThreadBinding(); * * for i in range(0, 16): * with T.sblock(): - * B = T.alloc_buffer(16, 16) + * B = T.alloc_tensor(16, 16) * for j in range(0, 16): * B[i, j] = A[i, j] + 1 * for j in range(0, 16): @@ -124,7 +124,7 @@ TVM_DLL Pass LiftThreadBinding(); * * for i in range(0, 16): * with T.sblock(): - * B = T.alloc_buffer(1, 16) + * B = T.alloc_tensor(1, 16) * for j in range(0, 16): * B[0, j] = A[i, j] + 1 * for j in range(0, 16): diff --git a/include/tvm/script/printer/doc.h b/include/tvm/script/printer/doc.h index 14a270b4bb34..eba976d265ac 100644 --- a/include/tvm/script/printer/doc.h +++ b/include/tvm/script/printer/doc.h @@ -819,7 +819,7 @@ class AssignDocNode : public StmtDocNode { /*! * \brief The right hand side of the assignment. * - * If null, this doc represents declaration, e.g. `A: T.Buffer((1,2))` + * If null, this doc represents declaration, e.g. `A: T.Tensor((1,2))` * */ ffi::Optional rhs; /*! \brief The type annotation of this assignment. */ diff --git a/include/tvm/tirx/builtin.h b/include/tvm/tirx/builtin.h index c54dfa008d5d..4c77ce995543 100644 --- a/include/tvm/tirx/builtin.h +++ b/include/tvm/tirx/builtin.h @@ -42,7 +42,7 @@ namespace tirx { /*! \brief Collection of builtin intrinsics as ops */ namespace builtin { /*! - * \brief Allocate a buffer: alloc_buffer(shape, dtype, scope) -> BufferType. + * \brief Allocate a buffer: alloc_tensor(shape, dtype, scope) -> TensorType. * * Arguments, in order: * - args[0]: shape, Tuple of integer extents (IntImm or symbolic integer expressions). @@ -50,12 +50,12 @@ namespace builtin { * - args[2]: scope, StringImm naming the storage scope. * * DictAttrs directly holds the allocation annotations, defaulting to an empty dictionary. - * The BufferType result agrees with the operands and retains buffer access/storage metadata. + * The TensorType result agrees with the operands and retains buffer access/storage metadata. * * \code * // Example pattern match code for a given Binding: * if (const auto* call = binding->value.as(); - * call && call->op.same_as(builtin::alloc_buffer())) { + * call && call->op.same_as(builtin::alloc_tensor())) { * tvm::Tuple shape = call->args[0].as_or_throw(); * DLDataType dtype = call->args[1].as_or_throw()->value; * ffi::String scope = call->args[2].as_or_throw()->value; @@ -63,9 +63,9 @@ namespace builtin { * } * \endcode */ -TVM_DLL const Op& alloc_buffer(); +TVM_DLL const Op& alloc_tensor(); /*! - * \brief Declare a buffer view: decl_buffer(data, shape, dtype, scope) -> BufferType. + * \brief Declare a buffer view: decl_tensor(data, shape, dtype, scope) -> TensorType. * * Arguments, in order: * - args[0]: data, Expr for the existing physical pointer backing the buffer view. @@ -73,13 +73,13 @@ TVM_DLL const Op& alloc_buffer(); * - args[2]: dtype, DataTypeImm with a DLDataType payload for the element type. * - args[3]: scope, StringImm naming the storage scope. * - * There are no attributes. The BufferType result agrees with the operands and retains + * There are no attributes. The TensorType result agrees with the operands and retains * buffer access/storage metadata. The operation binds a view without allocating memory. * * \code * // Example pattern match code for a given Binding: * if (const auto* call = binding->value.as(); - * call && call->op.same_as(builtin::decl_buffer())) { + * call && call->op.same_as(builtin::decl_tensor())) { * Expr data = call->args[0]; * tvm::Tuple shape = call->args[1].as_or_throw(); * DLDataType dtype = call->args[2].as_or_throw()->value; @@ -87,7 +87,7 @@ TVM_DLL const Op& alloc_buffer(); * } * \endcode */ -TVM_DLL const Op& decl_buffer(); +TVM_DLL const Op& decl_tensor(); /*! * \brief Return from a GPU thread without returning a function value. */ @@ -635,7 +635,7 @@ TVM_DLL const Op& buffer_offset(); /*! * \brief Project the physical pointer associated with a BufferVar definition. * - * The result pointer type is derived from the BufferType dtype and storage + * The result pointer type is derived from the TensorType dtype and storage * scope of the sole BufferVar argument. This operation is consumed by TIRx * lowering and code generation. */ diff --git a/include/tvm/tirx/expr.h b/include/tvm/tirx/expr.h index a9e7daad7ad5..cb46cd33034f 100644 --- a/include/tvm/tirx/expr.h +++ b/include/tvm/tirx/expr.h @@ -40,29 +40,29 @@ namespace tirx { class Stmt; /*! - * \brief Checked zero-state view over an ordinary VarNode with BufferType. + * \brief Checked zero-state view over an ordinary VarNode with TensorType. * * BufferVar does not introduce a runtime object or a second identity. It * safely widens to Var, and get() returns the underlying VarNode used by - * identity-sensitive maps. operator-> exposes the immutable BufferType + * identity-sensitive maps. operator-> exposes the immutable TensorType * access contract for concise compiler-side metadata access. */ class BufferVar : public Var { public: - /*! \brief Construct a fresh buffer variable from an explicit BufferType. */ - TVM_DLL explicit BufferVar(ffi::String name, BufferType type, Span span = Span()); + /*! \brief Construct a fresh buffer variable from an explicit TensorType. */ + TVM_DLL explicit BufferVar(ffi::String name, TensorType type, Span span = Span()); /*! \brief Create a checked buffer view over an existing ordinary Var. */ explicit BufferVar(Var var) : Var(std::move(var)) { - TVM_FFI_ICHECK(get() != nullptr && get()->ty.as()) - << "Expected a non-null Var with BufferType"; + TVM_FFI_ICHECK(get() != nullptr && get()->ty.as()) + << "Expected a non-null Var with TensorType"; } /*! \brief Return the ordinary variable view over the same identity. */ Var var() const { return ffi::GetRef(get()); } /*! \brief Return the buffer type carried by the ordinary variable. */ - BufferType type() const { return get()->ty.as_or_throw(); } + TensorType type() const { return get()->ty.as_or_throw(); } /*! \brief Return the buffer's diagnostic name. */ const ffi::String& name() const { return get()->name; } @@ -117,7 +117,7 @@ class BufferVar : public Var { * * If flattening changes the type, the result is a fresh BufferVar. Callers * that use it as a view over this buffer must bind the returned variable with - * a `Bind` of `flattened` to a `decl_buffer` Call over `this->data()`. + * a `Bind` of `flattened` to a `decl_tensor` Call over `this->data()`. */ BufferVar GetFlattenedBuffer() const; @@ -169,12 +169,12 @@ class BufferVar : public Var { explicit BufferVar(ffi::UnsafeInit tag) : Var(tag) {} TVM_FFI_DEFINE_DEFAULT_COPY_MOVE_AND_ASSIGN(BufferVar); - const BufferTypeNode* operator->() const { + const TensorTypeNode* operator->() const { const auto* var_node = static_cast(data_.get()); TVM_FFI_ICHECK(var_node != nullptr); - const auto* type_node = var_node->ty.as(); + const auto* type_node = var_node->ty.as(); TVM_FFI_ICHECK(type_node != nullptr) - << "Expected a Var with BufferType, but " << var_node->name << " has type " << var_node->ty; + << "Expected a Var with TensorType, but " << var_node->name << " has type " << var_node->ty; return type_node; } @@ -194,13 +194,13 @@ inline bool operator!=(const BufferVar& lhs, const BufferVar& rhs) { return !lhs /*! \brief Recover a checked buffer view from an ordinary VarNode pointer. */ inline BufferVar GetBufferVar(const VarNode* var) { return BufferVar(ffi::GetRef(var)); } -inline ffi::ObjectPtr CopyBufferType(const BufferVar& var) { - return ffi::make_object(*var.operator->()); +inline ffi::ObjectPtr CopyTensorType(const BufferVar& var) { + return ffi::make_object(*var.operator->()); } -inline BufferVar RebuildBufferVar(const BufferVar& var, ffi::ObjectPtr type, +inline BufferVar RebuildBufferVar(const BufferVar& var, ffi::ObjectPtr type, ffi::Optional name = std::nullopt) { - return BufferVar(name.value_or(var.name()), BufferType(std::move(type)), var.span()); + return BufferVar(name.value_or(var.name()), TensorType(std::move(type)), var.span()); } /*! @@ -213,7 +213,7 @@ inline BufferVar RebuildBufferVar(const BufferVar& var, ffi::ObjectPtr shape, PrimType dtype = PrimType::Float(32), +TVM_DLL BufferVar decl_tensor(ffi::Array shape, PrimType dtype = PrimType::Float(32), ffi::String name = "buffer", ffi::String storage_scope = "", Span span = Span()); @@ -277,7 +277,7 @@ struct TypeTraits : public ObjectRefTypeTraitsBase( details::ObjectUnsafe::ObjectPtrFromUnowned(src->v_obj).get()); - return details::AnyUnsafe::CheckAnyStrict(var->ExprNode::ty); + return details::AnyUnsafe::CheckAnyStrict(var->ExprNode::ty); } TVM_FFI_INLINE static std::optional TryCastFromAnyView(const TVMFFIAny* src) { diff --git a/include/tvm/tirx/function.h b/include/tvm/tirx/function.h index ec797186d6b5..fce9690444bb 100644 --- a/include/tvm/tirx/function.h +++ b/include/tvm/tirx/function.h @@ -135,7 +135,7 @@ class PrimFunc : public BaseFunc { * from __future__ import annotations * * @Ts.prim_func - * def mem_copy(A: T.Buffer((m, n), "float32"), B: T.Buffer((m, n), "float32"), + * def mem_copy(A: T.Tensor((m, n), "float32"), B: T.Tensor((m, n), "float32"), * m: T.int32, n: T.int32) -> None: * for i, j in T.grid(m, n): * with Ts.sblock(): @@ -147,15 +147,15 @@ class PrimFunc : public BaseFunc { * * \code{.py} * a, _, m, n = mem_copy.params - * func = mem_copy.specialize({a: tirx.decl_buffer((16, 16))}) + * func = mem_copy.specialize({a: tirx.decl_tensor((16, 16))}) * # or * func = mem_copy.specialize({n: 16, m: 16}) * \endcode * * \code{.py} * @Ts.prim_func - * def mem_copy_16_16(A: T.Buffer((16, 16), "float32"), - * B: T.Buffer((16, 16), "float32")) -> None: + * def mem_copy_16_16(A: T.Tensor((16, 16), "float32"), + * B: T.Tensor((16, 16), "float32")) -> None: * for i, j in T.grid(16, 16): * with Ts.sblock(): * vi, vj = Ts.axis.remap("SS", [i, j]) diff --git a/include/tvm/tirx/script/ir_builder/frame.h b/include/tvm/tirx/script/ir_builder/frame.h index 97f47f2aa53c..49f77c912ce2 100644 --- a/include/tvm/tirx/script/ir_builder/frame.h +++ b/include/tvm/tirx/script/ir_builder/frame.h @@ -500,7 +500,7 @@ class ElseFrame : public TIRFrame { TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(ElseFrame, TIRFrame, ElseFrameNode); }; -class DeclBufferFrameNode : public TIRFrameNode { +class DeclTensorFrameNode : public TIRFrameNode { public: /*! \brief The declared buffer. */ tvm::tirx::BufferVar buffer; @@ -511,24 +511,24 @@ class DeclBufferFrameNode : public TIRFrameNode { static void RegisterReflection() { namespace refl = tvm::ffi::reflection; - refl::ObjectDef() - .def_ro("buffer", &DeclBufferFrameNode::buffer) - .def_ro("data", &DeclBufferFrameNode::data) - .def_ro("allocated", &DeclBufferFrameNode::allocated); + refl::ObjectDef() + .def_ro("buffer", &DeclTensorFrameNode::buffer) + .def_ro("data", &DeclTensorFrameNode::data) + .def_ro("allocated", &DeclTensorFrameNode::allocated); } - TVM_FFI_DECLARE_OBJECT_INFO_FINAL("script.ir_builder.tirx.DeclBufferFrame", DeclBufferFrameNode, + TVM_FFI_DECLARE_OBJECT_INFO_FINAL("script.ir_builder.tirx.DeclTensorFrame", DeclTensorFrameNode, TIRFrameNode); public: void ExitWithScope() final; }; -class DeclBufferFrame : public TIRFrame { +class DeclTensorFrame : public TIRFrame { public: - explicit DeclBufferFrame(ffi::ObjectPtr data) : TIRFrame(data) { + explicit DeclTensorFrame(ffi::ObjectPtr data) : TIRFrame(data) { TVM_FFI_ICHECK(data != nullptr); } - TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(DeclBufferFrame, TIRFrame, DeclBufferFrameNode); + TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(DeclTensorFrame, TIRFrame, DeclTensorFrameNode); }; } // namespace tirx diff --git a/include/tvm/tirx/script/ir_builder/ir.h b/include/tvm/tirx/script/ir_builder/ir.h index a04b1cddf5a0..0c397337fe55 100644 --- a/include/tvm/tirx/script/ir_builder/ir.h +++ b/include/tvm/tirx/script/ir_builder/ir.h @@ -317,7 +317,7 @@ ElseFrame Else(); * \param layout The layout of the buffer. * \return The declaration frame. */ -DeclBufferFrame DeclBuffer(ffi::Array shape, PrimType dtype, ffi::String buffer_name, +DeclTensorFrame DeclTensor(ffi::Array shape, PrimType dtype, ffi::String buffer_name, ffi::Optional data, ffi::Optional> strides, ffi::Optional elem_offset, ffi::String storage_scope, int align, int offset_factor, @@ -332,7 +332,7 @@ DeclBufferFrame DeclBuffer(ffi::Array shape, PrimType dtype, ffi::Stri * \param annotations Optional annotations for the allocation. * \return The allocated buffer. */ -BufferVar AllocBuffer(ffi::Array shape, PrimType dtype = PrimType::Float(32), +BufferVar AllocTensor(ffi::Array shape, PrimType dtype = PrimType::Float(32), ffi::String storage_scope = "global", ffi::Optional> annotations = std::nullopt); diff --git a/include/tvm/tirx/stmt.h b/include/tvm/tirx/stmt.h index 492103e49a91..adbd2f467cc8 100644 --- a/include/tvm/tirx/stmt.h +++ b/include/tvm/tirx/stmt.h @@ -750,7 +750,7 @@ constexpr const char* thread_extent = "thread_extent"; constexpr const char* virtual_thread = "virtual_thread"; constexpr const char* async_wait_queue_scope = "async_wait_queue_scope"; constexpr const char* async_wait_inflight_count = "async_wait_inflight_count"; -/*! \brief Annotation key on AllocBuffer marking the allocation as volatile. */ +/*! \brief Annotation key on AllocTensor marking the allocation as volatile. */ constexpr const char* kVolatile = "tirx.volatile"; /*! \brief Mark buffer initial addr alignment in bytes */ constexpr const char* buffer_data_alignment = "buffer_data_alignment"; diff --git a/include/tvm/tirx/tile_primitive.h b/include/tvm/tirx/tile_primitive.h index 8fbea6becb99..5f39a7c10365 100644 --- a/include/tvm/tirx/tile_primitive.h +++ b/include/tvm/tirx/tile_primitive.h @@ -92,7 +92,7 @@ constexpr const char* kDeviceInitStmt = "device_init_stmt"; * which will be inserted at the beginning of the kernel */ constexpr const char* kHostInitStmt = "host_init_stmt"; -/*! \brief Statements to be inserted after a specific buffer's definition (DeclBuffer/AllocBuffer). +/*! \brief Statements to be inserted after a specific buffer's definition (DeclTensor/AllocTensor). * Stored as Map>. */ constexpr const char* kPostBufferDefStmt = "post_buffer_def_stmt"; diff --git a/include/tvm/tirx/transform.h b/include/tvm/tirx/transform.h index d8bdc625effc..437d1849859d 100644 --- a/include/tvm/tirx/transform.h +++ b/include/tvm/tirx/transform.h @@ -337,7 +337,7 @@ TVM_DLL Pass TilePrimitiveDispatch(); TVM_DLL Pass LowerTIRxCleanup(); /*! - * \brief Lower opaque constructs in TIRX programs: AllocBuffer, For(thread_binding), + * \brief Lower opaque constructs in TIRX programs: allocation calls, For(thread_binding), * unit loop elimination. This is the tirx-specific counterpart of * s_tir::LowerOpaqueBlock, without any SBlock handling. * \return The pass. diff --git a/include/tvm/tirx/type.h b/include/tvm/tirx/type.h index 57e9f4ec857a..c9a993b28c9b 100644 --- a/include/tvm/tirx/type.h +++ b/include/tvm/tirx/type.h @@ -54,14 +54,14 @@ inline DLDataType DefaultIndexType() { } /*! - * \brief Structural type of a TIRx buffer variable. + * \brief Structural type of a TIRx tensor variable. * - * A buffer value is an ordinary VarNode whose ExprNode::ty is BufferType. - * BufferType owns the immutable access contract. The physical pointer is + * A tensor value is an ordinary VarNode whose ExprNode::ty is TensorType. + * TensorType owns the immutable access contract. The physical pointer is * deliberately not stored here; it is obtained with buffer_data(BufferVar) * and is bound by the surrounding buffer definition. */ -class BufferTypeNode : public TypeNode { +class TensorTypeNode : public TypeNode { public: /*! \brief dtype in the content of the tensor */ PrimType dtype = PrimType::Void(); @@ -98,24 +98,24 @@ class BufferTypeNode : public TypeNode { ffi::Array allocated_addr; /*! \brief constructor */ - BufferTypeNode() {} + TensorTypeNode() {} static void RegisterReflection() { namespace refl = tvm::ffi::reflection; - refl::ObjectDef() - .def_ro("dtype", &BufferTypeNode::dtype) - .def_ro("storage_scope", &BufferTypeNode::storage_scope) + refl::ObjectDef() + .def_ro("dtype", &TensorTypeNode::dtype) + .def_ro("storage_scope", &TensorTypeNode::storage_scope) // TODO(tqchen): use SEqHashDefSimple after the next pypi tvm-ffi release - .def_ro("shape", &BufferTypeNode::shape, refl::AttachFieldFlag::SEqHashDefPattern()) + .def_ro("shape", &TensorTypeNode::shape, refl::AttachFieldFlag::SEqHashDefPattern()) // TODO(tqchen): use SEqHashDefSimple after the next pypi tvm-ffi release - .def_ro("strides", &BufferTypeNode::strides, refl::AttachFieldFlag::SEqHashDefPattern()) + .def_ro("strides", &TensorTypeNode::strides, refl::AttachFieldFlag::SEqHashDefPattern()) // TODO(tqchen): use SEqHashDefSimple after the next pypi tvm-ffi release - .def_ro("elem_offset", &BufferTypeNode::elem_offset, + .def_ro("elem_offset", &TensorTypeNode::elem_offset, refl::AttachFieldFlag::SEqHashDefPattern()) - .def_ro("data_alignment", &BufferTypeNode::data_alignment) - .def_ro("offset_factor", &BufferTypeNode::offset_factor) - .def_ro("layout", &BufferTypeNode::layout) - .def_ro("allocated_addr", &BufferTypeNode::allocated_addr); + .def_ro("data_alignment", &TensorTypeNode::data_alignment) + .def_ro("offset_factor", &TensorTypeNode::offset_factor) + .def_ro("layout", &TensorTypeNode::layout) + .def_ro("allocated_addr", &TensorTypeNode::allocated_addr); } /*! \return preferred index type for this buffer node */ @@ -146,22 +146,22 @@ class BufferTypeNode : public TypeNode { */ ffi::Array ElemOffset(ffi::Array index, bool inner = false) const; - TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tirx.BufferType", BufferTypeNode, TypeNode); + TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tirx.TensorType", TensorTypeNode, TypeNode); }; /*! - * \brief Managed reference to BufferTypeNode. + * \brief Managed reference to TensorTypeNode. */ -class BufferType : public Type { +class TensorType : public Type { public: - TVM_DLL BufferType(ffi::String storage_scope, PrimType dtype, ffi::Array shape, + TVM_DLL TensorType(ffi::String storage_scope, PrimType dtype, ffi::Array shape, ffi::Array strides, PrimExpr elem_offset, int data_alignment, int offset_factor, ffi::Optional layout = std::nullopt, ffi::Array allocated_addr = {}, Span span = Span()); - TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(BufferType, Type, BufferTypeNode); + TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(TensorType, Type, TensorTypeNode); - explicit BufferType(ffi::ObjectPtr n) : Type(ffi::UnsafeInit{}) { + explicit TensorType(ffi::ObjectPtr n) : Type(ffi::UnsafeInit{}) { TVM_FFI_ICHECK(n != nullptr); data_ = std::move(n); } diff --git a/include/tvm/topi/detail/extern.h b/include/tvm/topi/detail/extern.h index 14b16ace93d5..4c5fe2759c03 100644 --- a/include/tvm/topi/detail/extern.h +++ b/include/tvm/topi/detail/extern.h @@ -72,11 +72,11 @@ inline ffi::Array make_extern(const ffi::Array>& ou ffi::Array input_placeholders; for (auto t : inputs) { - input_placeholders.push_back(tvm::tirx::decl_buffer(t->shape, t->dtype, t->op->name)); + input_placeholders.push_back(tvm::tirx::decl_tensor(t->shape, t->dtype, t->op->name)); } ffi::Array output_placeholders; for (size_t i = 0; i < out_shapes.size(); ++i) { - output_placeholders.push_back(tvm::tirx::decl_buffer(out_shapes[i], out_types[i], name)); + output_placeholders.push_back(tvm::tirx::decl_tensor(out_shapes[i], out_types[i], name)); } auto body = fextern(input_placeholders, output_placeholders); diff --git a/python/tvm/backend/cuda/lang/alloc_pool.py b/python/tvm/backend/cuda/lang/alloc_pool.py index f3d7eba2659f..119b4000d71b 100644 --- a/python/tvm/backend/cuda/lang/alloc_pool.py +++ b/python/tvm/backend/cuda/lang/alloc_pool.py @@ -290,7 +290,7 @@ def alloc(self, shape, dtype="float32", *, layout=None, cols=None): if layout is None: assert len(shape) == 2, "TMEMPool.alloc() requires layout= for non-2D TMEM buffers" layout = _default_tmem_layout(shape[0], shape[1]) - res = ir.decl_buffer(shape, dtype, scope="tmem", allocated_addr=col_start, layout=layout) + res = ir.decl_tensor(shape, dtype, scope="tmem", allocated_addr=col_start, layout=layout) self.offset = col_end self.max_offset = max(self.max_offset, self.offset) return res @@ -400,7 +400,7 @@ class SMEMPool: Parameters ---------- ptr : Var or None, optional - If omitted, an ``alloc_buffer([0], "uint8", scope="shared.dyn")`` is + If omitted, an ``alloc_tensor([0], "uint8", scope="shared.dyn")`` is created automatically and ``commit()`` must be called after all allocations to emit the size annotation. If a ``Var`` is provided, the caller manages the backing buffer and @@ -410,7 +410,7 @@ class SMEMPool: def __init__(self, ptr=_POOL_UNSET): ir = _get_ir() if ptr is _POOL_UNSET: - self.buf = ir.alloc_buffer([0], "uint8", scope="shared.dyn") + self.buf = ir.alloc_tensor([0], "uint8", scope="shared.dyn") self.ptr = self.buf.data self._owns_buffer = True else: @@ -432,7 +432,7 @@ def alloc( ir = _get_ir() if align > 0: self.offset = (self.offset + align - 1) // align * align - res = ir.decl_buffer( + res = ir.decl_tensor( shape, dtype, data=self.ptr, diff --git a/python/tvm/backend/cuda/lang/pipeline.py b/python/tvm/backend/cuda/lang/pipeline.py index f43563bd2166..8743674442a4 100644 --- a/python/tvm/backend/cuda/lang/pipeline.py +++ b/python/tvm/backend/cuda/lang/pipeline.py @@ -117,7 +117,7 @@ def _map_buffer_into_cta(ptr, rank, depth): T.evaluate(T.ptx.mapa.u64(mapped[0], ptr, T.uint32(rank))) remote_ptr = TIRVar("remote_mbar_ptr", ptr_ty) T.bind(T.reinterpret(ptr_ty, mapped[0]), var=remote_ptr) - return T.decl_buffer([depth], "uint64", data=remote_ptr, scope="shared") + return T.decl_tensor([depth], "uint64", data=remote_ptr, scope="shared") def _mbarrier_arrive_remote(bar, pred=None, count=None): diff --git a/python/tvm/backend/cuda/tile_primitive/copy_async/dsmem.py b/python/tvm/backend/cuda/tile_primitive/copy_async/dsmem.py index 468e64734b13..25c08c913de8 100644 --- a/python/tvm/backend/cuda/tile_primitive/copy_async/dsmem.py +++ b/python/tvm/backend/cuda/tile_primitive/copy_async/dsmem.py @@ -167,13 +167,13 @@ def impl(): for loop_vars in T.grid(*outer_extents): src_elem_offset, dst_elem_offset = T.meta_var(compute_offsets(loop_vars)) - src_buf_w = T.decl_buffer( + src_buf_w = T.decl_tensor( src_buf.shape, src_buf.dtype, src_buf.data, elem_offset=src_buf.elem_offset + src_elem_offset, scope=src_buf.scope(), layout=src_tile, ) - dst_buf_w = T.decl_buffer( + dst_buf_w = T.decl_tensor( dst_buf.shape, dst_buf.dtype, dst_buf.data, elem_offset=dst_buf.elem_offset + dst_elem_offset, scope=dst_buf.scope(), diff --git a/python/tvm/backend/cuda/tile_primitive/copy_async/tcgen05_cp.py b/python/tvm/backend/cuda/tile_primitive/copy_async/tcgen05_cp.py index 32b14cd2cf2c..74c4ad7e22d3 100644 --- a/python/tvm/backend/cuda/tile_primitive/copy_async/tcgen05_cp.py +++ b/python/tvm/backend/cuda/tile_primitive/copy_async/tcgen05_cp.py @@ -676,7 +676,7 @@ def _get_or_create_desc(sctx, s_buf, ldo, sdo, swizzle): if cached is not None: return cached - desc_buf = tvm.tirx.decl_buffer((1,), "uint64", name="cp_desc", scope="local") + desc_buf = tvm.tirx.decl_tensor((1,), "uint64", name="cp_desc", scope="local") encode_call = T.cuda.tcgen05.encode_matrix_descriptor( desc_buf.data, T.reinterpret("handle", T.uint64(0)), ldo, sdo, swizzle ) @@ -685,7 +685,7 @@ def _get_or_create_desc(sctx, s_buf, ldo, sdo, swizzle): Bind( desc_buf, Call( - "tirx.alloc_buffer", + "tirx.alloc_tensor", [ Tuple(desc_buf.ty.shape), DataTypeImm(desc_buf.ty.dtype.dtype), diff --git a/python/tvm/backend/cuda/tile_primitive/copy_async/tma.py b/python/tvm/backend/cuda/tile_primitive/copy_async/tma.py index 96bb8c1ab758..c39d71011a46 100644 --- a/python/tvm/backend/cuda/tile_primitive/copy_async/tma.py +++ b/python/tvm/backend/cuda/tile_primitive/copy_async/tma.py @@ -2168,10 +2168,10 @@ def emit_at(shared_ptr, coordinates): def shared_ptr(element_offset=0): # Keep the sliced offset in the pointer index instead of the Buffer's - # elem_offset. The latter is part of a flat DeclBuffer definition; + # elem_offset. The latter is part of a flat DeclTensor definition; # after loop unrolling, CSE may otherwise lift an offset containing a # locally bound coordinate above that coordinate's Bind statement. - smem_view = T.decl_buffer( + smem_view = T.decl_tensor( (1,), spec.smem_buffer.dtype, spec.smem_buffer.data, diff --git a/python/tvm/backend/cuda/tile_primitive/elementwise/reg.py b/python/tvm/backend/cuda/tile_primitive/elementwise/reg.py index f4bbf688a471..d77180b03f6d 100644 --- a/python/tvm/backend/cuda/tile_primitive/elementwise/reg.py +++ b/python/tvm/backend/cuda/tile_primitive/elementwise/reg.py @@ -354,7 +354,7 @@ def _make_views_meta(per_op_carved, per_thread_total): iter strides at codegen time. """ return { - op_br: T.decl_buffer( + op_br: T.decl_tensor( (per_thread_total,), op_br.source.dtype, op_br.source.data, diff --git a/python/tvm/backend/cuda/tile_primitive/gemm_async/tcgen05.py b/python/tvm/backend/cuda/tile_primitive/gemm_async/tcgen05.py index d76eda7c08fe..01ee954c0a11 100644 --- a/python/tvm/backend/cuda/tile_primitive/gemm_async/tcgen05.py +++ b/python/tvm/backend/cuda/tile_primitive/gemm_async/tcgen05.py @@ -1076,8 +1076,8 @@ def _atom_off(dim): _krp = Evaluate(tirx_op.tvm_kernel_replace_point()) def _make_lo_uniform(desc_buf): - desc_lo = tvm.tirx.decl_buffer((1,), "uint32", name=f"{desc_buf.name}_lo", scope="local") - desc_hi = tvm.tirx.decl_buffer((1,), "uint32", name=f"{desc_buf.name}_hi", scope="local") + desc_lo = tvm.tirx.decl_tensor((1,), "uint32", name=f"{desc_buf.name}_lo", scope="local") + desc_hi = tvm.tirx.decl_tensor((1,), "uint32", name=f"{desc_buf.name}_hi", scope="local") unpack = T.ptx.mov.b64(desc_lo[0], desc_hi[0], desc_buf[0]) shuffle = T.ptx.shfl_sync.idx.b32( desc_lo[0], @@ -1092,7 +1092,7 @@ def _make_lo_uniform(desc_buf): Bind( desc_lo, Call( - "tirx.alloc_buffer", + "tirx.alloc_tensor", [ Tuple(desc_lo.ty.shape), DataTypeImm(desc_lo.ty.dtype.dtype), @@ -1105,7 +1105,7 @@ def _make_lo_uniform(desc_buf): Bind( desc_hi, Call( - "tirx.alloc_buffer", + "tirx.alloc_tensor", [ Tuple(desc_hi.ty.shape), DataTypeImm(desc_hi.ty.dtype.dtype), @@ -1126,7 +1126,7 @@ def _make_desc(smem_buf, ldo, sdo, swizzle_val, name): # issuer, so make the descriptor low word uniform there. A # single-thread caller is already elected: a full-mask shuffle in that # scope is invalid, and the same thread consumes the descriptor anyway. - desc_buf = tvm.tirx.decl_buffer((1,), "uint64", name=name, scope="local") + desc_buf = tvm.tirx.decl_tensor((1,), "uint64", name=name, scope="local") encode_call = tvm.tirx.call_intrin( "", "tirx.cuda.tcgen05_encode_matrix_descriptor", @@ -1140,7 +1140,7 @@ def _make_desc(smem_buf, ldo, sdo, swizzle_val, name): Bind( desc_buf, Call( - "tirx.alloc_buffer", + "tirx.alloc_tensor", [ Tuple(desc_buf.ty.shape), DataTypeImm(desc_buf.ty.dtype.dtype), diff --git a/python/tvm/backend/cuda/tile_primitive/permute_layout/warp_xor_swizzle.py b/python/tvm/backend/cuda/tile_primitive/permute_layout/warp_xor_swizzle.py index 3b179bcedea7..fee4ea95c037 100644 --- a/python/tvm/backend/cuda/tile_primitive/permute_layout/warp_xor_swizzle.py +++ b/python/tvm/backend/cuda/tile_primitive/permute_layout/warp_xor_swizzle.py @@ -332,7 +332,7 @@ def _iter_off(iter_idx, strides): def impl(): warp_size = T.meta_var(32) lane_id = T.meta_var(tid_x % warp_size) - regs = T.alloc_buffer((P,), bits_dtype, scope="local") + regs = T.alloc_tensor((P,), bits_dtype, scope="local") base_src = T.meta_var(src_buf.ptr_to(list(src_st))) base_dst = T.meta_var(dst_buf.ptr_to(list(dst_st))) # Phase 1: read via L_src @@ -358,7 +358,7 @@ def impl(): def impl(): warp_size = T.meta_var(32) lane_id = T.meta_var(tid_x % warp_size) - regs = T.alloc_buffer((P,), dtype, scope="local") + regs = T.alloc_tensor((P,), dtype, scope="local") # Phase 1: read via L_src for r in T.unroll(0, P): j = T.meta_var(r ^ ((lane_id >> shift) & mask)) @@ -400,7 +400,7 @@ def impl(): # # After (BLK_SFA=128, P=4, k=2, shift=3): # lane_id = threadIdx.x % 32 -# regs = T.alloc_buffer((4,), "uint32", scope="local") +# regs = T.alloc_tensor((4,), "uint32", scope="local") # for r in T.unroll(4): # j = r ^ ((lane_id >> 3) & 0x3) # flat = lane_id + j * 32 diff --git a/python/tvm/backend/cuda/tile_primitive/reduction/local.py b/python/tvm/backend/cuda/tile_primitive/reduction/local.py index c6fe3a471128..fd6af4882570 100644 --- a/python/tvm/backend/cuda/tile_primitive/reduction/local.py +++ b/python/tvm/backend/cuda/tile_primitive/reduction/local.py @@ -362,7 +362,7 @@ def impl(): # the mediated layout explicitly. Bare local() is physical-order. src_local = src.local(*src_local_shape, layout=src.layout.storage()) dst_local = dst.local(*dst_local_shape, layout=dst.layout.storage()) - old_val = T.alloc_buffer([1], dtype, scope="local") + old_val = T.alloc_tensor([1], dtype, scope="local") for spa in T.serial(dst_local_total): dst_idx = T.meta_var(get_indices(spa, dst_local_st, dst_local_ext)) diff --git a/python/tvm/backend/cuda/tile_primitive/reduction/shared.py b/python/tvm/backend/cuda/tile_primitive/reduction/shared.py index 74334c2de00d..3b4e7b9d0cb7 100644 --- a/python/tvm/backend/cuda/tile_primitive/reduction/shared.py +++ b/python/tvm/backend/cuda/tile_primitive/reduction/shared.py @@ -192,7 +192,7 @@ def sync(): @T.prim_func def impl(): tid_in_scope = get_tid_in_scope() - thread_data = T.alloc_buffer([1], dtype=dtype, scope="local") + thread_data = T.alloc_tensor([1], dtype=dtype, scope="local") group_id = T.meta_var(T.floordiv(tid_in_scope, group_size)) lane_in_grp = T.meta_var(tid_in_scope % group_size) for step in T.serial(T.ceildiv(spatial_len, spatial_par)): diff --git a/python/tvm/backend/cuda/tile_primitive/reduction/sm100_packed.py b/python/tvm/backend/cuda/tile_primitive/reduction/sm100_packed.py index bec0bf777469..a0630fa5e489 100644 --- a/python/tvm/backend/cuda/tile_primitive/reduction/sm100_packed.py +++ b/python/tvm/backend/cuda/tile_primitive/reduction/sm100_packed.py @@ -91,7 +91,7 @@ def _emit_reduction_local_thread_packed_add_sum( # fmt: off @T.prim_func(check_well_formed=False) def impl(): - local_sum = T.alloc_buffer([8], dtype, scope="local") + local_sum = T.alloc_tensor([8], dtype, scope="local") # add.f32x2's operands are .b64 register pairs, so each packed add is # mov (pack) -> add -> mov (unpack). nvcc emitted exactly these movs for # the old make_float2/float2_x glue too; writing them keeps every @@ -173,7 +173,7 @@ def _emit_reduction_local_thread_3input_maxmin( # fmt: off @T.prim_func(check_well_formed=False) def impl(): - temp = T.alloc_buffer([4], dtype, scope="local") + temp = T.alloc_tensor([4], dtype, scope="local") # First pass: process first 8 elements into 4 temps for i in T.unroll(4): if accum and i == 0: diff --git a/python/tvm/backend/trn/transform/naive_allocator.py b/python/tvm/backend/trn/transform/naive_allocator.py index 4a768558d9f9..b704d5baff52 100644 --- a/python/tvm/backend/trn/transform/naive_allocator.py +++ b/python/tvm/backend/trn/transform/naive_allocator.py @@ -56,7 +56,7 @@ def _get_alloc_pool_start(stmt) -> int: def collect_alloc_buffer(op: Bind): nonlocal alloc_pool_start - if not isinstance(op.value, Call) or op.value.op != Op.get("tirx.alloc_buffer"): + if not isinstance(op.value, Call) or op.value.op != Op.get("tirx.alloc_tensor"): return buffer = op.var if len(buffer.ty.allocated_addr) == 0: @@ -79,7 +79,7 @@ def _allocate_missing_buffers(stmt, alloc_pool_start: int): def allocate_buffer(op: Bind): nonlocal alloc_offset - if not isinstance(op.value, Call) or op.value.op != Op.get("tirx.alloc_buffer"): + if not isinstance(op.value, Call) or op.value.op != Op.get("tirx.alloc_tensor"): return op buffer = op.var shape = op.value.args[0].fields diff --git a/python/tvm/backend/trn/transform/private_buffer_alloc.py b/python/tvm/backend/trn/transform/private_buffer_alloc.py index ee3a01254299..25ad98281321 100644 --- a/python/tvm/backend/trn/transform/private_buffer_alloc.py +++ b/python/tvm/backend/trn/transform/private_buffer_alloc.py @@ -90,7 +90,7 @@ def visit_attr(op: AttrStmt): allocation = Bind( buffer, Call( - "tirx.alloc_buffer", + "tirx.alloc_tensor", [ Tuple(buffer.ty.shape), DataTypeImm(buffer.ty.dtype.dtype), diff --git a/python/tvm/relax/analysis/analysis.py b/python/tvm/relax/analysis/analysis.py index c4ef71493ae2..5a258fbb843e 100644 --- a/python/tvm/relax/analysis/analysis.py +++ b/python/tvm/relax/analysis/analysis.py @@ -528,7 +528,7 @@ def check_well_formed(obj: IRModule | Function, check_ty: bool = True) -> bool: def _get_prim_func_default_dtype(func: PrimFunc): - """Detect default index dtype from BufferType-annotated parameters.""" + """Detect default index dtype from TensorType-annotated parameters.""" for param in func.params: if tirx.is_buffer_var(param): for value in param.shape: diff --git a/python/tvm/relax/backend/gpu_generic/cumsum.py b/python/tvm/relax/backend/gpu_generic/cumsum.py index 626ce3d1a28a..1ca3eb1ac6ee 100644 --- a/python/tvm/relax/backend/gpu_generic/cumsum.py +++ b/python/tvm/relax/backend/gpu_generic/cumsum.py @@ -101,9 +101,9 @@ def block_inclusive_inside_block( batch: T.int64, cur_len: T.int64, num_blocks: T.int64, - source: T.Buffer, - output: T.Buffer, - tmp_buf: T.Buffer, + source: T.Tensor, + output: T.Tensor, + tmp_buf: T.Tensor, src_offset: T.int64, tmp_offset: T.int64, ): @@ -158,8 +158,8 @@ def update_cross_block( batch: T.int64, cur_len: T.int64, num_blocks: T.int64, - source: T.Buffer, - output: T.Buffer, + source: T.Tensor, + output: T.Tensor, src_offset: T.int64, out_offset: T.int64, ): @@ -178,10 +178,10 @@ def update_cross_block( n = T.dynamic("n") @Ts.prim_func(private=True) - def cumsum(A: T.Buffer([m, n], dtype=in_dtype), Out: T.Buffer([m, n], dtype=out_dtype)): + def cumsum(A: T.Tensor([m, n], dtype=in_dtype), Out: T.Tensor([m, n], dtype=out_dtype)): T.func_attr({"tirx.is_scheduled": True}) # prevent further scheduling - Tmp = T.alloc_buffer([m, n], dtype=out_dtype) + Tmp = T.alloc_tensor([m, n], dtype=out_dtype) # LowerIntrin may implement signed FloorDiv using a sign-bit shift. Keep # hierarchy counting division-free so WebGPU can narrow indices to int32. total_rounds: T.let[T.int64] = _get_total_rounds(n, log_block_n, index_bits) @@ -260,8 +260,8 @@ def gpu_3d_axis_1_cumsum( @Ts.prim_func(private=True) def cumsum( - A: T.Buffer([outer, scan, inner], dtype=in_dtype), - Out: T.Buffer([outer, scan, inner], dtype=out_dtype), + A: T.Tensor([outer, scan, inner], dtype=in_dtype), + Out: T.Tensor([outer, scan, inner], dtype=out_dtype), ): T.func_attr({"tirx.is_scheduled": True}) diff --git a/python/tvm/relax/backend/gpu_generic/sampling.py b/python/tvm/relax/backend/gpu_generic/sampling.py index 911593b456aa..4c4a8e5abd9d 100644 --- a/python/tvm/relax/backend/gpu_generic/sampling.py +++ b/python/tvm/relax/backend/gpu_generic/sampling.py @@ -102,8 +102,8 @@ def gpu_multinomial_from_uniform( def block_cumsum( ty: T.int64, tx: T.int64, - source_local: T.Buffer, - output_shared: T.Buffer, + source_local: T.Tensor, + output_shared: T.Tensor, ): """cumsum inside block (SM)""" # Inclusive scan inside thread @@ -136,8 +136,8 @@ def compare_bool_not_equal(a: T.bool, b: T.bool) -> T.bool: def block_adjacent_difference_left( ty: T.int64, tx: T.int64, - source_local: T.Buffer, - output_local: T.Buffer, + source_local: T.Tensor, + output_local: T.Tensor, ): with Ts.sblock(): shared_buf = Ts.sblock_alloc_buffer((TX * TY,), "bool", scope="shared") @@ -162,11 +162,11 @@ def block_reduce_with_mask( ty: T.int64, tx: T.int64, init_value, - data_local: T.Buffer, - output_local: T.Buffer, + data_local: T.Tensor, + output_local: T.Tensor, dtype: str, reduce_op: Callable, # T.macro - mask_local: T.Buffer | None = None, + mask_local: T.Tensor | None = None, ): with Ts.sblock(): local_sum = Ts.sblock_alloc_buffer((), dtype, scope="local") @@ -265,10 +265,10 @@ def single_batch_sampling( @Ts.prim_func def parallel_sampling_from_prob( - prob: T.Buffer((n, vocab_size), prob_dtype), - uniform_samples: T.Buffer((batch_size, 1), sample_dtype), - row_indices: T.Buffer((batch_size, 1), sample_indices_dtype), - token_ids: T.Buffer((batch_size, 1), dtype), + prob: T.Tensor((n, vocab_size), prob_dtype), + uniform_samples: T.Tensor((batch_size, 1), sample_dtype), + row_indices: T.Tensor((batch_size, 1), sample_indices_dtype), + token_ids: T.Tensor((batch_size, 1), dtype), ): T.func_attr({"tirx.is_scheduled": True}) # match buffers @@ -324,10 +324,10 @@ def generic_get_sample_index( @Ts.prim_func(private=True) def _get_sample_index( - prob: T.Buffer((batch, vocab_size), prob_dtype), - usample: T.Buffer((out_batch, 1), sample_dtype), - sample_indices: T.Buffer((out_batch, 1), sample_indices_dtype), - output_index: T.Buffer((out_batch, 1), dtype), + prob: T.Tensor((batch, vocab_size), prob_dtype), + usample: T.Tensor((out_batch, 1), sample_dtype), + sample_indices: T.Tensor((out_batch, 1), sample_indices_dtype), + output_index: T.Tensor((out_batch, 1), dtype), ): for ax0, ax1 in T.grid(out_batch, vocab_size): with Ts.sblock("T_get_sample_index"): diff --git a/python/tvm/relax/block_builder.py b/python/tvm/relax/block_builder.py index ff6298607dc0..5d304b47e20f 100644 --- a/python/tvm/relax/block_builder.py +++ b/python/tvm/relax/block_builder.py @@ -477,9 +477,9 @@ def te_func(args, args_dict, msg): class Module: @Ts.prim_func def te_func( - rxplaceholder: T.Buffer([n, m], dtype="float32"), - rxplaceholder_1: T.Buffer([n, m], dtype="float32"), - compute: T.Buffer([128, 128], dtype="float32"), + rxplaceholder: T.Tensor([n, m], dtype="float32"), + rxplaceholder_1: T.Tensor([n, m], dtype="float32"), + compute: T.Tensor([128, 128], dtype="float32"), ) -> None: # function attr dict T.func_attr({"tirx.noalias": True}) @@ -528,8 +528,8 @@ def te_func(A): class Module: @Ts.prim_func def te_func( - rxplaceholder: T.Buffer([n + T.int64(1)], dtype="float32"), - compute: T.Buffer([n + T.int64(1)], dtype="float32"), + rxplaceholder: T.Tensor([n + T.int64(1)], dtype="float32"), + compute: T.Tensor([n + T.int64(1)], dtype="float32"), n: T.int64, ) -> None: diff --git a/python/tvm/relax/frontend/nn/llm/_decode_kernels.py b/python/tvm/relax/frontend/nn/llm/_decode_kernels.py index e7671175239a..83bba4dfe270 100644 --- a/python/tvm/relax/frontend/nn/llm/_decode_kernels.py +++ b/python/tvm/relax/frontend/nn/llm/_decode_kernels.py @@ -67,15 +67,15 @@ def _attention_decode_cpu(num_kv_heads, num_qo_heads, head_dim, qkv_dtype, slidi length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") @Ts.prim_func def batch_decode_paged_kv( - Q: T.Buffer((B, H_qo, D), qkv_dtype), - pages: T.Buffer((max_num_pages, 2, H_kv, page_size, D), qkv_dtype), - page_table_indptr: T.Buffer((B + 1,), 'int32', elem_offset=page_indptr_elem_offset), - page_table_values: T.Buffer((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), + Q: T.Tensor((B, H_qo, D), qkv_dtype), + pages: T.Tensor((max_num_pages, 2, H_kv, page_size, D), qkv_dtype), + page_table_indptr: T.Tensor((B + 1,), 'int32', elem_offset=page_indptr_elem_offset), + page_table_values: T.Tensor((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), length_info: _length_info_buffer(B, sliding_window, length_info_elem_offset), # [b] when sliding window = False, or otherwise [3, b] - k_rope_pos_offset: T.Buffer((B,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), - q_rope_position: T.Buffer((B,), 'int32', elem_offset=q_rope_position_elem_offset), - output: T.Buffer((B, H_qo, D), qkv_dtype), - lse: T.Buffer((B, H_qo), 'float32'), + k_rope_pos_offset: T.Tensor((B,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), + q_rope_position: T.Tensor((B,), 'int32', elem_offset=q_rope_position_elem_offset), + output: T.Tensor((B, H_qo, D), qkv_dtype), + lse: T.Tensor((B, H_qo), 'float32'), rotary_mode: T.int32, rope_scale: T.float32, rope_theta: T.float32, @@ -214,15 +214,15 @@ def _attention_decode(num_kv_heads, num_qo_heads, head_dim, qkv_dtype, sliding_w length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") @Ts.prim_func def batch_decode_paged_kv( - Q: T.Buffer((B, H_qo, D), qkv_dtype), - pages: T.Buffer((max_num_pages, 2, H_kv, page_size, D), qkv_dtype, elem_offset=pages_elem_offset), - page_table_indptr: T.Buffer((B + 1,), 'int32', elem_offset=page_indptr_elem_offset), - page_table_values: T.Buffer((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), + Q: T.Tensor((B, H_qo, D), qkv_dtype), + pages: T.Tensor((max_num_pages, 2, H_kv, page_size, D), qkv_dtype, elem_offset=pages_elem_offset), + page_table_indptr: T.Tensor((B + 1,), 'int32', elem_offset=page_indptr_elem_offset), + page_table_values: T.Tensor((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), length_info: _length_info_buffer(B, sliding_window, length_info_elem_offset), # [b] when sliding window = False, or otherwise [3, b] - k_rope_pos_offset: T.Buffer((B,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), - q_rope_position: T.Buffer((B,), 'int32', elem_offset=q_rope_position_elem_offset), - output: T.Buffer((B, H_qo, D), qkv_dtype), - lse: T.Buffer((B, H_qo), 'float32'), + k_rope_pos_offset: T.Tensor((B,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), + q_rope_position: T.Tensor((B,), 'int32', elem_offset=q_rope_position_elem_offset), + output: T.Tensor((B, H_qo, D), qkv_dtype), + lse: T.Tensor((B, H_qo), 'float32'), rotary_mode: T.int32, rope_scale: T.float32, rope_theta: T.float32, @@ -400,10 +400,10 @@ def _merge_state_inplace_cpu(v_dtype): D = T.dynamic("D", "int32") @Ts.prim_func def merge_state_inplace_cpu( - V: T.Buffer((N, H, D), v_dtype), - S: T.Buffer((N, H), 'float32'), - V_other: T.Buffer((N, H, D), v_dtype), - S_other: T.Buffer((N, H), 'float32'), + V: T.Tensor((N, H, D), v_dtype), + S: T.Tensor((N, H), 'float32'), + V_other: T.Tensor((N, H, D), v_dtype), + S_other: T.Tensor((N, H), 'float32'), ): T.func_attr({"tirx.is_scheduled": True}) @@ -445,10 +445,10 @@ def _merge_state_inplace(num_heads, head_dim, v_dtype, target: Target, global_sy D = T.dynamic("D", "int32") @Ts.prim_func def merge_state_inplace( - V: T.Buffer((N, H, D), v_dtype), - S: T.Buffer((N, H), 'float32'), - V_other: T.Buffer((N, H, D), v_dtype), - S_other: T.Buffer((N, H), 'float32'), + V: T.Tensor((N, H, D), v_dtype), + S: T.Tensor((N, H), 'float32'), + V_other: T.Tensor((N, H, D), v_dtype), + S_other: T.Tensor((N, H), 'float32'), ): T.func_attr({"tirx.is_scheduled": True}) diff --git a/python/tvm/relax/frontend/nn/llm/_kernel_common.py b/python/tvm/relax/frontend/nn/llm/_kernel_common.py index 78c7b03004a0..293547f2d798 100644 --- a/python/tvm/relax/frontend/nn/llm/_kernel_common.py +++ b/python/tvm/relax/frontend/nn/llm/_kernel_common.py @@ -107,7 +107,7 @@ class RopeMode(enum.IntEnum): NORMAL = 1 INLINE = 2 -def _rope(buffer: T.Buffer, offset: tirx.Var, rotary_dim: int, theta: tirx.Var, scale: tirx.Var, indices: tuple[tirx.Var, ...], qkv_dtype: str, rope_scaling: dict[str, Any]): +def _rope(buffer: T.Tensor, offset: tirx.Var, rotary_dim: int, theta: tirx.Var, scale: tirx.Var, indices: tuple[tirx.Var, ...], qkv_dtype: str, rope_scaling: dict[str, Any]): d = indices[-1] cos_freq, sin_freq, var_map = switch_rope_freq_func(rope_scaling)(offset * scale, d, rotary_dim, theta, "float32") cos = cos_freq * buffer[indices].astype("float32") @@ -138,9 +138,9 @@ def _causal_or_sliding_cross_mask(causal, row, col, kv_len, qo_len, sliding_wind def _length_info_buffer(batch_size, sliding_window, elem_offset): return ( - T.Buffer( (3, batch_size), "int32", elem_offset=elem_offset) + T.Tensor( (3, batch_size), "int32", elem_offset=elem_offset) if sliding_window - else T.Buffer( (batch_size,), "int32", elem_offset=elem_offset) + else T.Tensor( (batch_size,), "int32", elem_offset=elem_offset) ) def _get_kv_chunk_len(num_pages, page_size, seq_id, length_info, sliding_window): @@ -206,7 +206,7 @@ def _make_prefill_macros(tile_x, tile_y, tile_z, tile_o, bdx, num_warps, group_s """ @T.macro def init_states( - m_smem: T.Buffer, d_smem: T.Buffer, O_local: T.Buffer, ty: T.int32, tx: T.int32, + m_smem: T.Tensor, d_smem: T.Tensor, O_local: T.Tensor, ty: T.int32, tx: T.int32, ): for i in T.serial(T.ceildiv(tile_x, bdx * num_warps)): row: T.let[T.int32] = i * bdx * num_warps + ty * bdx + tx @@ -221,7 +221,7 @@ def init_states( @T.macro def compute_s_gemm( - Q_smem: T.Buffer, K_smem: T.Buffer, S_local: T.Buffer, S_smem: T.Buffer, sm_scale: T.float32, + Q_smem: T.Tensor, K_smem: T.Tensor, S_local: T.Tensor, S_smem: T.Tensor, sm_scale: T.float32, ): with Ts.sblock(): for li, lj, lk in T.grid(tile_x, tile_z, tile_y): @@ -239,8 +239,8 @@ def compute_s_gemm( @T.macro def softmax_update_causal( - S_smem: T.Buffer, m_smem: T.Buffer, d_smem: T.Buffer, m_prev_smem: T.Buffer, - m_new: T.Buffer, m_prev: T.Buffer, d_new: T.Buffer, + S_smem: T.Tensor, m_smem: T.Tensor, d_smem: T.Tensor, m_prev_smem: T.Tensor, + m_new: T.Tensor, m_prev: T.Tensor, d_new: T.Tensor, ty: T.int32, tx: T.int32, LH_start: T.int32, L_kv_start: T.int32, causal: T.int32, kv_len: T.int32, qo_len: T.int32, sliding_window_size: T.int32, ): @@ -282,8 +282,8 @@ def softmax_update_causal( @T.macro def compute_o_gemm( - S_smem: T.Buffer, V_smem: T.Buffer, O_local: T.Buffer, - m_prev_smem: T.Buffer, m_smem: T.Buffer, + S_smem: T.Tensor, V_smem: T.Tensor, O_local: T.Tensor, + m_prev_smem: T.Tensor, m_smem: T.Tensor, ): with Ts.sblock(): for li, lj, lk in T.grid(tile_x, tile_o, tile_z): @@ -295,8 +295,8 @@ def compute_o_gemm( @T.macro def paged_store_output_lse( - output: T.Buffer, lse: T.Buffer, O_local: T.Buffer, m_smem: T.Buffer, d_smem: T.Buffer, - q_indptr: T.Buffer, b_idx: T.int32, by: T.int32, LH_start: T.int32, + output: T.Tensor, lse: T.Tensor, O_local: T.Tensor, m_smem: T.Tensor, d_smem: T.Tensor, + q_indptr: T.Tensor, b_idx: T.int32, by: T.int32, LH_start: T.int32, ): """Paged-style (q_indptr-based) O_store + lse_store epilogue. @@ -320,8 +320,8 @@ def paged_store_output_lse( @T.macro def advance_tile_batch( - tile_id: T.Buffer, batch_idx: T.Buffer, batch_tiles: T.Buffer, batch_rows: T.Buffer, - q_indptr: T.Buffer, batch_size: T.int32, + tile_id: T.Tensor, batch_idx: T.Tensor, batch_tiles: T.Tensor, batch_rows: T.Tensor, + q_indptr: T.Tensor, batch_size: T.int32, ): """Advance tile_id/batch_idx past exhausted batches. @@ -338,8 +338,8 @@ def advance_tile_batch( @T.macro def softmax_update_valid_length( - S_smem: T.Buffer, m_smem: T.Buffer, d_smem: T.Buffer, m_prev_smem: T.Buffer, - m_new: T.Buffer, m_prev: T.Buffer, d_new: T.Buffer, + S_smem: T.Tensor, m_smem: T.Tensor, d_smem: T.Tensor, m_prev_smem: T.Tensor, + m_new: T.Tensor, m_prev: T.Tensor, d_new: T.Tensor, ty: T.int32, tx: T.int32, LH_start: T.int32, L_kv_start: T.int32, valid_len: T.int32, qo_len: T.int32, kv_len: T.int32, ): @@ -379,8 +379,8 @@ def softmax_update_valid_length( @T.macro def softmax_update_causal_padded_left( - S_smem: T.Buffer, m_smem: T.Buffer, d_smem: T.Buffer, m_prev_smem: T.Buffer, - m_new: T.Buffer, m_prev: T.Buffer, d_new: T.Buffer, + S_smem: T.Tensor, m_smem: T.Tensor, d_smem: T.Tensor, m_prev_smem: T.Tensor, + m_new: T.Tensor, m_prev: T.Tensor, d_new: T.Tensor, ty: T.int32, tx: T.int32, LH_start: T.int32, L_kv_start: T.int32, valid_len: T.int32, qo_len: T.int32, kv_len: T.int32, ): diff --git a/python/tvm/relax/frontend/nn/llm/_page_kernels.py b/python/tvm/relax/frontend/nn/llm/_page_kernels.py index ff53afa5361b..bd754d2945fd 100644 --- a/python/tvm/relax/frontend/nn/llm/_page_kernels.py +++ b/python/tvm/relax/frontend/nn/llm/_page_kernels.py @@ -48,10 +48,10 @@ def _kv_cache_transpose_append(num_key_value_heads, head_dim, dtype, page_size: position_map_elem_offset = T.dynamic("position_map_elem_offset", "int32") @Ts.prim_func def tir_kv_cache_transpose_append( - pages: T.Buffer((num_pages, 2, num_key_value_heads, page_size, head_dim), dtype, elem_offset=pages_elem_offset), - k_data: T.Buffer((ntoken, num_key_value_heads, head_dim), dtype), - v_data: T.Buffer((ntoken, num_key_value_heads, head_dim), dtype), - position_map: T.Buffer((ntoken,), 'int32', elem_offset=position_map_elem_offset), + pages: T.Tensor((num_pages, 2, num_key_value_heads, page_size, head_dim), dtype, elem_offset=pages_elem_offset), + k_data: T.Tensor((ntoken, num_key_value_heads, head_dim), dtype), + v_data: T.Tensor((ntoken, num_key_value_heads, head_dim), dtype), + position_map: T.Tensor((ntoken,), 'int32', elem_offset=position_map_elem_offset), ): T.func_attr({"tirx.noalias": True}) @@ -81,9 +81,9 @@ def _kv_cache_transpose_append_mla(d_qk: int, dtype, page_size: int = 16): position_map_elem_offset = T.dynamic("position_map_elem_offset", "int32") @Ts.prim_func def tir_kv_cache_transpose_append_mla( - pages: T.Buffer((num_pages, page_size, d_qk), dtype, elem_offset=pages_elem_offset), - kv_data: T.Buffer((ntoken, d_qk), dtype), - position_map: T.Buffer((ntoken,), 'int32', elem_offset=position_map_elem_offset), + pages: T.Tensor((num_pages, page_size, d_qk), dtype, elem_offset=pages_elem_offset), + kv_data: T.Tensor((ntoken, d_qk), dtype), + position_map: T.Tensor((ntoken,), 'int32', elem_offset=position_map_elem_offset), ): T.func_attr({"tirx.noalias": True}) @@ -108,10 +108,10 @@ def _kv_cache_debug_get_kv(num_hidden_layers, num_key_value_heads, head_dim, dty position_map_elem_offset = T.dynamic("position_map_elem_offset") @Ts.prim_func def tir_kv_cache_debug_get_kv( - pages: T.Buffer((num_pages, 2, num_key_value_heads, page_size, head_dim), dtype, elem_offset=pages_elem_offset), - position_map: T.Buffer((seqlen,), 'int32', elem_offset=position_map_elem_offset), - k_data: T.Buffer((num_hidden_layers, seqlen, num_key_value_heads, head_dim), dtype), - v_data: T.Buffer((num_hidden_layers, seqlen, num_key_value_heads, head_dim), dtype), + pages: T.Tensor((num_pages, 2, num_key_value_heads, page_size, head_dim), dtype, elem_offset=pages_elem_offset), + position_map: T.Tensor((seqlen,), 'int32', elem_offset=position_map_elem_offset), + k_data: T.Tensor((num_hidden_layers, seqlen, num_key_value_heads, head_dim), dtype), + v_data: T.Tensor((num_hidden_layers, seqlen, num_key_value_heads, head_dim), dtype), layer_id: T.int64, ): T.func_attr({"tirx.noalias": True}) @@ -137,9 +137,9 @@ def _kv_cache_debug_get_kv_mla(num_hidden_layers, d_qk, dtype): position_map_elem_offset = T.dynamic("position_map_elem_offset") @Ts.prim_func def tir_kv_cache_debug_get_kv_mla( - pages: T.Buffer((num_pages, page_size, d_qk), dtype, elem_offset=pages_elem_offset), - position_map: T.Buffer((seqlen,), 'int32', elem_offset=position_map_elem_offset), - compressed_kv_with_k_pe_data: T.Buffer((num_hidden_layers, seqlen, d_qk), dtype), + pages: T.Tensor((num_pages, page_size, d_qk), dtype, elem_offset=pages_elem_offset), + position_map: T.Tensor((seqlen,), 'int32', elem_offset=position_map_elem_offset), + compressed_kv_with_k_pe_data: T.Tensor((num_hidden_layers, seqlen, d_qk), dtype), layer_id: T.int64, ): T.func_attr({"tirx.noalias": True}) @@ -160,7 +160,7 @@ def _copy_single_page(num_heads, page_size, head_dim, dtype, target: Target): num_pages = T.dynamic("num_pages", "int32") pages_elem_offset = T.dynamic("pages_elem_offset") @Ts.prim_func - def copy_single_page(pages: T.Buffer((num_pages, 2, num_heads, page_size, head_dim), dtype, elem_offset=pages_elem_offset), src_page_id: T.int64, tgt_page_id: T.int64, copy_length: T.int64): + def copy_single_page(pages: T.Tensor((num_pages, 2, num_heads, page_size, head_dim), dtype, elem_offset=pages_elem_offset), src_page_id: T.int64, tgt_page_id: T.int64, copy_length: T.int64): T.func_attr({"tirx.is_scheduled": True}) for b in T.thread_binding((copy_length * num_heads * head_dim + tx - 1) // tx, thread="blockIdx.x"): @@ -181,7 +181,7 @@ def _copy_single_page_mla(page_size, head_dim, dtype, target: Target): num_pages = T.dynamic("num_pages", "int32") pages_elem_offset = T.dynamic("pages_elem_offset") @Ts.prim_func - def copy_single_page_mla(pages: T.Buffer((num_pages, page_size, head_dim), dtype, elem_offset=pages_elem_offset), src_page_id: T.int64, tgt_page_id: T.int64, copy_length: T.int64): + def copy_single_page_mla(pages: T.Tensor((num_pages, page_size, head_dim), dtype, elem_offset=pages_elem_offset), src_page_id: T.int64, tgt_page_id: T.int64, copy_length: T.int64): T.func_attr({"tirx.is_scheduled": True}) for b in T.thread_binding((copy_length * head_dim + tx - 1) // tx, thread="blockIdx.x"): @@ -199,7 +199,7 @@ def _copy_single_page_cpu(num_heads, page_size, head_dim, dtype): num_pages = T.dynamic("num_pages", "int32") @Ts.prim_func - def copy_single_page_cpu(pages: T.Buffer((num_pages, 2, num_heads, page_size, head_dim), dtype), src_page_id: T.int64, tgt_page_id: T.int64, copy_length: T.int64): + def copy_single_page_cpu(pages: T.Tensor((num_pages, 2, num_heads, page_size, head_dim), dtype), src_page_id: T.int64, tgt_page_id: T.int64, copy_length: T.int64): T.func_attr({"tirx.is_scheduled": True}) for b in T.serial((copy_length * num_heads * head_dim + tx - 1) // tx): @@ -223,7 +223,7 @@ def _compact_kv_copy(num_heads, head_dim, dtype, target: Target, page_size: int 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(pages: T.Buffer((num_pages, 2, num_heads, page_size, head_dim), dtype, elem_offset=pages_elem_offset), copy_length_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=copy_length_indptr_elem_offset), copy_src_dst_pos: T.Buffer((2, total_copy_length), 'int32', elem_offset=copy_src_dst_pos_elem_offset), batch_size: T.int32): # noqa: F821 + def compact_kv_copy(pages: T.Tensor((num_pages, 2, num_heads, page_size, head_dim), dtype, elem_offset=pages_elem_offset), copy_length_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=copy_length_indptr_elem_offset), copy_src_dst_pos: T.Tensor((2, total_copy_length), 'int32', elem_offset=copy_src_dst_pos_elem_offset), batch_size: T.int32): # noqa: F821 T.func_attr({"tirx.is_scheduled": True}) with Ts.sblock("root"): @@ -249,7 +249,7 @@ def _compact_kv_copy_cpu(num_heads, head_dim, dtype, page_size: int = 16): 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(pages: T.Buffer((num_pages, 2, num_heads, page_size, head_dim), dtype), copy_length_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=copy_length_indptr_elem_offset), copy_src_dst_pos: T.Buffer((2, total_copy_length), 'int32', elem_offset=copy_src_dst_pos_elem_offset), batch_size: T.int32): # noqa: F821 + def compact_kv_copy_cpu(pages: T.Tensor((num_pages, 2, num_heads, page_size, head_dim), dtype), copy_length_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=copy_length_indptr_elem_offset), copy_src_dst_pos: T.Tensor((2, total_copy_length), 'int32', elem_offset=copy_src_dst_pos_elem_offset), batch_size: T.int32): # noqa: F821 T.func_attr({"tirx.is_scheduled": True}) with Ts.sblock("root"): diff --git a/python/tvm/relax/frontend/nn/llm/_prefill_kernels.py b/python/tvm/relax/frontend/nn/llm/_prefill_kernels.py index e5232b39ffde..b5122615cc06 100644 --- a/python/tvm/relax/frontend/nn/llm/_prefill_kernels.py +++ b/python/tvm/relax/frontend/nn/llm/_prefill_kernels.py @@ -81,16 +81,16 @@ def _attention_prefill_cpu( length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") @Ts.prim_func def batch_prefill_paged_kv_cpu( - q: T.Buffer((total_len, h_q, d), dtype), # [total_len, h_q, d] - q_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] - pages: T.Buffer((max_num_pages, 2, h_kv, page_size, d), dtype), # [max_num_pages, 2, h_kv, page_size, d] - page_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=page_indptr_elem_offset), # [batch_size + 1] - page_values: T.Buffer((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), # [nnz_pages] + q: T.Tensor((total_len, h_q, d), dtype), # [total_len, h_q, d] + q_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] + pages: T.Tensor((max_num_pages, 2, h_kv, page_size, d), dtype), # [max_num_pages, 2, h_kv, page_size, d] + page_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=page_indptr_elem_offset), # [batch_size + 1] + page_values: T.Tensor((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), # [nnz_pages] length_info: _length_info_buffer(batch_size, sliding_window, length_info_elem_offset), # [b] when sliding window = False, or otherwise [3, b] - k_rope_pos_offset: T.Buffer((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] - q_rope_position: T.Buffer((total_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_len] - output: T.Buffer((total_len, h_q, d), dtype), # [total_len, h_q, d] - lse: T.Buffer((total_len, h_q), 'float32'), # [total_len, h_q] + k_rope_pos_offset: T.Tensor((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] + q_rope_position: T.Tensor((total_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_len] + output: T.Tensor((total_len, h_q, d), dtype), # [total_len, h_q, d] + lse: T.Tensor((total_len, h_q), 'float32'), # [total_len, h_q] causal: T.int32, rotary_mode: T.int32, rope_scale: T.float32, @@ -235,16 +235,16 @@ def _attention_prefill( length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") @Ts.prim_func def batch_prefill_paged_kv( - q: T.Buffer((total_len, h_q, d), dtype), # [total_len, h_q, d] - q_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] - pages: T.Buffer((max_num_pages, 2, h_kv, page_size, d), dtype, elem_offset=pages_elem_offset), # [max_num_pages, 2, h_kv, page_size, d] - page_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=page_indptr_elem_offset), # [batch_size + 1] - page_values: T.Buffer((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), # [nnz_pages] + q: T.Tensor((total_len, h_q, d), dtype), # [total_len, h_q, d] + q_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] + pages: T.Tensor((max_num_pages, 2, h_kv, page_size, d), dtype, elem_offset=pages_elem_offset), # [max_num_pages, 2, h_kv, page_size, d] + page_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=page_indptr_elem_offset), # [batch_size + 1] + page_values: T.Tensor((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), # [nnz_pages] length_info: _length_info_buffer(batch_size, sliding_window, length_info_elem_offset), # [b] when sliding window = False, or otherwise [3, b] - k_rope_pos_offset: T.Buffer((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] - q_rope_position: T.Buffer((total_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_len] - output: T.Buffer((total_len, h_q, d), dtype), # [total_len, h_q, d] - lse: T.Buffer((total_len, h_q), 'float32'), # [total_len, h_q] + k_rope_pos_offset: T.Tensor((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] + q_rope_position: T.Tensor((total_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_len] + output: T.Tensor((total_len, h_q, d), dtype), # [total_len, h_q, d] + lse: T.Tensor((total_len, h_q), 'float32'), # [total_len, h_q] causal: T.int32, rotary_mode: T.int32, rope_scale: T.float32, @@ -378,11 +378,11 @@ def _attention_sequence_prefill(h_kv, h_q, d, dtype, target: Target, causal=0, s kv_len = T.dynamic("kv_len", "int32") @Ts.prim_func def batch_sequence_prefill_kv( # pylint: disable=too-many-branches - q: T.Buffer((batch_size, qo_len, h_q, d), dtype), # [total_len, h_q, d] - k: T.Buffer((batch_size, kv_len, h_kv, d), dtype), # [total_len, h_kv, d] - v: T.Buffer((batch_size, kv_len, h_kv, d), dtype), # [total_len, h_kv, d] - output: T.Buffer((batch_size, qo_len, h_q, d), dtype), # [total_len, h_q, d] - lse: T.Buffer((batch_size, qo_len, h_q), dtype) # [total_len, h_q] + q: T.Tensor((batch_size, qo_len, h_q, d), dtype), # [total_len, h_q, d] + k: T.Tensor((batch_size, kv_len, h_kv, d), dtype), # [total_len, h_kv, d] + v: T.Tensor((batch_size, kv_len, h_kv, d), dtype), # [total_len, h_kv, d] + output: T.Tensor((batch_size, qo_len, h_q, d), dtype), # [total_len, h_q, d] + lse: T.Tensor((batch_size, qo_len, h_q), dtype) # [total_len, h_q] ): # pylint: disable=unused-variable @@ -540,12 +540,12 @@ def _kv_col_valid(col, valid_len, kv_len): kv_len = T.dynamic("kv_len", "int32") @Ts.prim_func def batch_sequence_prefill_kv_masked( # pylint: disable=too-many-branches - q: T.Buffer((batch_size, qo_len, h_q, d), dtype), # [batch_size, qo_len, h_q, d] - k: T.Buffer((batch_size, kv_len, h_kv, d), dtype), # [batch_size, kv_len, h_kv, d] - v: T.Buffer((batch_size, kv_len, h_kv, d), dtype), # [batch_size, kv_len, h_kv, d] - valid_lens: T.Buffer((batch_size,), 'int32'), # [batch_size], int32 - output: T.Buffer((batch_size, qo_len, h_q, d), dtype), # [batch_size, qo_len, h_q, d] - lse: T.Buffer((batch_size, qo_len, h_q), dtype) # [batch_size, qo_len, h_q] + q: T.Tensor((batch_size, qo_len, h_q, d), dtype), # [batch_size, qo_len, h_q, d] + k: T.Tensor((batch_size, kv_len, h_kv, d), dtype), # [batch_size, kv_len, h_kv, d] + v: T.Tensor((batch_size, kv_len, h_kv, d), dtype), # [batch_size, kv_len, h_kv, d] + valid_lens: T.Tensor((batch_size,), 'int32'), # [batch_size], int32 + output: T.Tensor((batch_size, qo_len, h_q, d), dtype), # [batch_size, qo_len, h_q, d] + lse: T.Tensor((batch_size, qo_len, h_q), dtype) # [batch_size, qo_len, h_q] ): batch_tiles: T.let[T.int32] = T.ceildiv(qo_len * group_size, tile_x) @@ -650,15 +650,15 @@ def _attention_prefill_ragged_cpu(h_kv, h_q, d_qk, d_v, dtype, rope_scaling: dic 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 - q: T.Buffer((qo_len, h_q, d_qk), dtype), # [total_len, h_q, d_qk] - q_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] - k: T.Buffer((kv_len, h_kv, d_qk), dtype), # [total_len, h_kv, d_qk] - v: T.Buffer((kv_len, h_kv, d_v), dtype), # [total_len, h_kv, d_v] - kv_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=kv_indptr_elem_offset), # [batch_size + 1] - q_rope_position: T.Buffer((qo_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_q_len] - k_rope_pos_offset: T.Buffer((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] - output: T.Buffer((qo_len, h_q, d_v), dtype), # [total_len, h_q, d_v] - lse: T.Buffer((qo_len, h_q), 'float32'), # [total_len, h_q] + q: T.Tensor((qo_len, h_q, d_qk), dtype), # [total_len, h_q, d_qk] + q_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] + k: T.Tensor((kv_len, h_kv, d_qk), dtype), # [total_len, h_kv, d_qk] + v: T.Tensor((kv_len, h_kv, d_v), dtype), # [total_len, h_kv, d_v] + kv_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=kv_indptr_elem_offset), # [batch_size + 1] + q_rope_position: T.Tensor((qo_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_q_len] + k_rope_pos_offset: T.Tensor((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] + output: T.Tensor((qo_len, h_q, d_v), dtype), # [total_len, h_q, d_v] + lse: T.Tensor((qo_len, h_q), 'float32'), # [total_len, h_q] causal: T.int32, rotary_mode: T.int32, rope_scale: T.float32, @@ -759,15 +759,15 @@ def _attention_prefill_ragged(h_kv, h_q, d_qk, d_v, dtype, rope_scaling: dict[st 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 - q: T.Buffer((qo_len, h_q, d_qk), dtype), # [total_len, h_q, d_qk] - q_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] - k: T.Buffer((kv_len, h_kv, d_qk), dtype), # [total_len, h_kv, d_qk] - v: T.Buffer((kv_len, h_kv, d_v), dtype), # [total_len, h_kv, d_v] - kv_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=kv_indptr_elem_offset), # [batch_size + 1] - q_rope_position: T.Buffer((qo_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_q_len] - k_rope_pos_offset: T.Buffer((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] - output: T.Buffer((qo_len, h_q, d_v), dtype), # [total_len, h_q, d_v] - lse: T.Buffer((qo_len, h_q), 'float32'), # [total_len, h_q] + q: T.Tensor((qo_len, h_q, d_qk), dtype), # [total_len, h_q, d_qk] + q_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] + k: T.Tensor((kv_len, h_kv, d_qk), dtype), # [total_len, h_kv, d_qk] + v: T.Tensor((kv_len, h_kv, d_v), dtype), # [total_len, h_kv, d_v] + kv_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=kv_indptr_elem_offset), # [batch_size + 1] + q_rope_position: T.Tensor((qo_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_q_len] + k_rope_pos_offset: T.Tensor((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] + output: T.Tensor((qo_len, h_q, d_v), dtype), # [total_len, h_q, d_v] + lse: T.Tensor((qo_len, h_q), 'float32'), # [total_len, h_q] causal: T.int32, rotary_mode: T.int32, rope_scale: T.float32, @@ -889,14 +889,14 @@ def _attention_prefill_mla(h_q, d_latent, d_rope, dtype, sliding_window: bool, t length_info_elem_offset = T.dynamic("length_info_elem_offset", "int32") @Ts.prim_func def batch_prefill_paged_kv_mla( - q: T.Buffer((total_len, h_q, d_qk), dtype), # [total_len, h_q, d_qk] - q_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] - pages: T.Buffer((max_num_pages, page_size, d_qk), dtype, elem_offset=pages_elem_offset), # [max_num_pages, page_size, d_qk] - page_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=page_indptr_elem_offset), # [batch_size + 1] - page_values: T.Buffer((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), # [nnz_pages] + q: T.Tensor((total_len, h_q, d_qk), dtype), # [total_len, h_q, d_qk] + q_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] + pages: T.Tensor((max_num_pages, page_size, d_qk), dtype, elem_offset=pages_elem_offset), # [max_num_pages, page_size, d_qk] + page_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=page_indptr_elem_offset), # [batch_size + 1] + page_values: T.Tensor((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), # [nnz_pages] length_info: _length_info_buffer(batch_size, sliding_window, length_info_elem_offset), # [b] when sliding window = False, or otherwise [3, b] - output: T.Buffer((total_len, h_q, d_latent), dtype), # [total_len, h_q, d_latent] - lse: T.Buffer((total_len, h_q), 'float32'), # [total_len, h_q] + output: T.Tensor((total_len, h_q, d_latent), dtype), # [total_len, h_q, d_latent] + lse: T.Tensor((total_len, h_q), 'float32'), # [total_len, h_q] causal: T.int32, sm_scale: T.float32, ): diff --git a/python/tvm/relax/frontend/nn/llm/position_embedding.py b/python/tvm/relax/frontend/nn/llm/position_embedding.py index 511a2bf70cec..b1d5e3b38f31 100644 --- a/python/tvm/relax/frontend/nn/llm/position_embedding.py +++ b/python/tvm/relax/frontend/nn/llm/position_embedding.py @@ -170,7 +170,7 @@ def rope_freq_longrope( # pylint: disable=too-many-arguments dtype: str, max_position_embeddings: int, original_max_position_embeddings: int, - ext_factors: T.Buffer | None = None, + ext_factors: T.Tensor | None = None, ): """Compute the inverse frequency of RoPE for longrope scaling.""" scale = max_position_embeddings / original_max_position_embeddings @@ -363,7 +363,7 @@ def llama_rope( # pylint: disable=too-many-arguments scale = tirx.const(scale, dtype) def _rope( # pylint: disable=too-many-arguments - x: T.Buffer, + x: T.Tensor, b: tirx.Var, s: tirx.Var, h: tirx.Var, @@ -396,10 +396,10 @@ def _rope( # pylint: disable=too-many-arguments @Ts.prim_func(private=True) def fused_rope( # pylint: disable=too-many-locals - qkv: T.Buffer((batch_size, seq_len, fused_heads, head_dim), dtype), - q: T.Buffer((batch_size, seq_len, num_q_heads, head_dim), dtype), - k: T.Buffer((batch_size, seq_len, num_kv_heads, head_dim), dtype), - v: T.Buffer((batch_size, seq_len, num_kv_heads, head_dim), dtype), + qkv: T.Tensor((batch_size, seq_len, fused_heads, head_dim), dtype), + q: T.Tensor((batch_size, seq_len, num_q_heads, head_dim), dtype), + k: T.Tensor((batch_size, seq_len, num_kv_heads, head_dim), dtype), + v: T.Tensor((batch_size, seq_len, num_kv_heads, head_dim), dtype), total_seq_len: T.int64, ): T.func_attr( @@ -490,12 +490,12 @@ def llama_rope_with_position_map( # pylint: disable=too-many-arguments original_max_position_embeddings = 0 def _rope( # pylint: disable=too-many-arguments - x: T.Buffer, + x: T.Tensor, s: tirx.Var, h: tirx.Var, d: tirx.Var, pos: tirx.Var, - ext_factors: T.Buffer | None = None, + ext_factors: T.Tensor | None = None, ): kwargs = {} if ext_factors is not None: @@ -526,11 +526,11 @@ def _rope( # pylint: disable=too-many-arguments @Ts.prim_func def fused_rope( # pylint: disable=too-many-locals - qkv: T.Buffer((seq_len, fused_heads, head_dim), dtype), - position_map: T.Buffer((seq_len,), "int32", elem_offset=position_map_elem_offset), - q: T.Buffer((seq_len, num_q_heads, head_dim), dtype), - k: T.Buffer((seq_len, num_kv_heads, head_dim), dtype), - v: T.Buffer((seq_len, num_kv_heads, head_dim), dtype), + qkv: T.Tensor((seq_len, fused_heads, head_dim), dtype), + position_map: T.Tensor((seq_len,), "int32", elem_offset=position_map_elem_offset), + q: T.Tensor((seq_len, num_q_heads, head_dim), dtype), + k: T.Tensor((seq_len, num_kv_heads, head_dim), dtype), + v: T.Tensor((seq_len, num_kv_heads, head_dim), dtype), apply_rope: T.int64, ): T.func_attr( @@ -563,12 +563,12 @@ def fused_rope( # pylint: disable=too-many-locals @Ts.prim_func def fused_rope_longrope_scaling( # pylint: disable=too-many-locals - qkv: T.Buffer((seq_len, fused_heads, head_dim), dtype), - position_map: T.Buffer((seq_len,), "int32", elem_offset=position_map_elem_offset), - q: T.Buffer((seq_len, num_q_heads, head_dim), dtype), - k: T.Buffer((seq_len, num_kv_heads, head_dim), dtype), - v: T.Buffer((seq_len, num_kv_heads, head_dim), dtype), - ext_factors: T.Buffer((rotary_dim,), "float32"), # type: ignore + qkv: T.Tensor((seq_len, fused_heads, head_dim), dtype), + position_map: T.Tensor((seq_len,), "int32", elem_offset=position_map_elem_offset), + q: T.Tensor((seq_len, num_q_heads, head_dim), dtype), + k: T.Tensor((seq_len, num_kv_heads, head_dim), dtype), + v: T.Tensor((seq_len, num_kv_heads, head_dim), dtype), + ext_factors: T.Tensor((rotary_dim,), "float32"), # type: ignore ): T.func_attr( { @@ -578,8 +578,8 @@ def fused_rope_longrope_scaling( # pylint: disable=too-many-locals ) # long factors is the first half, short factors is the second half - long_factors = T.decl_buffer((rotary_dim // 2,), "float32", data=ext_factors.data) - short_factors = T.decl_buffer( + long_factors = T.decl_tensor((rotary_dim // 2,), "float32", data=ext_factors.data) + short_factors = T.decl_tensor( (rotary_dim // 2,), "float32", data=ext_factors.data, @@ -706,12 +706,12 @@ def llama4_rope_with_position_map( # pylint: disable=too-many-arguments original_max_position_embeddings = 0 def _rope( # pylint: disable=too-many-arguments - x: T.Buffer, + x: T.Tensor, s: tirx.Var, h: tirx.Var, d: tirx.Var, pos: tirx.Var, - ext_factors: T.Buffer | None = None, + ext_factors: T.Tensor | None = None, ): kwargs = {} if ext_factors is not None: @@ -743,11 +743,11 @@ def _rope( # pylint: disable=too-many-arguments @Ts.prim_func(private=True) def fused_rope( # pylint: disable=too-many-locals - qkv: T.Buffer((seq_len, fused_heads, head_dim), dtype), - position_map: T.Buffer((seq_len,), "int32", elem_offset=position_map_elem_offset), - q: T.Buffer((seq_len, num_q_heads, head_dim), dtype), - k: T.Buffer((seq_len, num_kv_heads, head_dim), dtype), - v: T.Buffer((seq_len, num_kv_heads, head_dim), dtype), + qkv: T.Tensor((seq_len, fused_heads, head_dim), dtype), + position_map: T.Tensor((seq_len,), "int32", elem_offset=position_map_elem_offset), + q: T.Tensor((seq_len, num_q_heads, head_dim), dtype), + k: T.Tensor((seq_len, num_kv_heads, head_dim), dtype), + v: T.Tensor((seq_len, num_kv_heads, head_dim), dtype), apply_rope: T.int64, ): T.func_attr( @@ -780,12 +780,12 @@ def fused_rope( # pylint: disable=too-many-locals @Ts.prim_func def fused_rope_longrope_scaling( # pylint: disable=too-many-locals - qkv: T.Buffer((seq_len, fused_heads, head_dim), dtype), - position_map: T.Buffer((seq_len,), "int32", elem_offset=position_map_elem_offset), - q: T.Buffer((seq_len, num_q_heads, head_dim), dtype), - k: T.Buffer((seq_len, num_kv_heads, head_dim), dtype), - v: T.Buffer((seq_len, num_kv_heads, head_dim), dtype), - ext_factors: T.Buffer((rotary_dim,), "float32"), # type: ignore + qkv: T.Tensor((seq_len, fused_heads, head_dim), dtype), + position_map: T.Tensor((seq_len,), "int32", elem_offset=position_map_elem_offset), + q: T.Tensor((seq_len, num_q_heads, head_dim), dtype), + k: T.Tensor((seq_len, num_kv_heads, head_dim), dtype), + v: T.Tensor((seq_len, num_kv_heads, head_dim), dtype), + ext_factors: T.Tensor((rotary_dim,), "float32"), # type: ignore ): T.func_attr( { @@ -795,8 +795,8 @@ def fused_rope_longrope_scaling( # pylint: disable=too-many-locals ) # long factors is the first half, short factors is the second half - long_factors = T.decl_buffer((rotary_dim // 2,), "float32", data=ext_factors.data) - short_factors = T.decl_buffer( + long_factors = T.decl_tensor((rotary_dim // 2,), "float32", data=ext_factors.data) + short_factors = T.decl_tensor( (rotary_dim // 2,), "float32", data=ext_factors.data, diff --git a/python/tvm/relax/frontend/nn/llm/tree_attn.py b/python/tvm/relax/frontend/nn/llm/tree_attn.py index 4473a8406a7b..a49b8bff30a7 100644 --- a/python/tvm/relax/frontend/nn/llm/tree_attn.py +++ b/python/tvm/relax/frontend/nn/llm/tree_attn.py @@ -100,16 +100,16 @@ def tree_attn_cpu(h_kv, h_q, d, dtype, rope_scaling: dict[str, Any]): 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 - q: T.Buffer((qo_len, h_q, d), dtype), # [total_len, h_q, d] - q_indptr: T.Buffer((batch_size_plus_1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] - k: T.Buffer((kv_len, h_kv, d), dtype), # [total_len, h_kv, d] - v: T.Buffer((kv_len, h_kv, d), dtype), # [total_len, h_kv, d] - kv_indptr: T.Buffer((batch_size_plus_1,), 'int32', elem_offset=kv_indptr_elem_offset), # [batch_size + 1], kv_indptr should be the same as q_indptr in this case - q_rope_position: T.Buffer((qo_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_q_len] - mn_indptr: T.Buffer((batch_size_plus_1,), 'int32', elem_offset=mn_indptr_elem_offset), # [batch_size + 1] - mask: T.Buffer((tree_size, 2), 'int32', elem_offset=mask_elem_offset), # [mn_indptr[batch_size]] - output: T.Buffer((qo_len, h_q, d), dtype), # [total_len, h_q, d] - lse: T.Buffer((qo_len, h_q), 'float32'), # [total_len, h_q] + q: T.Tensor((qo_len, h_q, d), dtype), # [total_len, h_q, d] + q_indptr: T.Tensor((batch_size_plus_1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] + k: T.Tensor((kv_len, h_kv, d), dtype), # [total_len, h_kv, d] + v: T.Tensor((kv_len, h_kv, d), dtype), # [total_len, h_kv, d] + kv_indptr: T.Tensor((batch_size_plus_1,), 'int32', elem_offset=kv_indptr_elem_offset), # [batch_size + 1], kv_indptr should be the same as q_indptr in this case + q_rope_position: T.Tensor((qo_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_q_len] + mn_indptr: T.Tensor((batch_size_plus_1,), 'int32', elem_offset=mn_indptr_elem_offset), # [batch_size + 1] + mask: T.Tensor((tree_size, 2), 'int32', elem_offset=mask_elem_offset), # [mn_indptr[batch_size]] + output: T.Tensor((qo_len, h_q, d), dtype), # [total_len, h_q, d] + lse: T.Tensor((qo_len, h_q), 'float32'), # [total_len, h_q] rotary_mode: T.int32, rope_scale: T.float32, rope_theta: T.float32, @@ -282,16 +282,16 @@ def tree_attn(h_kv, h_q, d, dtype, rope_scaling: dict[str, Any], target: Target) batch_size_plus_1 = T.dynamic("batch_size_plus_1", "int32") @Ts.prim_func def batch_tree_attn( # pylint: disable=too-many-branches - q: T.Buffer((qo_len, h_q, d), dtype), # [total_len, h_q, d] - q_indptr: T.Buffer((batch_size_plus_1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] - k: T.Buffer((kv_len, h_kv, d), dtype), # [total_len, h_kv, d] - v: T.Buffer((kv_len, h_kv, d), dtype), # [total_len, h_kv, d] - kv_indptr: T.Buffer((batch_size_plus_1,), 'int32', elem_offset=kv_indptr_elem_offset), # [batch_size + 1], kv_indptr should be the same as q_indptr in this case - q_rope_position: T.Buffer((qo_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_q_len] - mn_indptr: T.Buffer((batch_size_plus_1,), 'int32', elem_offset=mn_indptr_elem_offset), # [batch_size + 1] - mask: T.Buffer((tree_size, 2), 'int32', elem_offset=mask_elem_offset), # [mn_indptr[batch_size]] - output: T.Buffer((qo_len, h_q, d), dtype), # [total_len, h_q, d] - lse: T.Buffer((qo_len, h_q), 'float32'), # [total_len, h_q] + q: T.Tensor((qo_len, h_q, d), dtype), # [total_len, h_q, d] + q_indptr: T.Tensor((batch_size_plus_1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] + k: T.Tensor((kv_len, h_kv, d), dtype), # [total_len, h_kv, d] + v: T.Tensor((kv_len, h_kv, d), dtype), # [total_len, h_kv, d] + kv_indptr: T.Tensor((batch_size_plus_1,), 'int32', elem_offset=kv_indptr_elem_offset), # [batch_size + 1], kv_indptr should be the same as q_indptr in this case + q_rope_position: T.Tensor((qo_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_q_len] + mn_indptr: T.Tensor((batch_size_plus_1,), 'int32', elem_offset=mn_indptr_elem_offset), # [batch_size + 1] + mask: T.Tensor((tree_size, 2), 'int32', elem_offset=mask_elem_offset), # [mn_indptr[batch_size]] + output: T.Tensor((qo_len, h_q, d), dtype), # [total_len, h_q, d] + lse: T.Tensor((qo_len, h_q), 'float32'), # [total_len, h_q] rotary_mode: T.int32, rope_scale: T.float32, rope_theta: T.float32, @@ -597,22 +597,22 @@ def tree_attn_with_paged_kv_cache_cpu(h_kv, h_q, d, dtype, rope_scaling: dict[st total_tree_order_len = T.dynamic("total_tree_order_len", "int32") @Ts.prim_func def tree_attn_paged_kv_cpu( - q: T.Buffer((total_len, h_q, d), dtype), # [total_len, h_q, d] - q_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] - pages: T.Buffer((max_num_pages, 2, h_kv, 16, d), dtype), # [max_num_pages, 2, h_kv, page_size, d] - page_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=page_indptr_elem_offset), # [batch_size + 1] - page_values: T.Buffer((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), # [nnz_pages] + q: T.Tensor((total_len, h_q, d), dtype), # [total_len, h_q, d] + q_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] + pages: T.Tensor((max_num_pages, 2, h_kv, 16, d), dtype), # [max_num_pages, 2, h_kv, page_size, d] + page_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=page_indptr_elem_offset), # [batch_size + 1] + page_values: T.Tensor((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), # [nnz_pages] length_info: _length_info_buffer(batch_size, sliding_window, length_info_elem_offset), # [b] when sliding window = False, or otherwise [3, b] - k_rope_pos_offset: T.Buffer((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] - q_rope_position: T.Buffer((total_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_len] - output: T.Buffer((total_len, h_q, d), dtype), # [total_len, h_q, d] - lse: T.Buffer((total_len, h_q), 'float32'), # [total_len, h_q] + k_rope_pos_offset: T.Tensor((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] + q_rope_position: T.Tensor((total_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_len] + output: T.Tensor((total_len, h_q, d), dtype), # [total_len, h_q, d] + lse: T.Tensor((total_len, h_q), 'float32'), # [total_len, h_q] rotary_mode: T.int32, rope_scale: T.float32, rope_theta: T.float32, sm_scale: T.float32, - tree_order_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=tree_order_indptr_elem_offset), # [batch_size + 1] - tree_order: T.Buffer((total_tree_order_len, 2), 'int32', elem_offset=tree_order_elem_offset), # [total_len, 2] + tree_order_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=tree_order_indptr_elem_offset), # [batch_size + 1] + tree_order: T.Tensor((total_tree_order_len, 2), 'int32', elem_offset=tree_order_elem_offset), # [total_len, 2] ): T.func_attr({"global_symbol": global_symbol}) @@ -772,22 +772,22 @@ def tree_attn_with_paged_kv_cache( total_tree_order_len = T.dynamic("total_tree_order_len", "int32") @Ts.prim_func def tree_attn_paged_kv( - q: T.Buffer((total_len, h_q, d), dtype), # [total_len, h_q, d] - q_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] - pages: T.Buffer((max_num_pages, 2, h_kv, 16, d), dtype), # [max_num_pages, 2, h_kv, page_size, d] - page_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=page_indptr_elem_offset), # [batch_size + 1] - page_values: T.Buffer((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), # [nnz_pages] + q: T.Tensor((total_len, h_q, d), dtype), # [total_len, h_q, d] + q_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=q_indptr_elem_offset), # [batch_size + 1] + pages: T.Tensor((max_num_pages, 2, h_kv, 16, d), dtype), # [max_num_pages, 2, h_kv, page_size, d] + page_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=page_indptr_elem_offset), # [batch_size + 1] + page_values: T.Tensor((nnz_pages,), 'int32', elem_offset=page_values_elem_offset), # [nnz_pages] length_info: _length_info_buffer(batch_size, sliding_window, length_info_elem_offset), # [b] when sliding window = False, or otherwise [3, b] - k_rope_pos_offset: T.Buffer((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] - q_rope_position: T.Buffer((total_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_len] - output: T.Buffer((total_len, h_q, d), dtype), # [total_len, h_q, d] - lse: T.Buffer((total_len, h_q), 'float32'), # [total_len, h_q] + k_rope_pos_offset: T.Tensor((batch_size,), 'int32', elem_offset=k_rope_pos_offset_elem_offset), # [b] + q_rope_position: T.Tensor((total_len,), 'int32', elem_offset=q_rope_position_elem_offset), # [total_len] + output: T.Tensor((total_len, h_q, d), dtype), # [total_len, h_q, d] + lse: T.Tensor((total_len, h_q), 'float32'), # [total_len, h_q] rotary_mode: T.int32, rope_scale: T.float32, rope_theta: T.float32, sm_scale: T.float32, - tree_order_indptr: T.Buffer((batch_size + 1,), 'int32', elem_offset=tree_order_indptr_elem_offset), # [batch_size + 1] - tree_order: T.Buffer((total_tree_order_len, 2), 'int32', elem_offset=tree_order_elem_offset), # [total_len, 2] + tree_order_indptr: T.Tensor((batch_size + 1,), 'int32', elem_offset=tree_order_indptr_elem_offset), # [batch_size + 1] + tree_order: T.Tensor((total_tree_order_len, 2), 'int32', elem_offset=tree_order_elem_offset), # [total_len, 2] ): # pylint: disable=unused-variable, too-many-branches T.func_attr({"global_symbol": global_symbol}) diff --git a/python/tvm/relax/frontend/nn/op.py b/python/tvm/relax/frontend/nn/op.py index 6fd8b02cea5d..90a93e58e60d 100644 --- a/python/tvm/relax/frontend/nn/op.py +++ b/python/tvm/relax/frontend/nn/op.py @@ -2800,10 +2800,10 @@ def _cumsum_mask(cumsum_sorted, top_p, top_k, i, j): @Ts.prim_func(private=True) def _get_renorm_prob( - cumsum_sorted: T.Buffer((batch, vocab_size), prob_dtype), - top_p: T.Buffer((batch, 1), prob_dtype), - top_k: T.Buffer((batch, 1), index_dtype), - renorm_prob: T.Buffer((batch, 1), prob_dtype), + cumsum_sorted: T.Tensor((batch, vocab_size), prob_dtype), + top_p: T.Tensor((batch, 1), prob_dtype), + top_k: T.Tensor((batch, 1), index_dtype), + renorm_prob: T.Tensor((batch, 1), prob_dtype), ): for ax0, ax1 in T.grid(batch, vocab_size): with Ts.sblock("T_get_renorm_prob"): @@ -2822,12 +2822,12 @@ def _get_renorm_prob( @Ts.prim_func(private=True) def _get_index_from_sorted( - cumsum_sorted: T.Buffer((batch, vocab_size), prob_dtype), - indices: T.Buffer((batch, vocab_size), index_dtype), - renorm_prob: T.Buffer((batch, 1), prob_dtype), - usample: T.Buffer((kernel_out_batch, 1), prob_dtype), - sample_indices: T.Buffer((kernel_out_batch, 1), sample_indices_dtype), - output_index: T.Buffer((kernel_out_batch, 1), index_dtype), + cumsum_sorted: T.Tensor((batch, vocab_size), prob_dtype), + indices: T.Tensor((batch, vocab_size), index_dtype), + renorm_prob: T.Tensor((batch, 1), prob_dtype), + usample: T.Tensor((kernel_out_batch, 1), prob_dtype), + sample_indices: T.Tensor((kernel_out_batch, 1), sample_indices_dtype), + output_index: T.Tensor((kernel_out_batch, 1), index_dtype), ): for ax0, ax1 in T.grid(kernel_out_batch, vocab_size): with Ts.sblock("T_get_index_from_sorted"): @@ -2909,11 +2909,11 @@ def _cumsum_mask(cumsum_sorted, top_p, top_k, i, j): @Ts.prim_func(private=True) def _get_renorm_cutoff( - sorted_prob: T.Buffer((kernel_batch, vocab_size), prob_dtype), - cumsum_sorted: T.Buffer((kernel_batch, vocab_size), prob_dtype), - top_p: T.Buffer((kernel_batch, 1), prob_dtype), - top_k: T.Buffer((kernel_batch, 1), top_k_dtype), - cutoff: T.Buffer((kernel_batch, 1), prob_dtype), + sorted_prob: T.Tensor((kernel_batch, vocab_size), prob_dtype), + cumsum_sorted: T.Tensor((kernel_batch, vocab_size), prob_dtype), + top_p: T.Tensor((kernel_batch, 1), prob_dtype), + top_k: T.Tensor((kernel_batch, 1), top_k_dtype), + cutoff: T.Tensor((kernel_batch, 1), prob_dtype), ): for ax0, ax1 in T.grid(kernel_batch, vocab_size): with Ts.sblock("T_get_renorm_cutoff"): diff --git a/python/tvm/relax/frontend/tflite/tflite_frontend.py b/python/tvm/relax/frontend/tflite/tflite_frontend.py index 4f0543df02ee..c1828267dfe1 100644 --- a/python/tvm/relax/frontend/tflite/tflite_frontend.py +++ b/python/tvm/relax/frontend/tflite/tflite_frontend.py @@ -1242,7 +1242,7 @@ def _get_relax_tensor_shape(self, tensor): For COMPLEX64 tensors, the trailing (2,) axis encodes the real/imag pair. Returns an empty tuple () for rank-0 tensors. Shape elements are Python ints - (not numpy scalars) so the result is safe to feed into TIRX ``T.Buffer(shape, ...)``. + (not numpy scalars) so the result is safe to feed into TIRX ``T.Tensor(shape, ...)``. """ tensor = self._unwrap_tflite_tensor(tensor) shape = to_int_list(tensor.ShapeAsNumpy()) if tensor.ShapeLength() > 0 else () @@ -8332,13 +8332,13 @@ def _build_tflite_rfft2d_primfunc(input_shape, output_pair_shape): @Ts.prim_func(private=True, check_well_formed=False) def kernel( - data: T.Buffer(input_shape, "float32"), output: T.Buffer(output_pair_shape, "float32") + data: T.Tensor(input_shape, "float32"), output: T.Tensor(output_pair_shape, "float32") ): # Flat 1D aliases of the multi-dim buffers. The kernel is rank-agnostic # over the leading batch dimensions, so collapsing the index space # avoids special-casing 2D / 3D / 4D input shapes. - data_flat = T.decl_buffer((input_total,), "float32", data=data.data) - output_flat = T.decl_buffer((output_complex_total * 2,), "float32", data=output.data) + data_flat = T.decl_tensor((input_total,), "float32", data=data.data) + output_flat = T.decl_tensor((output_complex_total * 2,), "float32", data=output.data) neg_two_pi_const = T.float32(neg_two_pi) for b_idx, out_y, out_x in T.grid(batch, height, out_width): @@ -8506,13 +8506,13 @@ def _stage_stmts(stage_count, length, indent, stride=1, base_expr="row_base"): "from tvm.script import s_tir as Ts\n" "@Ts.prim_func(private=True, check_well_formed=False)\n" "def kernel(\n" - f" data: T.Buffer({tuple(int(x) for x in input_shape)}, 'float32'),\n" - f" output: T.Buffer({tuple(int(x) for x in output_pair_shape)}, 'float32'),\n" + f" data: T.Tensor({tuple(int(x) for x in input_shape)}, 'float32'),\n" + f" output: T.Tensor({tuple(int(x) for x in output_pair_shape)}, 'float32'),\n" "):\n" - f" data_flat = T.decl_buffer(({input_total},), 'float32', data=data.data)\n" - f" output_flat = T.decl_buffer(({output_complex_total * 2},), 'float32', data=output.data)\n" - f" scratch_real = T.decl_buffer(({input_total},), 'float32')\n" - f" scratch_imag = T.decl_buffer(({input_total},), 'float32')\n" + f" data_flat = T.decl_tensor(({input_total},), 'float32', data=data.data)\n" + f" output_flat = T.decl_tensor(({output_complex_total * 2},), 'float32', data=output.data)\n" + f" scratch_real = T.decl_tensor(({input_total},), 'float32')\n" + f" scratch_imag = T.decl_tensor(({input_total},), 'float32')\n" f" for b_idx in T.serial({batch}):\n" f" with Ts.sblock('rfft2d_fft'):\n" f" v_b = Ts.axis.remap('S', [b_idx])\n" @@ -8615,9 +8615,9 @@ def _store_value(words, write_index): @Ts.prim_func(private=True) def kernel( - initial_state: T.Buffer((state_len,), "uint64"), - output_state: T.Buffer((state_len,), "uint64"), - output: T.Buffer(out_shape, out_dtype), + initial_state: T.Tensor((state_len,), "uint64"), + output_state: T.Tensor((state_len,), "uint64"), + output: T.Tensor(out_shape, out_dtype), ): # A single opaque structured block keeps the imperative kernel as a # well-formed block-structured PrimFunc, as required by the Relax @@ -8629,10 +8629,10 @@ def kernel( key_1 = _u32(state_key >> T.uint64(32)) output_state[0] = state_key output_state[1] = state_counter + T.uint64(num_blocks) - out_flat = T.decl_buffer((total,), out_dtype, data=output.data) - keys = T.decl_buffer((3,), "uint32", scope="local") - rotations = T.decl_buffer((8,), "uint32", scope="local") - ctr = T.decl_buffer((2,), "uint32", scope="local") + out_flat = T.decl_tensor((total,), out_dtype, data=output.data) + keys = T.decl_tensor((3,), "uint32", scope="local") + rotations = T.decl_tensor((8,), "uint32", scope="local") + ctr = T.decl_tensor((2,), "uint32", scope="local") keys[0] = key_0 keys[1] = key_1 keys[2] = key_0 ^ key_1 ^ T.uint32(parity) @@ -8665,9 +8665,9 @@ def kernel( @Ts.prim_func(private=True) def kernel( - initial_state: T.Buffer((state_len,), "uint64"), - output_state: T.Buffer((state_len,), "uint64"), - output: T.Buffer(out_shape, out_dtype), + initial_state: T.Tensor((state_len,), "uint64"), + output_state: T.Tensor((state_len,), "uint64"), + output: T.Tensor(out_shape, out_dtype), ): with Ts.sblock("rng_bit_generator"): state_key = initial_state[0] @@ -8676,10 +8676,10 @@ def kernel( key_1 = _u32(state_key >> T.uint64(32)) output_state[0] = state_key output_state[1] = state_counter + T.uint64(num_blocks) - out_flat = T.decl_buffer((total,), out_dtype, data=output.data) - ctr = T.decl_buffer((4,), "uint32", scope="local") - keys = T.decl_buffer((2,), "uint32", scope="local") - high_ctr = T.decl_buffer((2,), "uint32", scope="local") + out_flat = T.decl_tensor((total,), out_dtype, data=output.data) + ctr = T.decl_tensor((4,), "uint32", scope="local") + keys = T.decl_tensor((2,), "uint32", scope="local") + high_ctr = T.decl_tensor((2,), "uint32", scope="local") if state_len == 3: # PHILOX u64[3]: the third state word feeds the high counter and # is passed through to the output state unchanged. diff --git a/python/tvm/relax/transform/legalize_ops/grad.py b/python/tvm/relax/transform/legalize_ops/grad.py index 573710a46d92..c501fb56c07b 100644 --- a/python/tvm/relax/transform/legalize_ops/grad.py +++ b/python/tvm/relax/transform/legalize_ops/grad.py @@ -232,7 +232,7 @@ def gen_ir(output_grad_ptr, x_ptr, indices_ptr, out_ptr): return ib.get() shape = x.shape - out_buf = tirx.decl_buffer(shape, x.dtype, "out_buf", layout=None) + out_buf = tirx.decl_tensor(shape, x.dtype, "out_buf", layout=None) return te.extern( [shape], diff --git a/python/tvm/relax/transform/legalize_ops/inspect_op.py b/python/tvm/relax/transform/legalize_ops/inspect_op.py index ce09203fe48e..db4ee3f46aa0 100644 --- a/python/tvm/relax/transform/legalize_ops/inspect_op.py +++ b/python/tvm/relax/transform/legalize_ops/inspect_op.py @@ -73,9 +73,9 @@ def _get_tensor_stride_i(dlpack_handle: T.handle, axis: T.int64) -> T.int64: shape_ptr: T.let[T.handle("int64")] = T.tvm_struct_get( dlpack_handle, 0, int(TVMStructFieldKind.kDLTensorShape), T.handle("int64").ty ) - shape = T.decl_buffer(ndim, "int64", data=shape_ptr) + shape = T.decl_tensor(ndim, "int64", data=shape_ptr) - product = T.decl_buffer([], "int64") + product = T.decl_tensor([], "int64") product[()] = 1 # TODO(Lunderberg): Add a TIR lowering pass to allow @@ -87,7 +87,7 @@ def _get_tensor_stride_i(dlpack_handle: T.handle, axis: T.int64) -> T.int64: return product[()] else: - strides = T.decl_buffer(ndim, "int64", data=stride_ptr) + strides = T.decl_tensor(ndim, "int64", data=stride_ptr) stride: T.let[T.int64] = strides[axis] return stride diff --git a/python/tvm/relax/transform/transform.py b/python/tvm/relax/transform/transform.py index d18f88b8c27f..dfa92478b19b 100644 --- a/python/tvm/relax/transform/transform.py +++ b/python/tvm/relax/transform/transform.py @@ -1134,9 +1134,9 @@ def main( @Ts.prim_func def add( - A: T.Buffer((2, 3), "float32"), - B: T.Buffer((2, 3), "float32"), - T_add: T.Buffer((2, 3), "float32"), + A: T.Tensor((2, 3), "float32"), + B: T.Tensor((2, 3), "float32"), + T_add: T.Tensor((2, 3), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(2, 3): @@ -1148,9 +1148,9 @@ def add( @Ts.prim_func def multiply( - A: T.Buffer((2, 3), "float32"), - B: T.Buffer((2, 3), "float32"), - T_multiply: T.Buffer((2, 3), "float32"), + A: T.Tensor((2, 3), "float32"), + B: T.Tensor((2, 3), "float32"), + T_multiply: T.Tensor((2, 3), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(2, 3): diff --git a/python/tvm/s_tir/schedule/schedule.py b/python/tvm/s_tir/schedule/schedule.py index dcea7c2fe879..e8917c1a7b68 100644 --- a/python/tvm/s_tir/schedule/schedule.py +++ b/python/tvm/s_tir/schedule/schedule.py @@ -624,7 +624,7 @@ def merge(self, *loops: list[LoopRV]) -> LoopRV: @Ts.prim_func def before_merge( - A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), B: T.Tensor((128, 128)), C: T.Tensor((128, 128)) ) -> None: for i, j in T.grid(128, 128): @@ -652,7 +652,7 @@ def before_merge( @Ts.prim_func def after_fuse( - A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), B: T.Tensor((128, 128)), C: T.Tensor((128, 128)) ) -> None: # the 2 loops are merged into 1 @@ -700,7 +700,7 @@ def fuse(self, *loops: list[LoopRV], preserve_unit_iters: bool = True) -> LoopRV .. code-block:: python @Ts.prim_func - def before_fuse(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_fuse(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -722,7 +722,7 @@ def before_fuse(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_fuse(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_fuse(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: # the 2 loops are fused into 1 @@ -790,7 +790,7 @@ def split( .. code-block:: python @Ts.prim_func - def before_split(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_split(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -812,7 +812,7 @@ def before_split(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_split(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_split(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: # the original loop is split into 2 loops @@ -877,7 +877,7 @@ def loop_partition( .. code-block:: python @Ts.prim_func - def before_partition(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_partition(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -898,7 +898,7 @@ def before_partition(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: .. code-block:: python - def after_partition(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_partition(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: # the original loop is partition into 3 loops @@ -971,7 +971,7 @@ def reorder(self, *ordered_loops: list[LoopRV]) -> None: .. code-block:: python @Ts.prim_func - def before_reorder(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_reorder(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -993,7 +993,7 @@ def before_reorder(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_reorder(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_reorder(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: # Here j and i are reordered @@ -1025,9 +1025,9 @@ def reorder_block_iter_var(self, block: SBlockRV, new_order: list[int]) -> None: @Ts.prim_func def matmul( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("C"): @@ -1050,9 +1050,9 @@ def matmul( @Ts.prim_func def matmul_after_reorder_block_iter_var( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ): for i, j, k in T.grid(128, 128, 128): with Ts.sblock("C"): @@ -1093,9 +1093,9 @@ def add_unit_loop(self, block_or_loop: LoopRV | SBlockRV) -> LoopRV: @Ts.prim_func def before_add_unit_loop( - A: T.Buffer((), "int32"), - B: T.Buffer((), "int32"), - C: T.Buffer((), "int32"), + A: T.Tensor((), "int32"), + B: T.Tensor((), "int32"), + C: T.Tensor((), "int32"), ) -> None: with Ts.sblock("C"): vi = Ts.axis.spatial(1, 0) @@ -1115,9 +1115,9 @@ def before_add_unit_loop( @Ts.prim_func def after_add_unit_loop( - A: T.Buffer((), "int32"), - B: T.Buffer((), "int32"), - C: T.Buffer((), "int32"), + A: T.Tensor((), "int32"), + B: T.Tensor((), "int32"), + C: T.Tensor((), "int32"), ) -> None: for u in T.serial(1): with Ts.sblock("C"): @@ -1152,7 +1152,7 @@ def parallel(self, loop: LoopRV) -> None: .. code-block:: python @Ts.prim_func - def before_parallel(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_parallel(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -1173,7 +1173,7 @@ def before_parallel(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_parallel(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_parallel(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i in T.parallel(0, 128): @@ -1208,7 +1208,7 @@ def vectorize(self, loop: LoopRV) -> None: .. code-block:: python @Ts.prim_func - def before_vectorize(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_vectorize(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -1229,7 +1229,7 @@ def before_vectorize(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_vectorize(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_vectorize(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i in T.serial(0, 128): @@ -1269,7 +1269,7 @@ def bind(self, loop: LoopRV, thread_axis: str) -> None: .. code-block:: python @Ts.prim_func - def before_bind(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_bind(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -1291,7 +1291,7 @@ def before_bind(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_bind(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_bind(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i in T.thread_binding(0, 128, thread = "blockIdx.x"): @@ -1320,7 +1320,7 @@ def unroll(self, loop: LoopRV) -> None: .. code-block:: python @Ts.prim_func - def before_unroll(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_unroll(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -1341,7 +1341,7 @@ def before_unroll(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_unroll(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_unroll(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i in T.unroll(0, 128): @@ -1399,7 +1399,7 @@ def cache_read( .. code-block:: python @Ts.prim_func - def before_cache_read(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_cache_read(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -1421,7 +1421,7 @@ def before_cache_read(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_cache_read(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_cache_read(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: A_local = Ts.sblock_alloc_buffer((128, 128), scope="local") @@ -1494,7 +1494,7 @@ def cache_write( .. code-block:: python @Ts.prim_func - def before_cache_write(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_cache_write(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -1516,7 +1516,7 @@ def before_cache_write(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None .. code-block:: python @Ts.prim_func - def after_cache_write(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_cache_write(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: B_local = Ts.sblock_alloc_buffer((128, 128), scope="local") @@ -1587,7 +1587,7 @@ def reindex_cache_read( .. code-block:: python @Ts.prim_func - def before_reindex_cache_read(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_reindex_cache_read(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -1609,7 +1609,7 @@ def before_reindex_cache_read(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) .. code-block:: python @Ts.prim_func - def after_reindex_cache_read(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_reindex_cache_read(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: A_local = Ts.sblock_alloc_buffer((128, 128), scope="local") @@ -1688,7 +1688,7 @@ def reindex_cache_write( @Ts.prim_func def before_reindex_cache_write( - A: T.Buffer((128, 128)), B: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), B: T.Tensor((128, 128)) ) -> None: for i, j in T.grid(128, 128): @@ -1710,7 +1710,7 @@ def before_reindex_cache_write( .. code-block:: python @Ts.prim_func - def after_cache_write(A: T.Buffer((128, 128)), B: T.Buffer((64, 2, 128))) -> None: + def after_cache_write(A: T.Tensor((128, 128)), B: T.Tensor((64, 2, 128))) -> None: B_local = Ts.sblock_alloc_buffer((128, 128), scope="local") @@ -1780,7 +1780,7 @@ def cache_inplace( .. code-block:: python @Ts.prim_func - def before_cache_inplace(data_io: T.Buffer((64), "int32")): + def before_cache_inplace(data_io: T.Tensor((64), "int32")): for i0 in T.serial(1): with Ts.sblock("A"): Ts.reads(data_io[:64]) @@ -1801,7 +1801,7 @@ def before_cache_inplace(data_io: T.Buffer((64), "int32")): .. code-block:: python @Ts.prim_func - def cache_inplace(data_io: T.Buffer(64, "int32")) -> None: + def cache_inplace(data_io: T.Tensor(64, "int32")) -> None: data_io_local = Ts.sblock_alloc_buffer([64], dtype="int32", scope="local") for i0 in T.serial(1): for ax0 in T.serial(64): @@ -1864,7 +1864,7 @@ def cache_index( .. code-block:: python @Ts.prim_func - def resize(A: T.Buffer((1, 3, 40, 40)), B: T.Buffer((1, 3, 80, 80))) -> None: + def resize(A: T.Tensor((1, 3, 40, 40)), B: T.Tensor((1, 3, 80, 80))) -> None: for i0, i1, i2, i3 in T.grid(1, 3, 80, 80): @@ -1887,7 +1887,7 @@ def resize(A: T.Buffer((1, 3, 40, 40)), B: T.Buffer((1, 3, 80, 80))) -> None: @Ts.prim_func def resize_cache_index( - A: T.Buffer((1, 3, 40, 40), "float32"), B: T.Buffer((1, 3, 80, 80), "float32") + A: T.Tensor((1, 3, 40, 40), "float32"), B: T.Tensor((1, 3, 80, 80), "float32") ) -> None: index_var_0 = Ts.sblock_alloc_buffer([80, 80], dtype="int32", strides=[1]) index_var_1 = Ts.sblock_alloc_buffer([80], dtype="int32", strides=[1]) @@ -1965,8 +1965,8 @@ def reindex(self, block: SBlockRV | str, buffer: tuple[str, int] | str | Buffer) @Ts.prim_func def before_reindex( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32") ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -1987,8 +1987,8 @@ def before_reindex( @Ts.prim_func def after_reindex( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32") ) -> None: A_reindex = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): @@ -2079,7 +2079,7 @@ def compute_at( @Ts.prim_func def before_compute_at( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -2109,7 +2109,7 @@ def before_compute_at( @Ts.prim_func def after_compute_at( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -2179,7 +2179,7 @@ def reverse_compute_at( @Ts.prim_func def before_reverse_compute_at( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -2209,7 +2209,7 @@ def before_reverse_compute_at( @Ts.prim_func def after_reverse_compute_at( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -2258,7 +2258,7 @@ def compute_inline(self, block: SBlockRV | str) -> None: .. code-block:: python @Ts.prim_func - def before_inline(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: + def before_inline(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) @@ -2284,7 +2284,7 @@ def before_inline(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_inline(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: + def after_inline(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -2327,7 +2327,7 @@ def reverse_compute_inline(self, block: SBlockRV | str) -> None: .. code-block:: python @Ts.prim_func - def before_inline(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: + def before_inline(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) @@ -2353,7 +2353,7 @@ def before_inline(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_inline(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: + def after_inline(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -2461,7 +2461,7 @@ def decompose_reduction(self, block: SBlockRV | str, loop: LoopRV) -> SBlockRV: @Ts.prim_func def before_decompose( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i, j, k in T.grid(128, 128, 128): @@ -2585,7 +2585,7 @@ def rfactor(self, loop: LoopRV, factor_axis: int) -> SBlockRV: .. code-block:: python @Ts.prim_func - def before_rfactor(A: T.Buffer((128, 128, 128)), B: T.Buffer((128,))) -> None: + def before_rfactor(A: T.Tensor((128, 128, 128)), B: T.Tensor((128,))) -> None: for ii, i, j in T.grid(128, 128, 128): with Ts.sblock("B"): vii, vi, vj = Ts.axis.remap("SRR", [ii, i, j]) @@ -2607,7 +2607,7 @@ def before_rfactor(A: T.Buffer((128, 128, 128)), B: T.Buffer((128,))) -> None: .. code-block:: python @Ts.prim_func - def after_rfactor(A: T.Buffer([128, 128, 128]), B: T.Buffer([128])) -> None: + def after_rfactor(A: T.Tensor([128, 128, 128]), B: T.Tensor([128])) -> None: B_rf = Ts.sblock_alloc_buffer([128, 128]) @@ -2683,7 +2683,7 @@ def storage_align( # pylint: disable=too-many-arguments .. code-block:: python @Ts.prim_func - def before_storage_align(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: + def before_storage_align(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) @@ -2709,7 +2709,7 @@ def before_storage_align(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> No .. code-block:: python @Ts.prim_func - def after_storage_align(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: + def after_storage_align(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) @@ -2727,7 +2727,7 @@ def after_storage_align(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> Non Note ---- - Storage_align requires the buffer to be an intermediate buffer defined via `alloc_buffer`. + Storage_align requires the buffer to be an intermediate buffer defined via `alloc_tensor`. """ block = self._normalize_block_arg(block) _ffi_api.ScheduleStorageAlign( # type: ignore # pylint: disable=no-member @@ -2759,7 +2759,7 @@ def set_scope( @Ts.prim_func def before_set_scope( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), dtype="float32") @@ -2786,7 +2786,7 @@ def before_set_scope( @Ts.prim_func def after_set_scope( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B_shared = Ts.sblock_alloc_buffer([128, 128], dtype="float32", scope="shared") @@ -2801,7 +2801,7 @@ def after_set_scope( Note ---- - `set_scope` requires the buffer to be an intermediate buffer defined via `alloc_buffer`. + `set_scope` requires the buffer to be an intermediate buffer defined via `alloc_tensor`. """ block = self._normalize_block_arg(block) if not isinstance(buffer_index, int): @@ -2838,7 +2838,7 @@ def unsafe_set_dtype(self, block: SBlockRV | str, buffer_index: int, dtype: str) @Ts.prim_func def before_set_dtype( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), dtype="float32") @@ -2865,7 +2865,7 @@ def before_set_dtype( @Ts.prim_func def after_set_dtype( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), dtype="float16") @@ -2881,7 +2881,7 @@ def after_set_dtype( Note ---- `unsafe_set_dtype` requires the buffer to be an intermediate buffer defined via - `alloc_buffer`. + `alloc_tensor`. """ block = self._normalize_block_arg(block) _ffi_api.ScheduleUnsafeSetDType( # type: ignore # pylint: disable=no-member @@ -2917,8 +2917,8 @@ def blockize( @Ts.prim_func def before_blockize( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32") ) -> None: for i_0, j_0, i_1, j_1 in T.grid(8, 8, 16, 16): with Ts.sblock("B"): @@ -2944,8 +2944,8 @@ def before_blockize( @Ts.prim_func def after_blockize( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32") )-> None: for i_0, j_0 in T.grid(8, 8): with Ts.sblock("B_o"): @@ -2996,9 +2996,9 @@ def tensorize( @Ts.prim_func def before_tensorize( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: # body # with Ts.sblock("root") @@ -3017,9 +3017,9 @@ def before_tensorize( @Ts.prim_func def mma_desc( - A: T.Buffer((16, 16), align=128, offset_factor=1), - B: T.Buffer((16, 16), align=128, offset_factor=1), - C: T.Buffer((16, 16), align=128, offset_factor=1), + A: T.Tensor((16, 16), align=128, offset_factor=1), + B: T.Tensor((16, 16), align=128, offset_factor=1), + C: T.Tensor((16, 16), align=128, offset_factor=1), ) -> None: with Ts.sblock("root"): @@ -3032,9 +3032,9 @@ def mma_desc( @Ts.prim_func def mma_intrin( - A: T.Buffer((16, 16), align=128, offset_factor=1), - B: T.Buffer((16, 16), align=128, offset_factor=1), - C: T.Buffer((16, 16), align=128, offset_factor=1), + A: T.Tensor((16, 16), align=128, offset_factor=1), + B: T.Tensor((16, 16), align=128, offset_factor=1), + C: T.Tensor((16, 16), align=128, offset_factor=1), ) -> None: with Ts.sblock("root"): @@ -3072,9 +3072,9 @@ def mma_intrin( @Ts.prim_func def after_tensorize( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: # body # with Ts.sblock("root") @@ -3155,7 +3155,7 @@ def annotate( .. code-block:: python @Ts.prim_func - def before_annotate(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_annotate(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -3176,7 +3176,7 @@ def before_annotate(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_annotate(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_annotate(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -3209,7 +3209,7 @@ def unannotate(self, block_or_loop: SBlockRV | LoopRV, ann_key: str) -> None: .. code-block:: python @Ts.prim_func - def before_unannotate(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_unannotate(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -3231,7 +3231,7 @@ def before_unannotate(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: .. code-block:: python @Ts.prim_func - def after_unannotate(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def after_unannotate(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -3405,7 +3405,7 @@ def transform_layout( @Ts.prim_func def before_transform_layout( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -3434,7 +3434,7 @@ def before_transform_layout( @Ts.prim_func def two_elementwise_transformed_intermediate_buffer( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((8, 8, 16, 16), "float32") @@ -3519,8 +3519,8 @@ def transform_block_layout(self, block: SBlockRV | str, index_map: IndexMap | Ca @Ts.prim_func def before_transform_block_layout( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32") ) -> None: for i, j in T.grid(16, 16): with Ts.sblock("B"): @@ -3541,8 +3541,8 @@ def before_transform_block_layout( @Ts.prim_func def after_transform_block_layout( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32") ) -> None: for i in range(256): with Ts.sblock("B"): @@ -3597,7 +3597,7 @@ def decompose_padding(self, block: SBlockRV | str, loop: LoopRV) -> SBlockRV: .. code-block:: python @Ts.prim_func - def before_decompose(x: T.Buffer(128, "int32"), y: T.Buffer(140, "int32")): + def before_decompose(x: T.Tensor(128, "int32"), y: T.Tensor(140, "int32")): for i in range(140): with Ts.sblock("block"): vi = Ts.axis.remap("S", [i]) @@ -3617,7 +3617,7 @@ def before_decompose(x: T.Buffer(128, "int32"), y: T.Buffer(140, "int32")): .. code-block:: python @Ts.prim_func - def after_decompose(x: T.Buffer(128, "int32"), y: T.Buffer(140, "int32")): + def after_decompose(x: T.Tensor(128, "int32"), y: T.Tensor(140, "int32")): for i in T.serial(140): with Ts.sblock("block_pad_const"): vi = Ts.axis.spatial(140, i) @@ -3667,9 +3667,9 @@ def pad_einsum(self, block: SBlockRV | str, padding: list[int]) -> None: @Ts.prim_func def before_pad_einsum( - A: T.Buffer((127, 127), "float32"), - B: T.Buffer((127, 127), "float32"), - C: T.Buffer((127, 127), "float32"), + A: T.Tensor((127, 127), "float32"), + B: T.Tensor((127, 127), "float32"), + C: T.Tensor((127, 127), "float32"), ) -> None: for i0, i1, i2 in T.grid(127, 127, 127): with Ts.sblock("C_shared"): @@ -3693,9 +3693,9 @@ def before_pad_einsum( @Ts.prim_func def main( - A: T.Buffer((127, 127), "float32"), - B: T.Buffer((127, 127), "float32"), - C: T.Buffer((127, 127), "float32"), + A: T.Tensor((127, 127), "float32"), + B: T.Tensor((127, 127), "float32"), + C: T.Tensor((127, 127), "float32"), ): # with Ts.sblock("root"): A_pad = Ts.sblock_alloc_buffer((128, 128)) @@ -3746,7 +3746,7 @@ def rolling_buffer(self, block: SBlockRV | str, write_buffer_index: int) -> None 1) The block is not an output block and has only RAW dependencies. - 2) The buffer to be an intermediate buffer defined via `alloc_buffer`. + 2) The buffer to be an intermediate buffer defined via `alloc_tensor`. 3) The LCA of the producer and consumer of the buffer is a for loop, typically, the producer and consumer of the buffer are cascaded through compute_at. @@ -3770,7 +3770,7 @@ def rolling_buffer(self, block: SBlockRV | str, write_buffer_index: int) -> None @Ts.prim_func def before_rolling_buffer( - A: T.Buffer((12, 12), "int8"), C: T.Buffer((8, 8), "int8") + A: T.Tensor((12, 12), "int8"), C: T.Tensor((8, 8), "int8") ) -> None: # body # with Ts.sblock("root") @@ -3807,8 +3807,8 @@ def before_rolling_buffer( @Ts.prim_func def after_rolling_buffer( - A: T.Buffer((12, 12), "int8"), - C: T.Buffer((8, 8), "int8") + A: T.Tensor((12, 12), "int8"), + C: T.Tensor((8, 8), "int8") ) -> None: # body # with Ts.sblock("root") @@ -3909,8 +3909,8 @@ def annotate_buffer_access( @Ts.prim_func def before_annotate_buffer_access( - A: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): @@ -3938,8 +3938,8 @@ def before_annotate_buffer_access( @Ts.prim_func def after_annotate_buffer_access( - A: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): diff --git a/python/tvm/s_tir/tensor_intrin/arm_cpu.py b/python/tvm/s_tir/tensor_intrin/arm_cpu.py index 4b75b49dfa71..399d2a2037c1 100644 --- a/python/tvm/s_tir/tensor_intrin/arm_cpu.py +++ b/python/tvm/s_tir/tensor_intrin/arm_cpu.py @@ -39,9 +39,9 @@ @Ts.prim_func def neon_4x4_i8i8i32_desc( - A: T.Buffer((4,), "int8", offset_factor=1), - B: T.Buffer((4, 4), "int8", offset_factor=1), - C: T.Buffer((4,), "int32", offset_factor=1), + A: T.Tensor((4,), "int8", offset_factor=1), + B: T.Tensor((4, 4), "int8", offset_factor=1), + C: T.Tensor((4,), "int32", offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:4], A[0:4], B[0:4, 0:4]) @@ -55,9 +55,9 @@ def neon_4x4_i8i8i32_desc( @Ts.prim_func def neon_4x4_i8i8i32_impl( - A: T.Buffer((4,), "int8", offset_factor=1), - B: T.Buffer((4, 4), "int8", offset_factor=1), - C: T.Buffer((4,), "int32", offset_factor=1), + A: T.Tensor((4,), "int8", offset_factor=1), + B: T.Tensor((4, 4), "int8", offset_factor=1), + C: T.Tensor((4,), "int32", offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:4], A[0:4], B[0:4, 0:4]) @@ -121,9 +121,9 @@ def get_dotprod_intrin(in_dtype, out_dtype): @Ts.prim_func def dot_prod_desc( - A: T.Buffer((4,), dtype=in_dtype, offset_factor=1), - B: T.Buffer((4, 4), dtype=in_dtype, offset_factor=1), - C: T.Buffer((4,), dtype=out_dtype, offset_factor=1), + A: T.Tensor((4,), dtype=in_dtype, offset_factor=1), + B: T.Tensor((4, 4), dtype=in_dtype, offset_factor=1), + C: T.Tensor((4,), dtype=out_dtype, offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:4], A[0:4], B[0:4, 0:4]) @@ -138,9 +138,9 @@ def dot_prod_desc( @Ts.prim_func def dot_prod_impl( - A: T.Buffer((4,), dtype=in_dtype, offset_factor=1), - B: T.Buffer((4, 4), dtype=in_dtype, offset_factor=1), - C: T.Buffer((4,), dtype=out_dtype, offset_factor=1), + A: T.Tensor((4,), dtype=in_dtype, offset_factor=1), + B: T.Tensor((4, 4), dtype=in_dtype, offset_factor=1), + C: T.Tensor((4,), dtype=out_dtype, offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:4], A[0:4], B[0:4, 0:4]) @@ -261,8 +261,8 @@ def get_sme_transpose_interleave_2svlx2svl_fp32_intrin(cols, rows): @Ts.prim_func def desc( - A: T.Buffer((SVF2, SVF2), dtype="float32", offset_factor=1), - A_t: T.Buffer((SVF2, SVF2), dtype="float32", offset_factor=1), + A: T.Tensor((SVF2, SVF2), dtype="float32", offset_factor=1), + A_t: T.Tensor((SVF2, SVF2), dtype="float32", offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(A[0:SVF2, 0:SVF2]) @@ -278,11 +278,11 @@ def impl(): with IRBuilder() as ib: with build_prim_func(): A = T.arg_( - "a", T.Buffer((SVF2, SVF2), "float32", offset_factor=1, strides=[T.int32(), 1]) + "a", T.Tensor((SVF2, SVF2), "float32", offset_factor=1, strides=[T.int32(), 1]) ) A_t = T.arg_( "a_t", - T.Buffer((SVF2, SVF2), "float32", offset_factor=1, strides=[T.int32(), 1]), + T.Tensor((SVF2, SVF2), "float32", offset_factor=1, strides=[T.int32(), 1]), ) with Ts.sblock("root"): @@ -391,8 +391,8 @@ def get_sme_transpose_interleave_block2_2svl_fp16_intrin(): @Ts.prim_func def desc( - A: T.Buffer((SVF2, SVF), dtype="float16", offset_factor=1), - A_t: T.Buffer((SVF, SVF2), dtype="float16", offset_factor=1), + A: T.Tensor((SVF2, SVF), dtype="float16", offset_factor=1), + A_t: T.Tensor((SVF, SVF2), dtype="float16", offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(A[0:SVF2, 0:SVF]) @@ -406,10 +406,10 @@ def impl(): with IRBuilder() as ib: with build_prim_func(): A = T.arg_( - "a", T.Buffer((SVF2, SVF), "float16", offset_factor=1, strides=[T.int32(), 1]) + "a", T.Tensor((SVF2, SVF), "float16", offset_factor=1, strides=[T.int32(), 1]) ) A_t = T.arg_( - "a_t", T.Buffer((SVF, SVF2), "float16", offset_factor=1, strides=[T.int32(), 1]) + "a_t", T.Tensor((SVF, SVF2), "float16", offset_factor=1, strides=[T.int32(), 1]) ) ptrue_fp16 = _create_ptrue_mask("float16") @@ -593,9 +593,9 @@ def get_sme_gemm_interleaved_mopa_2svlx2svl_intrin(M, K, in_dtype): @Ts.prim_func def desc( - A: T.Buffer((K, SVF2), dtype=in_dtype, offset_factor=1), - B: T.Buffer((K, SVF2), dtype=in_dtype, offset_factor=1), - C: T.Buffer((SVF2, SVF2), dtype="float32", offset_factor=1), + A: T.Tensor((K, SVF2), dtype=in_dtype, offset_factor=1), + B: T.Tensor((K, SVF2), dtype=in_dtype, offset_factor=1), + C: T.Tensor((SVF2, SVF2), dtype="float32", offset_factor=1), ): with Ts.sblock("root"): Ts.reads(C[0:SVF2, 0:SVF2], A[0:K, 0:SVF2], B[0:K, 0:SVF2]) @@ -611,13 +611,13 @@ def impl(): with IRBuilder() as ib: with build_prim_func(): A = T.arg_( - "a", T.Buffer((K, SVF2), in_dtype, offset_factor=1, strides=[T.int32(), 1]) + "a", T.Tensor((K, SVF2), in_dtype, offset_factor=1, strides=[T.int32(), 1]) ) B = T.arg_( - "b", T.Buffer((K, SVF2), in_dtype, offset_factor=1, strides=[T.int32(), 1]) + "b", T.Tensor((K, SVF2), in_dtype, offset_factor=1, strides=[T.int32(), 1]) ) C = T.arg_( - "c", T.Buffer((SVF2, SVF2), "float32", offset_factor=1, strides=[T.int32(), 1]) + "c", T.Tensor((SVF2, SVF2), "float32", offset_factor=1, strides=[T.int32(), 1]) ) ptrue = _create_ptrue_mask(in_dtype) @@ -722,7 +722,7 @@ def get_sme_init_intrin(): SVF2 = 2 * 4 * T.vscale() @Ts.prim_func - def desc(C: T.Buffer((SVF2, SVF2), "float32", offset_factor=1)) -> None: + def desc(C: T.Tensor((SVF2, SVF2), "float32", offset_factor=1)) -> None: with Ts.sblock("root"): Ts.reads() Ts.writes(C[0:SVF2, 0:SVF2]) @@ -732,7 +732,7 @@ def desc(C: T.Buffer((SVF2, SVF2), "float32", offset_factor=1)) -> None: C[v_m, v_n] = T.float32(0) @Ts.prim_func - def impl(C: T.Buffer((SVF2, SVF2), "float32", offset_factor=1)) -> None: + def impl(C: T.Tensor((SVF2, SVF2), "float32", offset_factor=1)) -> None: with Ts.sblock("root"): Ts.reads() Ts.writes(C[0:SVF2, 0:SVF2]) diff --git a/python/tvm/s_tir/tensor_intrin/cuda.py b/python/tvm/s_tir/tensor_intrin/cuda.py index f94edb47c399..dc536592f97b 100644 --- a/python/tvm/s_tir/tensor_intrin/cuda.py +++ b/python/tvm/s_tir/tensor_intrin/cuda.py @@ -152,10 +152,10 @@ def get_ldmatrix_intrin( @Ts.prim_func def ldmatrix_desc( - warp: T.Buffer( + warp: T.Tensor( (WARP_SIZE, local_size), dtype, align=64, offset_factor=offset_factor, scope="warp" ), - shared: T.Buffer( + shared: T.Tensor( (smem_tile_row, smem_tile_col), dtype, align=64, @@ -181,10 +181,10 @@ def ldmatrix_desc( @Ts.prim_func def ldmatrix_impl( - warp: T.Buffer( + warp: T.Tensor( (WARP_SIZE, local_size), dtype, align=64, offset_factor=offset_factor, scope="warp" ), - shared: T.Buffer( + shared: T.Tensor( (smem_tile_row, smem_tile_col), dtype, align=64, @@ -330,13 +330,13 @@ def swap_if_flag(i, j, flag): @Ts.prim_func def mma_sync_desc( - A: T.Buffer( + A: T.Tensor( (WARP_SIZE, local_size), a_dtype, align=64, offset_factor=A_offset_factor, scope="warp" ), - B: T.Buffer( + B: T.Tensor( (WARP_SIZE, local_size), b_dtype, align=64, offset_factor=B_offset_factor, scope="warp" ), - C: T.Buffer( + C: T.Tensor( (WARP_SIZE, local_size_out), out_dtype, align=64, @@ -375,13 +375,13 @@ def mma_sync_desc( @Ts.prim_func def mma_sync_impl( - A: T.Buffer( + A: T.Tensor( (WARP_SIZE, local_size), a_dtype, align=64, offset_factor=A_offset_factor, scope="warp" ), - B: T.Buffer( + B: T.Tensor( (WARP_SIZE, local_size), b_dtype, align=64, offset_factor=B_offset_factor, scope="warp" ), - C: T.Buffer( + C: T.Tensor( (WARP_SIZE, local_size_out), out_dtype, align=64, @@ -523,7 +523,7 @@ def get_mma_fill_intrin(dtype, local_size): index_map = shared_16x16_to_ldmatrix_32x8_layout @Ts.prim_func - def mma_fill_desc(C_warp: T.Buffer([WARP_SIZE, local_size], dtype=dtype, scope="warp")) -> None: + def mma_fill_desc(C_warp: T.Tensor([WARP_SIZE, local_size], dtype=dtype, scope="warp")) -> None: with Ts.sblock("root"): Ts.reads() Ts.writes(C_warp[0:WARP_SIZE, 0:local_size]) @@ -537,7 +537,7 @@ def mma_fill_desc(C_warp: T.Buffer([WARP_SIZE, local_size], dtype=dtype, scope=" @Ts.prim_func def mma_fill_impl( - C_warp: T.Buffer([WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1), + C_warp: T.Tensor([WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads() @@ -568,8 +568,8 @@ def get_mma_store_intrin(dtype, local_size, scope="global", use_mma_store_intrin @Ts.prim_func def mma_store_desc( - C_warp: T.Buffer([WARP_SIZE, local_size], dtype=dtype, scope="warp"), - C: T.Buffer([M_DIM, N_DIM], dtype=dtype, scope=scope), + C_warp: T.Tensor([WARP_SIZE, local_size], dtype=dtype, scope="warp"), + C: T.Tensor([M_DIM, N_DIM], dtype=dtype, scope=scope), ) -> None: with Ts.sblock("root"): Ts.reads(C_warp[0:WARP_SIZE, 0:local_size]) @@ -588,8 +588,8 @@ def mma_store_desc( @Ts.prim_func def mma_store_impl( - C_warp: T.Buffer([WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1), - C: T.Buffer( + C_warp: T.Tensor([WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1), + C: T.Tensor( [M_DIM, N_DIM], dtype=dtype, scope=scope, offset_factor=1, strides=[s0, s1] ), ) -> None: @@ -616,8 +616,8 @@ def mma_store_impl( @Ts.prim_func def mma_store_impl( - C_warp: T.Buffer([WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1), - C: T.Buffer( + C_warp: T.Tensor([WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1), + C: T.Tensor( [M_DIM, N_DIM], dtype=dtype, scope=scope, offset_factor=1, strides=[s0, s1] ), ) -> None: @@ -795,10 +795,10 @@ def get_wmma_load_intrin( @Ts.prim_func def wmma_load_desc( - A: T.Buffer( + A: T.Tensor( (frag_m, frag_n), dtype, align=64, offset_factor=offset_factor, scope=shared_scope ), - C: T.Buffer( + C: T.Tensor( (frag_m, frag_n), dtype, align=64, @@ -821,7 +821,7 @@ def wmma_load_desc( @Ts.prim_func def wmma_load_impl( - A: T.Buffer( + A: T.Tensor( (frag_m, frag_n), dtype, align=64, @@ -829,7 +829,7 @@ def wmma_load_impl( scope=shared_scope, strides=[s1, s0], ), - C: T.Buffer( + C: T.Tensor( (frag_m, frag_n), dtype, align=64, @@ -866,7 +866,7 @@ def get_wmma_fill_intrin( @Ts.prim_func def wmma_fill_desc( - C: T.Buffer( + C: T.Tensor( (m_dim, n_dim), dtype, align=64, offset_factor=offset_factor, scope="wmma.accumulator" ), ) -> None: @@ -883,7 +883,7 @@ def wmma_fill_desc( @Ts.prim_func def wmma_fill_impl( - C: T.Buffer( + C: T.Tensor( (m_dim, n_dim), dtype, align=64, @@ -917,10 +917,10 @@ def get_wmma_store_intrin( @Ts.prim_func def wmma_store_desc( - A: T.Buffer( + A: T.Tensor( (m_dim, n_dim), dtype, align=64, offset_factor=offset_factor, scope="wmma.accumulator" ), - C: T.Buffer((m_dim, n_dim), dtype, align=64, offset_factor=offset_factor, scope=scope), + C: T.Tensor((m_dim, n_dim), dtype, align=64, offset_factor=offset_factor, scope=scope), ) -> None: with Ts.sblock("root"): Ts.reads(A[0:m_dim, 0:n_dim]) @@ -937,7 +937,7 @@ def wmma_store_desc( @Ts.prim_func def wmma_store_impl( - A: T.Buffer( + A: T.Tensor( (m_dim, n_dim), dtype, align=64, @@ -945,7 +945,7 @@ def wmma_store_impl( scope="wmma.accumulator", strides=[d1, d0], ), - C: T.Buffer( + C: T.Tensor( (m_dim, n_dim), dtype, align=64, @@ -996,17 +996,17 @@ def maybe_swap(i, j): @Ts.prim_func def wmma_sync_desc( - A: T.Buffer( + A: T.Tensor( (m_dim, k_dim), in_dtype, align=64, offset_factor=A_offset_factor, scope="wmma.matrix_a" ), - B: T.Buffer( + B: T.Tensor( maybe_swap(k_dim, n_dim), in_dtype, align=64, offset_factor=B_offset_factor, scope="wmma.matrix_b", ), - C: T.Buffer( + C: T.Tensor( (m_dim, n_dim), out_dtype, align=64, @@ -1034,7 +1034,7 @@ def wmma_sync_desc( @Ts.prim_func def wmma_sync_impl( - A: T.Buffer( + A: T.Tensor( (m_dim, k_dim), in_dtype, align=64, @@ -1042,7 +1042,7 @@ def wmma_sync_impl( scope="wmma.matrix_a", strides=[a1, a0], ), - B: T.Buffer( + B: T.Tensor( maybe_swap(k_dim, n_dim), in_dtype, align=64, @@ -1050,7 +1050,7 @@ def wmma_sync_impl( scope="wmma.matrix_b", strides=[b1, b0], ), - C: T.Buffer( + C: T.Tensor( (m_dim, n_dim), out_dtype, align=64, @@ -1421,7 +1421,7 @@ def get_mma_init_intrin( @Ts.prim_func def mma_init_desc( - dst: T.Buffer((m_dim, n_dim), dtype, align=64, offset_factor=1, scope="m16n8k8.matrixC"), + dst: T.Tensor((m_dim, n_dim), dtype, align=64, offset_factor=1, scope="m16n8k8.matrixC"), ) -> None: with Ts.sblock("root"): Ts.reads() @@ -1433,7 +1433,7 @@ def mma_init_desc( @Ts.prim_func def mma_init_impl( - dst: T.Buffer((m_dim, n_dim), dtype, align=64, offset_factor=1, scope="m16n8k8.matrixC"), + dst: T.Tensor((m_dim, n_dim), dtype, align=64, offset_factor=1, scope="m16n8k8.matrixC"), ) -> None: with Ts.sblock("root"): Ts.reads() @@ -1469,8 +1469,8 @@ def get_mma_load_intrin( @Ts.prim_func def mma_load_desc( - src: T.Buffer((frag_m, frag_n), dtype, align=64, offset_factor=1, scope=shared_scope), - dst: T.Buffer((frag_m, frag_n), dtype, align=64, offset_factor=1, scope=mma_fragment_scope), + src: T.Tensor((frag_m, frag_n), dtype, align=64, offset_factor=1, scope=shared_scope), + dst: T.Tensor((frag_m, frag_n), dtype, align=64, offset_factor=1, scope=mma_fragment_scope), ) -> None: with Ts.sblock("root"): Ts.reads(src[0:frag_m, 0:frag_n]) @@ -1487,10 +1487,10 @@ def mma_load_desc( @Ts.prim_func def mma_load_impl( - src: T.Buffer( + src: T.Tensor( (frag_m, frag_n), dtype, align=64, offset_factor=1, scope=shared_scope, strides=[s0, s1] ), - dst: T.Buffer( + dst: T.Tensor( (frag_m, frag_n), dtype, align=64, @@ -1539,11 +1539,11 @@ def maybe_swap(i, j): @Ts.prim_func def mma_sync_desc( - A: T.Buffer((m_dim, k_dim), in_dtype, align=64, offset_factor=1, scope="m16n8k8.matrixA"), - B: T.Buffer( + A: T.Tensor((m_dim, k_dim), in_dtype, align=64, offset_factor=1, scope="m16n8k8.matrixA"), + B: T.Tensor( (B_shape_0, B_shape_1), in_dtype, align=64, offset_factor=1, scope="m16n8k8.matrixB" ), - C: T.Buffer((m_dim, n_dim), out_dtype, align=64, offset_factor=1, scope="m16n8k8.matrixC"), + C: T.Tensor((m_dim, n_dim), out_dtype, align=64, offset_factor=1, scope="m16n8k8.matrixC"), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:m_dim, 0:n_dim], A[0:m_dim, 0:k_dim], B[0:B_shape_0, 0:B_shape_1]) @@ -1565,7 +1565,7 @@ def mma_sync_desc( @Ts.prim_func def mma_sync_impl( - A: T.Buffer( + A: T.Tensor( (m_dim, k_dim), in_dtype, align=64, @@ -1573,7 +1573,7 @@ def mma_sync_impl( scope="m16n8k8.matrixA", strides=[a0, a1], ), - B: T.Buffer( + B: T.Tensor( (B_shape_0, B_shape_1), in_dtype, align=64, @@ -1581,7 +1581,7 @@ def mma_sync_impl( scope="m16n8k8.matrixB", strides=[b0, b1], ), - C: T.Buffer( + C: T.Tensor( (m_dim, n_dim), out_dtype, align=64, @@ -1623,8 +1623,8 @@ def get_mma_store_dummy_intrin( @Ts.prim_func def mma_store_desc( - src: T.Buffer((m_dim, n_dim), dtype, align=64, offset_factor=1, scope="m16n8k8.matrixC"), - dst: T.Buffer((m_dim, n_dim), dtype, align=64, offset_factor=1, scope="shared.dyn"), + src: T.Tensor((m_dim, n_dim), dtype, align=64, offset_factor=1, scope="m16n8k8.matrixC"), + dst: T.Tensor((m_dim, n_dim), dtype, align=64, offset_factor=1, scope="shared.dyn"), ) -> None: with Ts.sblock("root"): Ts.reads(src[0:m_dim, 0:n_dim]) diff --git a/python/tvm/s_tir/tensor_intrin/dot_product_common.py b/python/tvm/s_tir/tensor_intrin/dot_product_common.py index 14513f8b4c6a..5662407e2d1e 100644 --- a/python/tvm/s_tir/tensor_intrin/dot_product_common.py +++ b/python/tvm/s_tir/tensor_intrin/dot_product_common.py @@ -31,9 +31,9 @@ def get_dp4a_intrin(dtype_a, dtype_b, dtype_c): @Ts.prim_func def dp4a_desc( - A: T.Buffer((4,), dtype_a, offset_factor=1, align=4, scope="shared"), - B: T.Buffer((4,), dtype_b, offset_factor=1, align=4, scope="shared"), - C: T.Buffer((1,), dtype_c, offset_factor=1, align=4, scope="local"), + A: T.Tensor((4,), dtype_a, offset_factor=1, align=4, scope="shared"), + B: T.Tensor((4,), dtype_b, offset_factor=1, align=4, scope="shared"), + C: T.Tensor((1,), dtype_c, offset_factor=1, align=4, scope="local"), ) -> None: with Ts.sblock("root"): Ts.reads(C[0], A[0:4], B[0:4]) @@ -45,9 +45,9 @@ def dp4a_desc( @Ts.prim_func def dp4a_impl( - A: T.Buffer((4,), dtype_a, offset_factor=1, align=4, scope="shared"), - B: T.Buffer((4,), dtype_b, offset_factor=1, align=4, scope="shared"), - C: T.Buffer((1,), dtype_c, offset_factor=1, align=4, scope="local"), + A: T.Tensor((4,), dtype_a, offset_factor=1, align=4, scope="shared"), + B: T.Tensor((4,), dtype_b, offset_factor=1, align=4, scope="shared"), + C: T.Tensor((1,), dtype_c, offset_factor=1, align=4, scope="local"), ) -> None: with Ts.sblock("root"): Ts.reads(C[0], A[0:4], B[0:4]) diff --git a/python/tvm/s_tir/tensor_intrin/hexagon.py b/python/tvm/s_tir/tensor_intrin/hexagon.py index a1c477dfe27e..45b453a5f832 100644 --- a/python/tvm/s_tir/tensor_intrin/hexagon.py +++ b/python/tvm/s_tir/tensor_intrin/hexagon.py @@ -31,8 +31,8 @@ def generate_dma_load_intrin( @Ts.prim_func def sync_dma_load_desc( - A: T.Buffer(size, dtype, offset_factor=1, scope="global"), - C: T.Buffer(size, dtype, offset_factor=1, scope="global.vtcm"), + A: T.Tensor(size, dtype, offset_factor=1, scope="global"), + C: T.Tensor(size, dtype, offset_factor=1, scope="global.vtcm"), ) -> None: with Ts.sblock("root"): Ts.reads(A[0:size]) @@ -44,8 +44,8 @@ def sync_dma_load_desc( @Ts.prim_func def sync_dma_load_impl( - A: T.Buffer(size, dtype, offset_factor=1, scope="global"), - C: T.Buffer(size, dtype, offset_factor=1, scope="global.vtcm"), + A: T.Tensor(size, dtype, offset_factor=1, scope="global"), + C: T.Tensor(size, dtype, offset_factor=1, scope="global.vtcm"), ) -> None: with Ts.sblock("root"): Ts.reads(A[0:size]) @@ -80,9 +80,9 @@ def sync_dma_load_impl( def generate_dot_product_32x4_u8u8i32(mem_scope="global"): @Ts.prim_func def dot_product_32x4_u8u8i32_desc( - A: T.Buffer((4,), "uint8", offset_factor=1, scope=mem_scope), - B: T.Buffer((32, 4), "uint8", offset_factor=1, scope=mem_scope), - C: T.Buffer((32,), "int32", offset_factor=1, scope=mem_scope), + A: T.Tensor((4,), "uint8", offset_factor=1, scope=mem_scope), + B: T.Tensor((32, 4), "uint8", offset_factor=1, scope=mem_scope), + C: T.Tensor((32,), "int32", offset_factor=1, scope=mem_scope), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:32], A[0:4], B[0:32, 0:4]) @@ -95,9 +95,9 @@ def dot_product_32x4_u8u8i32_desc( @Ts.prim_func def dot_product_32x4_u8u8i32_vrmpy( - A: T.Buffer((4,), "uint8", offset_factor=1, scope=mem_scope), - B: T.Buffer((32, 4), "uint8", offset_factor=1, scope=mem_scope), - C: T.Buffer((32,), "int32", offset_factor=1, scope=mem_scope), + A: T.Tensor((4,), "uint8", offset_factor=1, scope=mem_scope), + B: T.Tensor((32, 4), "uint8", offset_factor=1, scope=mem_scope), + C: T.Tensor((32,), "int32", offset_factor=1, scope=mem_scope), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:32], A[0:4], B[0:32, 0:4]) @@ -123,9 +123,9 @@ def dot_product_32x4_u8u8i32_vrmpy( def generate_dot_product_32x4_u8i8i32(mem_scope="global"): @Ts.prim_func def dot_product_32x4_u8i8i32_desc( - A: T.Buffer((4,), "uint8", offset_factor=1, scope=mem_scope), - B: T.Buffer((32, 4), "int8", offset_factor=1, scope=mem_scope), - C: T.Buffer((32,), "int32", offset_factor=1, scope=mem_scope), + A: T.Tensor((4,), "uint8", offset_factor=1, scope=mem_scope), + B: T.Tensor((32, 4), "int8", offset_factor=1, scope=mem_scope), + C: T.Tensor((32,), "int32", offset_factor=1, scope=mem_scope), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:32], A[0:4], B[0:32, 0:4]) @@ -138,9 +138,9 @@ def dot_product_32x4_u8i8i32_desc( @Ts.prim_func def dot_product_32x4_u8i8i32_vrmpy( - A: T.Buffer((4,), "uint8", offset_factor=1, scope=mem_scope), - B: T.Buffer((32, 4), "int8", offset_factor=1, scope=mem_scope), - C: T.Buffer((32,), "int32", offset_factor=1, scope=mem_scope), + A: T.Tensor((4,), "uint8", offset_factor=1, scope=mem_scope), + B: T.Tensor((32, 4), "int8", offset_factor=1, scope=mem_scope), + C: T.Tensor((32,), "int32", offset_factor=1, scope=mem_scope), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:32], A[0:4], B[0:32, 0:4]) @@ -166,9 +166,9 @@ def dot_product_32x4_u8i8i32_vrmpy( def generate_dot_product_32x2_i16i16i32(mem_scope="global"): @Ts.prim_func def dot_product_32x2_i16i16i32_desc( - A: T.Buffer((2,), "int16", offset_factor=1, scope=mem_scope), - B: T.Buffer((32, 2), "int16", offset_factor=1, scope=mem_scope), - C: T.Buffer((32,), "int32", offset_factor=1, scope=mem_scope), + A: T.Tensor((2,), "int16", offset_factor=1, scope=mem_scope), + B: T.Tensor((32, 2), "int16", offset_factor=1, scope=mem_scope), + C: T.Tensor((32,), "int32", offset_factor=1, scope=mem_scope), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:32], A[0:2], B[0:32, 0:2]) @@ -181,9 +181,9 @@ def dot_product_32x2_i16i16i32_desc( @Ts.prim_func def dot_product_32x2_i16i16i32_vdmpy( - A: T.Buffer((2,), "int16", offset_factor=1, scope=mem_scope), - B: T.Buffer((32, 2), "int16", offset_factor=1, scope=mem_scope), - C: T.Buffer((32,), "int32", offset_factor=1, scope=mem_scope), + A: T.Tensor((2,), "int16", offset_factor=1, scope=mem_scope), + B: T.Tensor((32, 2), "int16", offset_factor=1, scope=mem_scope), + C: T.Tensor((32,), "int32", offset_factor=1, scope=mem_scope), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:32], A[0:2], B[0:32, 0:2]) diff --git a/python/tvm/s_tir/tensor_intrin/metal.py b/python/tvm/s_tir/tensor_intrin/metal.py index f79c8ae6ab96..09c1beab9630 100644 --- a/python/tvm/s_tir/tensor_intrin/metal.py +++ b/python/tvm/s_tir/tensor_intrin/metal.py @@ -43,7 +43,7 @@ def get_make_filled_simdgroup_matrix_intrin( dtype: str, col: int = 8, row: int = 8 ) -> tuple[PrimFunc, PrimFunc]: @Ts.prim_func - def desc(A: T.Buffer((col, row), dtype, scope="metal.simdgroup", offset_factor=1)) -> None: + def desc(A: T.Tensor((col, row), dtype, scope="metal.simdgroup", offset_factor=1)) -> None: with Ts.sblock("root"): Ts.reads() Ts.writes(A[0:col, 0:row]) @@ -57,7 +57,7 @@ def desc(A: T.Buffer((col, row), dtype, scope="metal.simdgroup", offset_factor=1 @Ts.prim_func def impl( - A: T.Buffer((col, row), dtype, scope="metal.simdgroup", strides=[d1, d0], offset_factor=1), + A: T.Tensor((col, row), dtype, scope="metal.simdgroup", strides=[d1, d0], offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads() @@ -84,8 +84,8 @@ def get_simdgroup_load_intrin( @Ts.prim_func def desc( - A: T.Buffer((col, row), dtype, align=align, scope=scope, offset_factor=1), - C: T.Buffer((col, row), dtype, align=align, scope="metal.simdgroup", offset_factor=1), + A: T.Tensor((col, row), dtype, align=align, scope=scope, offset_factor=1), + C: T.Tensor((col, row), dtype, align=align, scope="metal.simdgroup", offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(A[0:col, 0:row]) @@ -106,8 +106,8 @@ def desc( @Ts.prim_func def impl( - A: T.Buffer((col, row), dtype, align=align, scope=scope, strides=[s1, s0], offset_factor=1), - C: T.Buffer( + A: T.Tensor((col, row), dtype, align=align, scope=scope, strides=[s1, s0], offset_factor=1), + C: T.Tensor( (col, row), dtype, align=align, @@ -143,8 +143,8 @@ def get_simdgroup_store_intrin( @Ts.prim_func def desc( - A: T.Buffer((col, row), dtype, align=align, scope="metal.simdgroup", offset_factor=1), - C: T.Buffer((col, row), dtype, align=align, scope=scope, offset_factor=1), + A: T.Tensor((col, row), dtype, align=align, scope="metal.simdgroup", offset_factor=1), + C: T.Tensor((col, row), dtype, align=align, scope=scope, offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(A[0:col, 0:row]) @@ -164,7 +164,7 @@ def desc( @Ts.prim_func def impl( - A: T.Buffer( + A: T.Tensor( (col, row), dtype, align=align, @@ -172,7 +172,7 @@ def impl( strides=[s1, s0], offset_factor=1, ), - C: T.Buffer((col, row), dtype, align=align, scope=scope, strides=[d1, d0], offset_factor=1), + C: T.Tensor((col, row), dtype, align=align, scope=scope, strides=[d1, d0], offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(A[0:col, 0:row]) @@ -195,9 +195,9 @@ def get_simdgroup_multiply_accumulate_intrin( ) -> tuple[PrimFunc, PrimFunc]: @Ts.prim_func def desc( - A: T.Buffer((m_dim, k_dim), dtype, scope="metal.simdgroup", offset_factor=1), - B: T.Buffer((k_dim, n_dim), dtype, scope="metal.simdgroup", offset_factor=1), - C: T.Buffer((m_dim, n_dim), dtype, scope="metal.simdgroup", offset_factor=1), + A: T.Tensor((m_dim, k_dim), dtype, scope="metal.simdgroup", offset_factor=1), + B: T.Tensor((k_dim, n_dim), dtype, scope="metal.simdgroup", offset_factor=1), + C: T.Tensor((m_dim, n_dim), dtype, scope="metal.simdgroup", offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:m_dim, 0:n_dim], A[0:m_dim, 0:k_dim], B[0:k_dim, 0:n_dim]) @@ -216,13 +216,13 @@ def desc( @Ts.prim_func def impl( - A: T.Buffer( + A: T.Tensor( (m_dim, k_dim), dtype, scope="metal.simdgroup", strides=[a1, a0], offset_factor=1 ), - B: T.Buffer( + B: T.Tensor( (k_dim, n_dim), dtype, scope="metal.simdgroup", strides=[b1, b0], offset_factor=1 ), - C: T.Buffer( + C: T.Tensor( (m_dim, n_dim), dtype, scope="metal.simdgroup", strides=[c1, c0], offset_factor=1 ), ) -> None: diff --git a/python/tvm/s_tir/tensor_intrin/riscv_cpu.py b/python/tvm/s_tir/tensor_intrin/riscv_cpu.py index 78e4e00f7eab..7a8c1c2ae094 100644 --- a/python/tvm/s_tir/tensor_intrin/riscv_cpu.py +++ b/python/tvm/s_tir/tensor_intrin/riscv_cpu.py @@ -76,9 +76,9 @@ def rvv_vec_dot_product_kernels( @Ts.prim_func def rvv_vec_dot_prod_desc( - A: T.Buffer((n_elems,), data_dtype, offset_factor=1), - B: T.Buffer((n_lanes, n_elems), weight_dtype, offset_factor=1), - C: T.Buffer((n_lanes,), out_dtype, offset_factor=1), + A: T.Tensor((n_elems,), data_dtype, offset_factor=1), + B: T.Tensor((n_lanes, n_elems), weight_dtype, offset_factor=1), + C: T.Tensor((n_lanes,), out_dtype, offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:n_lanes], A[0:n_elems], B[0:n_lanes, 0:n_elems]) @@ -108,9 +108,9 @@ def rvv_vec_dot_prod_desc( # fmt: off @Ts.prim_func def rvv_vec_dot_prod_impl( - A: T.Buffer((n_elems,), data_dtype, offset_factor=1), - B: T.Buffer((n_lanes, n_elems), weight_dtype, offset_factor=1), - C: T.Buffer((n_lanes,), out_dtype, offset_factor=1), + A: T.Tensor((n_elems,), data_dtype, offset_factor=1), + B: T.Tensor((n_lanes, n_elems), weight_dtype, offset_factor=1), + C: T.Tensor((n_lanes,), out_dtype, offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:n_lanes], A[0:n_elems], B[0:n_lanes, 0:n_elems]) diff --git a/python/tvm/s_tir/tensor_intrin/rocm.py b/python/tvm/s_tir/tensor_intrin/rocm.py index 5d43cb1fe1e1..d0c1b252d6f4 100644 --- a/python/tvm/s_tir/tensor_intrin/rocm.py +++ b/python/tvm/s_tir/tensor_intrin/rocm.py @@ -30,9 +30,9 @@ @Ts.prim_func def sdot4( - A: T.Buffer((4,), "int8", offset_factor=1, align=4, scope="shared"), - B: T.Buffer((4,), "int8", offset_factor=1, align=4, scope="shared"), - C: T.Buffer((1,), "int32", offset_factor=1, align=4, scope="local"), + A: T.Tensor((4,), "int8", offset_factor=1, align=4, scope="shared"), + B: T.Tensor((4,), "int8", offset_factor=1, align=4, scope="shared"), + C: T.Tensor((1,), "int32", offset_factor=1, align=4, scope="local"), ) -> None: with Ts.sblock("root"): Ts.reads(C[0], A[0:4], B[0:4]) @@ -123,7 +123,7 @@ def get_mma_fill_intrin(dtype, local_size): index_map = shared_16x16_to_local_64x4_layout_C @Ts.prim_func - def mma_fill_desc(C_warp: T.Buffer([WARP_SIZE, local_size], dtype=dtype, scope="warp")) -> None: + def mma_fill_desc(C_warp: T.Tensor([WARP_SIZE, local_size], dtype=dtype, scope="warp")) -> None: with Ts.sblock("root"): Ts.reads() Ts.writes(C_warp[0:WARP_SIZE, 0:local_size]) @@ -137,7 +137,7 @@ def mma_fill_desc(C_warp: T.Buffer([WARP_SIZE, local_size], dtype=dtype, scope=" @Ts.prim_func def mma_fill_impl( - C_warp: T.Buffer([WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1), + C_warp: T.Tensor([WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads() @@ -198,8 +198,8 @@ def get_mfma_load_intrin( @Ts.prim_func def mfma_load_desc( - reg: T.Buffer((WARP_SIZE, local_size), dtype, offset_factor=1, scope="warp"), - memory: T.Buffer(memory_shape, dtype, offset_factor=1, scope=scope), + reg: T.Tensor((WARP_SIZE, local_size), dtype, offset_factor=1, scope="warp"), + memory: T.Tensor(memory_shape, dtype, offset_factor=1, scope=scope), ) -> None: with Ts.sblock("root"): Ts.reads(memory[0:row_dim, 0:col_dim]) @@ -219,8 +219,8 @@ def mfma_load_desc( @Ts.prim_func def mfma_load_impl( - reg: T.Buffer((WARP_SIZE, local_size), dtype, align=64, offset_factor=1, scope="warp"), - memory: T.Buffer( + reg: T.Tensor((WARP_SIZE, local_size), dtype, align=64, offset_factor=1, scope="warp"), + memory: T.Tensor( memory_shape, dtype, align=64, offset_factor=1, scope=scope, strides=[s0, s1] ), ) -> None: @@ -268,9 +268,9 @@ def maybe_swap(i, j): @Ts.prim_func def mfma_sync_desc( - A: T.Buffer((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), - B: T.Buffer((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), - C: T.Buffer((WARP_SIZE, local_size_out), out_dtype, offset_factor=1, scope="warp"), + A: T.Tensor((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), + B: T.Tensor((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), + C: T.Tensor((WARP_SIZE, local_size_out), out_dtype, offset_factor=1, scope="warp"), ) -> None: with Ts.sblock("root"): Ts.reads( @@ -302,9 +302,9 @@ def mfma_sync_desc( @Ts.prim_func def mfma_sync_impl_float( - A: T.Buffer((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), - B: T.Buffer((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), - C: T.Buffer((WARP_SIZE, local_size_out), out_dtype, offset_factor=1, scope="warp"), + A: T.Tensor((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), + B: T.Tensor((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), + C: T.Tensor((WARP_SIZE, local_size_out), out_dtype, offset_factor=1, scope="warp"), ) -> None: with Ts.sblock("root"): Ts.reads( @@ -328,9 +328,9 @@ def mfma_sync_impl_float( @Ts.prim_func def mfma_sync_impl_integer( - A: T.Buffer((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), - B: T.Buffer((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), - C: T.Buffer((WARP_SIZE, local_size_out), out_dtype, offset_factor=1, scope="warp"), + A: T.Tensor((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), + B: T.Tensor((WARP_SIZE, local_size), in_dtype, offset_factor=1, scope="warp"), + C: T.Tensor((WARP_SIZE, local_size_out), out_dtype, offset_factor=1, scope="warp"), ) -> None: with Ts.sblock("root"): Ts.reads( @@ -373,8 +373,8 @@ def get_mfma_store_intrin(local_size=4, dtype="float32", scope="global"): @Ts.prim_func def mfma_store_desc( - C_warp: T.Buffer([WARP_SIZE, local_size], dtype=dtype, scope="warp"), - C: T.Buffer([M_DIM, N_DIM], dtype=dtype, scope=scope), + C_warp: T.Tensor([WARP_SIZE, local_size], dtype=dtype, scope="warp"), + C: T.Tensor([M_DIM, N_DIM], dtype=dtype, scope=scope), ) -> None: with Ts.sblock("root"): Ts.reads(C_warp[0:WARP_SIZE, 0:local_size]) @@ -392,8 +392,8 @@ def mfma_store_desc( @Ts.prim_func def mfma_store_impl( - C_warp: T.Buffer([WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1), - C: T.Buffer([M_DIM, N_DIM], dtype=dtype, scope=scope, offset_factor=1, strides=[s0, s1]), + C_warp: T.Tensor([WARP_SIZE, local_size], dtype=dtype, scope="warp", offset_factor=1), + C: T.Tensor([M_DIM, N_DIM], dtype=dtype, scope=scope, offset_factor=1, strides=[s0, s1]), ) -> None: with Ts.sblock("root"): Ts.reads(C_warp[0:WARP_SIZE, 0:local_size]) diff --git a/python/tvm/s_tir/tensor_intrin/x86.py b/python/tvm/s_tir/tensor_intrin/x86.py index 1a9a7b51e658..4f4d13515ac2 100644 --- a/python/tvm/s_tir/tensor_intrin/x86.py +++ b/python/tvm/s_tir/tensor_intrin/x86.py @@ -28,9 +28,9 @@ @Ts.prim_func def dot_product_16x4_u8i8i32_desc( - A: T.Buffer((4,), "uint8", offset_factor=1), - B: T.Buffer((16, 4), "int8", offset_factor=1), - C: T.Buffer((16,), "int32", offset_factor=1), + A: T.Tensor((4,), "uint8", offset_factor=1), + B: T.Tensor((16, 4), "int8", offset_factor=1), + C: T.Tensor((16,), "int32", offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:16], A[0:4], B[0:16, 0:4]) @@ -44,9 +44,9 @@ def dot_product_16x4_u8i8i32_desc( @Ts.prim_func def dot_product_16x4_u8i8i32_vnni( - A: T.Buffer((4,), "uint8", offset_factor=1), - B: T.Buffer((16, 4), "int8", offset_factor=1), - C: T.Buffer((16,), "int32", offset_factor=1), + A: T.Tensor((4,), "uint8", offset_factor=1), + B: T.Tensor((16, 4), "int8", offset_factor=1), + C: T.Tensor((16,), "int32", offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:16], A[0:4], B[0:16, 0:4]) @@ -70,9 +70,9 @@ def dot_product_16x4_u8i8i32_vnni( @Ts.prim_func def dot_product_16x4_u8i8i32_avx512( - A: T.Buffer((4,), "uint8", offset_factor=1), - B: T.Buffer((16, 4), "int8", offset_factor=1), - C: T.Buffer((16,), "int32", offset_factor=1), + A: T.Tensor((4,), "uint8", offset_factor=1), + B: T.Tensor((16, 4), "int8", offset_factor=1), + C: T.Tensor((16,), "int32", offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:16], A[0:4], B[0:16, 0:4]) diff --git a/python/tvm/script/parser/entry.py b/python/tvm/script/parser/entry.py index 537a22abaf66..14d0ef925b11 100644 --- a/python/tvm/script/parser/entry.py +++ b/python/tvm/script/parser/entry.py @@ -475,7 +475,7 @@ def capture(A, B): B[()] = A[x_value] # x_value resolved from enclosing scope @T.prim_func - def use(A: T.Buffer((1024,), "int32"), B: T.Buffer((), "int32")) -> None: + def use(A: T.Tensor((1024,), "int32"), B: T.Tensor((), "int32")) -> None: capture(A, B) # Produces B[()] = A[128] """ if function is not None and not inspect.isfunction(function): diff --git a/python/tvm/te/operation.py b/python/tvm/te/operation.py index af07e7f1d557..7fcbbe3c1041 100644 --- a/python/tvm/te/operation.py +++ b/python/tvm/te/operation.py @@ -308,7 +308,7 @@ def extern( raise ValueError("expect inputs to be tensor") if in_buffers is None: input_placeholders.append( - tvm.tirx.decl_buffer( + tvm.tirx.decl_tensor( t.shape, t.dtype, t.op.name, @@ -329,7 +329,7 @@ def extern( for shp, dt in zip(shape, dtype): output_placeholders.append( - tvm.tirx.decl_buffer( + tvm.tirx.decl_tensor( shp, dt, name, @@ -378,7 +378,7 @@ def extern_primfunc(input_tensors: list[_tensor.Tensor], primfunc: tvm.tirx.Prim B = te.placeholder((128, 128), name="B") @Ts.prim_func - def before_split(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: + def before_split(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): @@ -389,7 +389,7 @@ def before_split(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: C = te.extern_primfunc([A, B], func) """ - # Preserve the function parameter order while selecting BufferType annotations. + # Preserve the function parameter order while selecting TensorType annotations. dt_access_map = tvm.s_tir._ffi_api.DomainTouchedAccessMap(primfunc) ordered_buffers = [param for param in primfunc.params if tvm.tirx.is_buffer_var(param)] in_buffers = [buf for buf in ordered_buffers if len(dt_access_map[buf][0])] @@ -571,7 +571,7 @@ def create_prim_func( @Ts.prim_func def tir_matmul( - A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), B: T.Tensor((128, 128)), C: T.Tensor((128, 128)) ) -> None: for i, j, k in T.grid(128, 128, 128): diff --git a/python/tvm/tirx/__init__.py b/python/tvm/tirx/__init__.py index ca211a89454d..5af0fdab56c6 100644 --- a/python/tvm/tirx/__init__.py +++ b/python/tvm/tirx/__init__.py @@ -29,10 +29,10 @@ from .buffer import ( Buffer, BufferAccessKind, - BufferType, + TensorType, buffer_data, buffer_data_pointer_type, - decl_buffer, + decl_tensor, is_buffer_var, ) from .type import TensorMapType diff --git a/python/tvm/tirx/_buffer_view.py b/python/tvm/tirx/_buffer_view.py index 47d639947f3e..c16b5ca05cd8 100644 --- a/python/tvm/tirx/_buffer_view.py +++ b/python/tvm/tirx/_buffer_view.py @@ -40,7 +40,7 @@ def _redecl(buf: Buffer, shape, layout, *, dtype=None, elem_offset=None, addr_of addr = buf.allocated_addr[0] if addr_offset is not None: addr = addr + addr_offset - return tvm.tirx.script.ir_builder.decl_buffer( + return tvm.tirx.script.ir_builder.decl_tensor( shape, buf.dtype if dtype is None else dtype, None, @@ -53,7 +53,7 @@ def _redecl(buf: Buffer, shape, layout, *, dtype=None, elem_offset=None, addr_of layout, allocated_addr=addr, ) - return tvm.tirx.script.ir_builder.decl_buffer( + return tvm.tirx.script.ir_builder.decl_tensor( shape, buf.dtype if dtype is None else dtype, buf.data, diff --git a/python/tvm/tirx/bench.py b/python/tvm/tirx/bench.py index c4bab4f2d105..9dc425804fbd 100644 --- a/python/tvm/tirx/bench.py +++ b/python/tvm/tirx/bench.py @@ -917,7 +917,7 @@ class CudaProfiler: def __init__( self, - profiler_buffer: T.Buffer, + profiler_buffer: T.Tensor, write_stride: int, num_groups: int, default_leader: None | tvm.tirx.Expr | bool = None, @@ -935,8 +935,8 @@ def __init__( # Assume Expr-like input; use as-is self.profiler_enabled = profiler_enabled # type: ignore[assignment] - self.profiler_tag = T.alloc_buffer([1], "uint64", scope="local", align=8) - self.profiler_write_offset = T.alloc_buffer([1], "uint32", scope="local", align=8) + self.profiler_tag = T.alloc_tensor([1], "uint64", scope="local", align=8) + self.profiler_write_offset = T.alloc_tensor([1], "uint32", scope="local", align=8) def _leader(self, leader: None | tvm.tirx.Expr | bool): if leader is not None: diff --git a/python/tvm/tirx/buffer.py b/python/tvm/tirx/buffer.py index b2db4c86d75d..1161ef034809 100644 --- a/python/tvm/tirx/buffer.py +++ b/python/tvm/tirx/buffer.py @@ -30,9 +30,9 @@ _REARRANGE_PATTERN_UNSET = object() -@tvm_ffi.register_object("tirx.BufferType") -class BufferType(Type): - """The structural type carried by an ordinary buffer variable.""" +@tvm_ffi.register_object("tirx.TensorType") +class TensorType(Type): + """The structural type carried by an ordinary TIRx tensor variable.""" dtype: PrimType storage_scope: str @@ -46,14 +46,14 @@ class BufferType(Type): def is_buffer_var(value) -> bool: - """Return whether ``value`` is an ordinary Var carrying BufferType. + """Return whether ``value`` is an ordinary Var carrying TensorType. Use this predicate instead of ``isinstance(value, Buffer)``. ``Buffer`` is a source-compatibility alias for :class:`tvm.ir.Var` and therefore does not discriminate buffer variables from scalar or pointer variables. """ - return isinstance(value, tvm.ir.Var) and isinstance(value.ty, BufferType) + return isinstance(value, tvm.ir.Var) and isinstance(value.ty, TensorType) class BufferAccessKind(IntEnum): @@ -69,12 +69,12 @@ class _BufferMethods: Buffer provide a way to represent data layout specialization of data structure in TVM. - Do not construct directly, use :py:func:`~decl_buffer` instead. - See the documentation of :py:func:`decl_buffer` for more details. + Do not construct directly, use :py:func:`~decl_tensor` instead. + See the documentation of :py:func:`decl_tensor` for more details. See Also -------- - decl_buffer : Declare a buffer + decl_tensor : Declare a buffer """ def access_ptr(self, access_mask, ptr_type="handle", content_lanes=1, offset=0, extent=None): @@ -314,7 +314,7 @@ def view(self, *args, **kwargs) -> "Buffer": Returns ------- - view : DeclBufferFrame + view : DeclTensorFrame The corresponding view buffer. """ @@ -351,7 +351,7 @@ def local(self, *shape, layout=None) -> "Buffer": Returns ------- - local : DeclBufferFrame + local : DeclTensorFrame The corresponding local buffer. """ return _buffer_view.local(self, *shape, layout=layout) @@ -366,7 +366,7 @@ def permute(self, *dims) -> "Buffer": Returns ------- - permuted : DeclBufferFrame + permuted : DeclTensorFrame The buffer with permuted dimensions. """ return _buffer_view.permute(self, *dims) @@ -463,7 +463,7 @@ def chunk(self, spec) -> "_buffer_view.ChunkIndexer": # definitions precede this class to avoid a circular buffer/view import. -def decl_buffer( +def decl_tensor( shape, dtype=None, name="buffer", @@ -497,7 +497,7 @@ def decl_buffer( if not isinstance(data.ty.element_type, PrimType): raise TypeError("Buffer data must point to a primitive type") storage_scope = data.ty.storage_scope - buffer_type = _ffi_api.BufferType( # type: ignore + buffer_type = _ffi_api.TensorType( # type: ignore storage_scope, dtype, shape, @@ -516,7 +516,7 @@ def buffer_data(buffer): """Project the physical pointer associated with a buffer variable.""" if not is_buffer_var(buffer): - raise TypeError("buffer_data expects a Var with BufferType") + raise TypeError("buffer_data expects a Var with TensorType") return _ffi_api.BufferData(buffer) @@ -524,7 +524,7 @@ def buffer_data_pointer_type(buffer): """Return the pointer type produced by :func:`buffer_data`.""" if not is_buffer_var(buffer): - raise TypeError("buffer_data_pointer_type expects a Var with BufferType") + raise TypeError("buffer_data_pointer_type expects a Var with TensorType") return _ffi_api.BufferDataPointerType(buffer) @@ -542,13 +542,13 @@ def buffer_data_pointer_type(buffer): def _buffer_type_field(name): def getter(value): if not is_buffer_var(value): - raise AttributeError(f"{name} is only available on a Var with BufferType") + raise AttributeError(f"{name} is only available on a Var with TensorType") return getattr(value.ty, name) return property(getter) -# Preserve Buffer's public metadata surface while keeping BufferType as the +# Preserve Buffer's public metadata surface while keeping TensorType as the # single source of truth. for _name in ( "shape", @@ -564,8 +564,8 @@ def getter(value): def _buffer_dtype_property(value): if not is_buffer_var(value): - raise AttributeError("dtype is only available on a Var with BufferType") - # Preserve the pre-migration Python Buffer surface. BufferType stores a + raise AttributeError("dtype is only available on a Var with TensorType") + # Preserve the pre-migration Python Buffer surface. TensorType stores a # PrimType, while Python callers historically receive its runtime DataType. return value.ty.dtype.dtype @@ -577,7 +577,7 @@ def _buffer_dtype_property(value): # and builder code calls ``buffer_data(A)`` directly. def _buffer_data_property(value): if not is_buffer_var(value): - raise AttributeError("data is only available on a Var with BufferType") + raise AttributeError("data is only available on a Var with TensorType") return buffer_data(value) diff --git a/python/tvm/tirx/function.py b/python/tvm/tirx/function.py index 37becd80e6da..f4a60a7d7121 100644 --- a/python/tvm/tirx/function.py +++ b/python/tvm/tirx/function.py @@ -127,8 +127,8 @@ def specialize(self, param_map: Mapping[Var, Expr | Buffer]): @T.prim_func def mem_copy( - A: T.Buffer((m, n), "float32"), - B: T.Buffer((m, n), "float32"), + A: T.Tensor((m, n), "float32"), + B: T.Tensor((m, n), "float32"), m: T.int32, n: T.int32, ) -> None: @@ -141,7 +141,7 @@ def mem_copy( .. code-block:: python a, _, m, n = mem_copy.params - func = mem_copy.specialize({a: tirx.decl_buffer((16, 16))}) + func = mem_copy.specialize({a: tirx.decl_tensor((16, 16))}) # or func = mem_copy.specialize({n: 16, m: 16}) @@ -151,7 +151,7 @@ def mem_copy( @T.prim_func def mem_copy_16_16( - A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), B: T.Tensor((16, 16), "float32") ) -> None: for i, j in T.grid(16, 16): diff --git a/python/tvm/tirx/script/ir_builder/frame.py b/python/tvm/tirx/script/ir_builder/frame.py index c46a97af0bb7..ca41a6d1970d 100644 --- a/python/tvm/tirx/script/ir_builder/frame.py +++ b/python/tvm/tirx/script/ir_builder/frame.py @@ -88,8 +88,8 @@ class ThenFrame(TIRFrame): ... class ElseFrame(TIRFrame): ... -@_register_object("script.ir_builder.tirx.DeclBufferFrame") -class DeclBufferFrame(TIRFrame): +@_register_object("script.ir_builder.tirx.DeclTensorFrame") +class DeclTensorFrame(TIRFrame): def __enter__(self) -> Buffer: super().__enter__() return self.buffer diff --git a/python/tvm/tirx/script/ir_builder/ir.py b/python/tvm/tirx/script/ir_builder/ir.py index 4e8fda52f827..a31203640e9f 100644 --- a/python/tvm/tirx/script/ir_builder/ir.py +++ b/python/tvm/tirx/script/ir_builder/ir.py @@ -208,9 +208,9 @@ def _record_meta_resource(value: Any, skip_frames: int = 2) -> None: scope.record(value, frame_info) -@_register_mutable_decl("tirx.Buffer", syntax="parameter") +@_register_mutable_decl("tirx.Tensor", syntax="parameter") @_annotation_constructor -def _buffer_type( +def _tensor_type( shape: list[Expr] | tuple[Expr] | Expr | Integral, dtype: str = "float32", data: Var = None, @@ -224,8 +224,8 @@ def _buffer_type( allocated_addr: int | tuple[int, ...] | None = None, *, span=None, -) -> tir.BufferType: - """Construct a buffer type for annotations and explicit type-valued fields. +) -> tir.TensorType: + """Construct a tensor type for annotations and explicit type-valued fields. Parameters ---------- @@ -265,9 +265,9 @@ def _buffer_type( Returns ------- - res : BufferType - The buffer type. Function annotations introduce variables of this type; - allocation and declaration operations construct buffer variables. + res : TensorType + The tensor type. Function annotations introduce variables of this type; + allocation and declaration operations construct tensor variables. """ shape = (shape,) if is_prim_expr(shape) or isinstance(shape, Integral) else shape shape = tuple(shape) @@ -277,7 +277,7 @@ def _buffer_type( allocated_addr = [] if not isinstance(allocated_addr, list | tuple): allocated_addr = [allocated_addr] - result = _ffi_api.BufferType( # type: ignore[attr-defined] # pylint: disable=no-member + result = _ffi_api.TensorType( # type: ignore[attr-defined] # pylint: disable=no-member shape, dtype, data, @@ -548,8 +548,8 @@ def thread_id_in_wg( return tuple(ret) -@_register_mutable_decl("tirx.alloc_buffer") -def alloc_buffer( +@_register_mutable_decl("tirx.alloc_tensor") +def alloc_tensor( shape: list[Expr] | tuple[Expr] | Expr | Integral, dtype: str = "float32", data: Var | None = None, @@ -563,11 +563,11 @@ def alloc_buffer( allocated_addr: int | tuple[int, ...] | None = None, annotations: dict[str, Any] | None = None, ) -> Buffer: - """Statement-level buffer allocation (creates a buffer-returning allocation Call). + """Allocate a tensor and return its variable. - Emits a Bind statement with an allocation Call and returns the Buffer directly:: + Emits a Bind statement with a ``tirx.alloc_tensor`` Call:: - buf = T.alloc_buffer((128, 128)) + buf = T.alloc_tensor((128, 128)) Parameters ---------- @@ -634,7 +634,7 @@ def _normalize_ann_value(v): norm_annotations = {k: _normalize_ann_value(v) for k, v in (annotations or {}).items()} allocation = ir.Call( - "tirx.alloc_buffer", + "tirx.alloc_tensor", [ ir.Tuple(buf.shape), ir.DataTypeImm(DataType(buf.dtype)), @@ -652,7 +652,7 @@ def wg_reg_tile(elem_per_thread: int, dtype: str = "float32") -> Buffer: Sugar for the recurring pattern:: - T.alloc_buffer( + T.alloc_tensor( (128, elem_per_thread), dtype, layout=wg_local_layout(elem_per_thread), scope="local", @@ -661,7 +661,7 @@ def wg_reg_tile(elem_per_thread: int, dtype: str = "float32") -> Buffer: Used to stage a tcgen05 load: each of the 128 threads in a warpgroup owns one row of ``elem_per_thread`` contiguous elements. """ - return alloc_buffer( + return alloc_tensor( (128, elem_per_thread), dtype, layout=wg_local_layout(elem_per_thread), @@ -756,8 +756,8 @@ def __repr__(self): return f"DtypeConstructor({self._dtype_str!r})" -@_register_mutable_decl("tirx.decl_buffer") -def decl_buffer( +@_register_mutable_decl("tirx.decl_tensor") +def decl_tensor( shape, dtype="float32", data=None, @@ -770,10 +770,11 @@ def decl_buffer( layout=MISSING, allocated_addr=None, ) -> Buffer: - """Create a buffer declaration node. + """Declare a tensor backed by a pointer, or allocate its storage. - When ``data`` is provided, creates a DeclBuffer (alias to existing data). - When ``data`` is None, creates an AllocBuffer (new allocation). + With ``data``, bind a ``tirx.decl_tensor`` Call to the tensor variable. + Without ``data``, bind a ``tirx.alloc_tensor`` Call instead. The ``tmem`` + scope uses ``allocated_addr`` to declare externally allocated tensor memory. Parameters ---------- @@ -817,7 +818,7 @@ def decl_buffer( if strides is None: strides = [] dtype = _normalize_prim_type(dtype) - decl_frame = _ffi_api.DeclBuffer( # type: ignore[attr-defined] # pylint: disable=no-member + decl_frame = _ffi_api.DeclTensor( # type: ignore[attr-defined] # pylint: disable=no-member shape, dtype, "", @@ -830,7 +831,7 @@ def decl_buffer( _get_layout(layout, shape, scope), allocated_addr, ) - if isinstance(decl_frame, frame.DeclBufferFrame): + if isinstance(decl_frame, frame.DeclTensorFrame): decl_frame.add_callback(partial(decl_frame.__exit__, None, None, None)) buf = decl_frame.__enter__() else: @@ -839,17 +840,17 @@ def decl_buffer( return buf -alloc_shared = functools.partial(alloc_buffer, scope="shared") +alloc_shared = functools.partial(alloc_tensor, scope="shared") _register_mutable_decl("tirx.alloc_shared")(alloc_shared) -alloc_local = functools.partial(alloc_buffer, scope="local") +alloc_local = functools.partial(alloc_tensor, scope="local") _register_mutable_decl("tirx.alloc_local")(alloc_local) smem = _register_mutable_decl("tirx.smem")(alloc_shared) -tmem = functools.partial(alloc_buffer, scope="tmem") +tmem = functools.partial(alloc_tensor, scope="tmem") def alloc_tcgen05_ldst_frag(instr_shape, tensor_shape, dtype): @@ -949,7 +950,7 @@ def alloc_scalar( annotations: dict[str, Any] | None = None, ) -> TensorLoad: """Allocate a zero-dimensional buffer (scalar), with optional allocation annotations.""" - buf = alloc_buffer( + buf = alloc_tensor( shape=(1,), dtype=dtype, scope=scope, layout=TileLayout(S[1]), annotations=annotations ) assert is_buffer_var(buf) @@ -960,7 +961,7 @@ def alloc_scalar( @_register_mutable_decl("tirx.decl_scalar") def decl_scalar(dtype, data, scope, elem_offset=None, byte_offset=None) -> TensorLoad: """Declare a zero-dimensional buffer (scalar) from a pointer.""" - buf = decl_buffer( + buf = decl_tensor( shape=(1,), dtype=dtype, data=data, @@ -1688,10 +1689,9 @@ def Ptr(dtype, storage_scope="global", *, span=None): return _at(span, ptr(dtype, storage_scope)) -Buffer = _buffer_type +Tensor = _tensor_type __all__ = [ - "Buffer", "BufferLoad", "ComposeLayout", "DtypeConstructor", @@ -1708,16 +1708,17 @@ def Ptr(dtype, storage_scope="global", *, span=None): "Range", "S", "ScopeIdDef", + "Tensor", "TensorMap", "TileLayout", "Tuple", "Var", - "alloc_buffer", "alloc_cast_frag", "alloc_local", "alloc_scalar", "alloc_shared", "alloc_tcgen05_ldst_frag", + "alloc_tensor", "bf16", "bfloat16", "boolean", @@ -1726,8 +1727,8 @@ def Ptr(dtype, storage_scope="global", *, span=None): "cta_id", "cta_id_in_cluster", "cta_id_in_pair", - "decl_buffer", "decl_scalar", + "decl_tensor", "f16", "f32", "f64", diff --git a/python/tvm/tirx/script/ir_builder/parser_protocol.py b/python/tvm/tirx/script/ir_builder/parser_protocol.py index 1d0a9cc3eadd..dfdde9ff4b5a 100644 --- a/python/tvm/tirx/script/ir_builder/parser_protocol.py +++ b/python/tvm/tirx/script/ir_builder/parser_protocol.py @@ -166,7 +166,7 @@ def device_entry() -> None: @T.prim_func def kernel(...): - A = T.Buffer(...) + A = T.Tensor(...) T.device_entry() # device region starts here bx = T.cta_id([SM_COUNT]) # standalone scope-id def ... diff --git a/python/tvm/tirx/script/ir_builder/tirx.py b/python/tvm/tirx/script/ir_builder/tirx.py index 5ac7f0769bec..77d1882da549 100644 --- a/python/tvm/tirx/script/ir_builder/tirx.py +++ b/python/tvm/tirx/script/ir_builder/tirx.py @@ -28,7 +28,7 @@ from tvm.tirx.lang.alloc_pool import SMEMPool, TMEMPool from . import _ffi_api -from .ir import decl_buffer, meta_class +from .ir import decl_tensor, meta_class def _normalize_scope(scope) -> ExecScope: @@ -1700,7 +1700,7 @@ def reshape(buffer: Buffer, shape: list[Expr]): + " are not compatible" ) - return decl_buffer( + return decl_tensor( shape, buffer.ty.dtype, buffer_data(buffer), diff --git a/python/tvm/tirx/script/jit.py b/python/tvm/tirx/script/jit.py index b2bd39464dda..f33a0758b174 100644 --- a/python/tvm/tirx/script/jit.py +++ b/python/tvm/tirx/script/jit.py @@ -113,8 +113,8 @@ def jit( @T.jit def add( - A: T.Buffer((N,), "float32"), - B: T.Buffer((N,), "float32"), + A: T.Tensor((N,), "float32"), + B: T.Tensor((N,), "float32"), *, N: T.constexpr, ): @@ -125,8 +125,8 @@ def add( @T.jit def guarded( - optional: T.Optional(T.Buffer((1,), "int32")), - output: T.Buffer((1,), "int32"), + optional: T.Optional(T.Tensor((1,), "int32")), + output: T.Tensor((1,), "int32"), ): if T.constexpr(optional is not None): output[0] = optional[0] @@ -176,7 +176,7 @@ class TIRJit: type to what ``@T.prim_func`` produces today. Constexpr params are removed from the resulting PrimFunc's parameter list; - their values are baked into the IR (e.g. into ``T.Buffer((M, K), ...)`` + their values are baked into the IR (e.g. into ``T.Tensor((M, K), ...)`` shape annotations and into the body). """ diff --git a/python/tvm/tirx/stmt.py b/python/tvm/tirx/stmt.py index 9dc819aecff8..e24a4c48b15e 100644 --- a/python/tvm/tirx/stmt.py +++ b/python/tvm/tirx/stmt.py @@ -21,7 +21,7 @@ .. code-block:: python x = tvm.tirx.Var("n", "int32") - buffer = tvm.tirx.decl_buffer((16,), "float32") + buffer = tvm.tirx.decl_tensor((16,), "float32") st = tvm.tirx.stmt.BufferStore(buffer, 1, (x,)) assert isinstance(st, tvm.tirx.stmt.BufferStore) assert(st.buffer == buffer) diff --git a/python/tvm/tirx/tile_primitive.py b/python/tvm/tirx/tile_primitive.py index c881b895ab02..3118130d01c5 100644 --- a/python/tvm/tirx/tile_primitive.py +++ b/python/tvm/tirx/tile_primitive.py @@ -157,7 +157,7 @@ def add_init_stmt(self, stmt: Stmt, host: bool = False) -> None: _ffi_api.DispatchContextAddInitStmt(self, stmt, host) # pylint: disable=no-member def add_post_buffer_def_stmt(self, buffer: Buffer, stmt: Stmt) -> None: - """Add a statement to be inserted after a buffer's definition (DeclBuffer/AllocBuffer). + """Add a statement to be inserted after a buffer's definition (DeclTensor/AllocTensor). Parameters ---------- diff --git a/python/tvm/tirx/transform/common.py b/python/tvm/tirx/transform/common.py index bdbb43c174bc..6d664e4514ae 100644 --- a/python/tvm/tirx/transform/common.py +++ b/python/tvm/tirx/transform/common.py @@ -19,7 +19,7 @@ import tvm_ffi from tvm.ir import Call, Op, is_prim_expr -from tvm.tirx import Evaluate, Expr, Stmt, TilePrimitiveCall, Var, decl_buffer +from tvm.tirx import Evaluate, Expr, Stmt, TilePrimitiveCall, Var, decl_tensor from tvm.tirx.buffer import Buffer, is_buffer_var from tvm.tirx.layout import Iter, TileLayout @@ -111,7 +111,7 @@ def _mutate_buffer(self, buffer: Buffer): if unchanged: return buffer - new_buffer = decl_buffer( + new_buffer = decl_tensor( new_shape, buffer.ty.dtype, buffer.name, diff --git a/python/tvm/tirx/transform/transform.py b/python/tvm/tirx/transform/transform.py index aecaff75198e..365a894343dd 100644 --- a/python/tvm/tirx/transform/transform.py +++ b/python/tvm/tirx/transform/transform.py @@ -271,17 +271,17 @@ def MakePackedAPI(): """Transform the PrimFuncs in the module to a packed func API. Prior to this pass, the PrimFunc may have parameters annotated with - `BufferType`. This pass consumes those annotations to generate + `TensorType`. This pass consumes those annotations to generate arguments that implement the packed based TVM FFI API. - For static shapes, the `BufferType::shape`, `BufferType::strides`, - and `BufferType::elem_offset` fields are used to + For static shapes, the `TensorType::shape`, `TensorType::strides`, + and `TensorType::elem_offset` fields are used to generate runtime checks on the corresponding member variables in the user-provided `DLTensor*` or `tvm.runtime.tensor` argument. (e.g. A PrimFunc that accepts a buffer of shape `[16,32]` validates that the `DLTensor::shape` array is `[16,32]`.) - For dynamic Buffers, in which one or more of these `BufferType` fields + For dynamic Buffers, in which one or more of these `TensorType` fields use `tirx.Var` that are not defined by other PrimFunc parameters, these are instead used to define the variables based on the corresponding `DLTensor` members. (e.g. A PrimFunc that accepts a @@ -539,7 +539,7 @@ def LowerTIRx(): def LowerTIRxOpaque(): """Lower opaque constructs in TIRX programs. - Handles AllocBuffer lowering, For(thread_binding) to AttrStmt(thread_extent) + Handles allocation call lowering, For(thread_binding) to AttrStmt(thread_extent) conversion, unit loop elimination, and pragma annotation handling. This is the tirx-specific counterpart of s_tir.LowerOpaqueBlock, without any SBlock/SBlockRealize handling. diff --git a/python/tvm/topi/gpu/scan.py b/python/tvm/topi/gpu/scan.py index 2ac3bb2e0fd3..5e14c16a8e1e 100644 --- a/python/tvm/topi/gpu/scan.py +++ b/python/tvm/topi/gpu/scan.py @@ -137,9 +137,9 @@ def exclusive_scan_ir(data, output, reduction=None, binop=operator.add, identity tx = te.thread_axis("threadIdx.x") bx = te.thread_axis("blockIdx.x") blocks_per_batch = cast(ceil_div(scan_axis_size, max_threads * width), "int32") - start_buf = T.decl_buffer([1], "int32", scope="local") - middle_buf = T.decl_buffer([1], "int32", scope="local") - end_buf = T.decl_buffer([1], "int32", scope="local") + start_buf = T.decl_tensor([1], "int32", scope="local") + middle_buf = T.decl_tensor([1], "int32", scope="local") + end_buf = T.decl_tensor([1], "int32", scope="local") with T.frame_scope( [ T.attr(tx, "thread_extent", nthread_tx), @@ -231,10 +231,10 @@ def exclusive_scan_ir(data, output, reduction=None, binop=operator.add, identity tx = te.thread_axis("threadIdx.x") bx = te.thread_axis("blockIdx.x") blocks_per_batch = cast(ceil_div(scan_axis_size, max_threads * width), "int32") - start_buf = T.decl_buffer([1], "int32", scope="local") - middle_buf = T.decl_buffer([1], "int32", scope="local") - end_buf = T.decl_buffer([1], "int32", scope="local") - tmp_buf = T.decl_buffer([1], out_dtype, scope="local") + start_buf = T.decl_tensor([1], "int32", scope="local") + middle_buf = T.decl_tensor([1], "int32", scope="local") + end_buf = T.decl_tensor([1], "int32", scope="local") + tmp_buf = T.decl_tensor([1], out_dtype, scope="local") with T.frame_scope( [ T.attr(tx, "thread_extent", nthread_tx), @@ -412,10 +412,10 @@ def ir(data_buf, data_ex_scan_buf, reduction_buf): return ib.get() - data_buf = tvm.tirx.decl_buffer( + data_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "valid_indices_buf", data_alignment=8, layout=None ) - ex_scan_output_buf = tvm.tirx.decl_buffer( + ex_scan_output_buf = tvm.tirx.decl_tensor( ex_scan_output.shape, ex_scan_output.dtype, "ex_scan_output_buf", @@ -484,15 +484,15 @@ def scan_thrust( (N-1)-D tensor storing the reduction of each scan axis. Returned if return_reduction is True. """ - data_buf = tvm.tirx.decl_buffer( + data_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "data_buf", data_alignment=8, layout=None ) - output_buf = tvm.tirx.decl_buffer( + output_buf = tvm.tirx.decl_tensor( data.shape, output_dtype, "output_buf", data_alignment=8, layout=None ) workspace_buf = ( - tvm.tirx.decl_buffer( + tvm.tirx.decl_tensor( workspace.shape, workspace.dtype, "workspace_buf", data_alignment=8, layout=None ) if workspace is not None @@ -593,10 +593,10 @@ def do_scan(data, output_dtype): # TIR exclusive scan accepts only 2D or higher-rank inputs. data = expand_dims(data, axis=0) - data_buf = tvm.tirx.decl_buffer( + data_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "data_buf", data_alignment=8, layout=None ) - output_buf = tvm.tirx.decl_buffer( + output_buf = tvm.tirx.decl_tensor( data.shape, output_dtype, "output_buf", data_alignment=8, layout=None ) diff --git a/python/tvm/topi/gpu/scatter_elements.py b/python/tvm/topi/gpu/scatter_elements.py index 17d1491a6362..d46a75ca8264 100644 --- a/python/tvm/topi/gpu/scatter_elements.py +++ b/python/tvm/topi/gpu/scatter_elements.py @@ -174,7 +174,7 @@ def max_func(dst_ptr, dst_index, update): "scatter_elements reduction not in [update, add, mul, mean, min, max]:", reduction ) - out_buf = tirx.decl_buffer(data.shape, data.dtype, "out_buf", layout=None) + out_buf = tirx.decl_tensor(data.shape, data.dtype, "out_buf", layout=None) return te.extern( [data.shape], [data, indices, updates], diff --git a/python/tvm/topi/gpu/scatter_nd.py b/python/tvm/topi/gpu/scatter_nd.py index 01bdc800a107..f6b19014ae7f 100644 --- a/python/tvm/topi/gpu/scatter_nd.py +++ b/python/tvm/topi/gpu/scatter_nd.py @@ -176,7 +176,7 @@ def gen_ir(data_ptr, indices_ptr, updates_ptr, out_ptr): return ib.get() - out_buf = tirx.decl_buffer(data.shape, data.dtype, "out_buf", layout=None) + out_buf = tirx.decl_tensor(data.shape, data.dtype, "out_buf", layout=None) return te.extern( [data.shape], [data, indices, updates], diff --git a/python/tvm/topi/gpu/sort.py b/python/tvm/topi/gpu/sort.py index 5b44906f2acd..f37013402d8f 100644 --- a/python/tvm/topi/gpu/sort.py +++ b/python/tvm/topi/gpu/sort.py @@ -122,14 +122,14 @@ def _odd_even_sort( tid = 2 * tx start = bx * block_size - # Buffer declarations (DeclBuffer generates both Allocate + DeclBuffer nodes) - tmp_keys_swap = T.decl_buffer([block_size], keys_swap.dtype, scope="shared") - temp_keys = T.decl_buffer([1], keys_swap.dtype, scope="local") - temp_cond1 = T.decl_buffer([1], keys_swap.dtype, scope="local") - temp_cond2 = T.decl_buffer([1], keys_swap.dtype, scope="local") + # Tensor declarations without data allocate their storage. + tmp_keys_swap = T.decl_tensor([block_size], keys_swap.dtype, scope="shared") + temp_keys = T.decl_tensor([1], keys_swap.dtype, scope="local") + temp_cond1 = T.decl_tensor([1], keys_swap.dtype, scope="local") + temp_cond2 = T.decl_tensor([1], keys_swap.dtype, scope="local") if values_swap is not None: - tmp_values_swap = T.decl_buffer([block_size], values_swap.dtype, scope="shared") - temp_values = T.decl_buffer([1], values_swap.dtype, scope="local") + tmp_values_swap = T.decl_tensor([block_size], values_swap.dtype, scope="shared") + temp_values = T.decl_tensor([1], values_swap.dtype, scope="local") # Copy data to scratch space base_idx = by_val * size * axis_mul_after + bz @@ -411,10 +411,10 @@ def mergepath( step_count, even, ): - first_buf = T.decl_buffer([1], target_dtype, scope="local") - last_buf = T.decl_buffer([1], target_dtype, scope="local") - i_buf_buf = T.decl_buffer([1], target_dtype, scope="local") - j_buf_buf = T.decl_buffer([1], target_dtype, scope="local") + first_buf = T.decl_tensor([1], target_dtype, scope="local") + last_buf = T.decl_tensor([1], target_dtype, scope="local") + i_buf_buf = T.decl_tensor([1], target_dtype, scope="local") + j_buf_buf = T.decl_tensor([1], target_dtype, scope="local") first = first_buf last = last_buf i_buf = i_buf_buf @@ -478,12 +478,12 @@ def dual_mergepath( step_count, even, ): - outer_first_buf = T.decl_buffer([1], target_dtype, scope="local") - outer_last_buf = T.decl_buffer([1], target_dtype, scope="local") - first_buf = T.decl_buffer([1], target_dtype, scope="local") - last_buf = T.decl_buffer([1], target_dtype, scope="local") - i_buf_buf = T.decl_buffer([1], target_dtype, scope="local") - j_buf_buf = T.decl_buffer([1], target_dtype, scope="local") + outer_first_buf = T.decl_tensor([1], target_dtype, scope="local") + outer_last_buf = T.decl_tensor([1], target_dtype, scope="local") + first_buf = T.decl_tensor([1], target_dtype, scope="local") + last_buf = T.decl_tensor([1], target_dtype, scope="local") + i_buf_buf = T.decl_tensor([1], target_dtype, scope="local") + j_buf_buf = T.decl_tensor([1], target_dtype, scope="local") outer_first = outer_first_buf outer_last = outer_last_buf first = first_buf @@ -782,10 +782,10 @@ def sort(data, axis=-1, is_ascend=1): axes = swap(list(range(ndim)), axis) data = transpose(data, axes) - value_buf = tvm.tirx.decl_buffer( + value_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "value_buf", data_alignment=8, layout=None ) - value_buf_swap = tvm.tirx.decl_buffer( + value_buf_swap = tvm.tirx.decl_tensor( data.shape, data.dtype, "value_buf_swap", data_alignment=8, layout=None ) @@ -840,10 +840,10 @@ def sort_thrust(data, axis=-1, is_ascend=1, workspace=None): axes = swap(list(range(ndim)), axis) data = transpose(data, axes) - value_buf = tvm.tirx.decl_buffer( + value_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "value_buf", data_alignment=8, layout=None ) - indices_buf = tvm.tirx.decl_buffer(data.shape, dtype, "out_buf", data_alignment=8, layout=None) + indices_buf = tvm.tirx.decl_tensor(data.shape, dtype, "out_buf", data_alignment=8, layout=None) def f_compute(ins, outs): args = ["tvm.contrib.thrust.sort", ins[0], outs[0], outs[1], is_ascend] @@ -904,14 +904,14 @@ def argsort(data, axis=-1, is_ascend=1, dtype="float32", ret_type="indices"): axes = swap(list(range(ndim)), axis) data = transpose(data, axes) - value_buf = tvm.tirx.decl_buffer( + value_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "value_buf", data_alignment=8, layout=None ) - value_swap_buf = tvm.tirx.decl_buffer( + value_swap_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "value_swap_buf", data_alignment=8, layout=None ) - indices_buf = tvm.tirx.decl_buffer(data.shape, dtype, "out_buf", data_alignment=8, layout=None) - indices_swap_buf = tvm.tirx.decl_buffer( + indices_buf = tvm.tirx.decl_tensor(data.shape, dtype, "out_buf", data_alignment=8, layout=None) + indices_swap_buf = tvm.tirx.decl_tensor( data.shape, dtype, "out_swap_buf", data_alignment=8, layout=None ) @@ -1018,16 +1018,16 @@ def topk(data, k=1, axis=-1, ret_type="both", is_ascend=False, dtype="int64"): axes = swap(list(range(ndim)), axis) data = transpose(data, axes) - values_buf = tvm.tirx.decl_buffer( + values_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "values_buf", data_alignment=8, layout=None ) - values_swap_buf = tvm.tirx.decl_buffer( + values_swap_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "values_swap_buf", data_alignment=8, layout=None ) - indices_buf = tvm.tirx.decl_buffer( + indices_buf = tvm.tirx.decl_tensor( data.shape, dtype, "indices_buf", data_alignment=8, layout=None ) - indices_swap_buf = tvm.tirx.decl_buffer( + indices_swap_buf = tvm.tirx.decl_tensor( data.shape, dtype, "indies_swap_buf", data_alignment=8, layout=None ) @@ -1129,18 +1129,18 @@ def topk_thrust( axes = swap(list(range(ndim)), axis) data = transpose(data, axes) - data_buf = tvm.tirx.decl_buffer( + data_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "data_buf", data_alignment=8, layout=None ) if workspace is not None: - workspace_buf = tvm.tirx.decl_buffer( + workspace_buf = tvm.tirx.decl_tensor( workspace.shape, workspace.dtype, "workspace_buf", data_alignment=8, layout=None ) else: workspace_buf = None out_bufs = [ - tvm.tirx.decl_buffer(data.shape, data.dtype, "value_buf", data_alignment=8, layout=None), - tvm.tirx.decl_buffer(data.shape, dtype, "indices_buf", data_alignment=8, layout=None), + tvm.tirx.decl_tensor(data.shape, data.dtype, "value_buf", data_alignment=8, layout=None), + tvm.tirx.decl_tensor(data.shape, dtype, "indices_buf", data_alignment=8, layout=None), ] def f_compute(ins, outs): diff --git a/python/tvm/topi/index_put.py b/python/tvm/topi/index_put.py index 029c681587f6..22950e6dd4e3 100644 --- a/python/tvm/topi/index_put.py +++ b/python/tvm/topi/index_put.py @@ -157,7 +157,7 @@ def add_func(dst_ptr, dst_index, update): in_buffers.extend(indices) in_buffers.append(values) - out_buf = tirx.decl_buffer(data.shape, data.dtype, "out_buf", layout=None) + out_buf = tirx.decl_tensor(data.shape, data.dtype, "out_buf", layout=None) return te.extern( [data.shape], in_buffers, diff --git a/python/tvm/topi/scan.py b/python/tvm/topi/scan.py index a39632ee0d0d..cae4baa666e8 100644 --- a/python/tvm/topi/scan.py +++ b/python/tvm/topi/scan.py @@ -25,7 +25,7 @@ from tvm.tirx.script import ir_builder as T from ..te import extern -from ..tirx import decl_buffer +from ..tirx import decl_tensor from . import utils from .math import cast @@ -156,7 +156,7 @@ def gen_ir(data_buf, out_buf): return ib.get() - out_buf = decl_buffer(shape, dtype, "out_buf") + out_buf = decl_tensor(shape, dtype, "out_buf") return extern( [shape], diff --git a/python/tvm/topi/scatter.py b/python/tvm/topi/scatter.py index 1fcd77792372..7c7800bd8e1d 100644 --- a/python/tvm/topi/scatter.py +++ b/python/tvm/topi/scatter.py @@ -193,7 +193,7 @@ def gen_ir(data_ptr, indices_ptr, updates_ptr, out_ptr): return ib.get() - out_buf = tirx.decl_buffer(data.shape, data.dtype, "out_buf", layout=None) + out_buf = tirx.decl_tensor(data.shape, data.dtype, "out_buf", layout=None) return te.extern( [data.shape], [data, indices, updates], diff --git a/python/tvm/topi/scatter_elements.py b/python/tvm/topi/scatter_elements.py index f3cc09481e67..827be26b956c 100644 --- a/python/tvm/topi/scatter_elements.py +++ b/python/tvm/topi/scatter_elements.py @@ -182,7 +182,7 @@ def max_func(dst_ptr, dst_index, update): "scatter_elements reduction not in [update, add, mul, mean, min, max]:", reduction ) - out_buf = tirx.decl_buffer(data.shape, data.dtype, "out_buf", layout=None) + out_buf = tirx.decl_tensor(data.shape, data.dtype, "out_buf", layout=None) return te.extern( [data.shape], [data, indices, updates], diff --git a/python/tvm/topi/searchsorted.py b/python/tvm/topi/searchsorted.py index d4d02634639a..05ace4eefd15 100644 --- a/python/tvm/topi/searchsorted.py +++ b/python/tvm/topi/searchsorted.py @@ -38,8 +38,8 @@ def binary_search(sequence_offset, search_range, sorted_sequence, value, right, Note that we index N-D Buffer by 1-D linearlized indices. """ - lo_buf = T.decl_buffer([1], out_dtype, scope="local") - hi_buf = T.decl_buffer([1], out_dtype, scope="local") + lo_buf = T.decl_tensor([1], out_dtype, scope="local") + hi_buf = T.decl_tensor([1], out_dtype, scope="local") lo = lo_buf hi = hi_buf diff --git a/python/tvm/topi/signal.py b/python/tvm/topi/signal.py index ba76b3894676..7cd3de87acd5 100644 --- a/python/tvm/topi/signal.py +++ b/python/tvm/topi/signal.py @@ -138,7 +138,7 @@ def gen_ir( return ib.get() - output_buf = tirx.decl_buffer(output_shape, data.dtype, "output_buf", layout=None) + output_buf = tirx.decl_tensor(output_shape, data.dtype, "output_buf", layout=None) loop_kind = "vectorize" if ir.is_prim_var(output_shape[2]): # any_dim loop_kind = "serial" diff --git a/python/tvm/topi/sort.py b/python/tvm/topi/sort.py index 2f51f58f29e7..d9b7dbc56d58 100644 --- a/python/tvm/topi/sort.py +++ b/python/tvm/topi/sort.py @@ -48,10 +48,10 @@ def sort(data, axis=-1, is_ascend=1): Sorted index tensor. """ - data_buf = tvm.tirx.decl_buffer( + data_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "data_buf", data_alignment=8, layout=None ) - out_buf = tvm.tirx.decl_buffer(data.shape, data.dtype, "out_buf", data_alignment=8, layout=None) + out_buf = tvm.tirx.decl_tensor(data.shape, data.dtype, "out_buf", data_alignment=8, layout=None) out = te.extern( data.shape, [data], @@ -113,14 +113,14 @@ def argsort(data, valid_count=None, axis=-1, is_ascend=1, dtype="float32"): tvm_out = tvm.runtime.tensor(np.zeros(dshape, dtype=data.dtype.dtype), dev) f(tvm_data, tvm_out) """ - data_buf = tvm.tirx.decl_buffer( + data_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "data_buf", data_alignment=8, layout=None ) if valid_count is not None: - valid_count_buf = tvm.tirx.decl_buffer( + valid_count_buf = tvm.tirx.decl_tensor( valid_count.shape, valid_count.dtype, "valid_count_buf", data_alignment=4, layout=None ) - out_buf = tvm.tirx.decl_buffer( + out_buf = tvm.tirx.decl_tensor( data.shape, "int32", "out_buf", data_alignment=8, layout=None ) out = te.extern( @@ -136,7 +136,7 @@ def argsort(data, valid_count=None, axis=-1, is_ascend=1, dtype="float32"): tag="argsort_nms_cpu", ) else: - out_buf = tvm.tirx.decl_buffer(data.shape, dtype, "out_buf", data_alignment=8, layout=None) + out_buf = tvm.tirx.decl_tensor(data.shape, dtype, "out_buf", data_alignment=8, layout=None) out = te.extern( data.shape, [data], @@ -184,7 +184,7 @@ def topk(data, k=1, axis=-1, ret_type="both", is_ascend=False, dtype="int64"): The computed result. """ assert ret_type in ["both", "values", "indices"] - data_buf = tvm.tirx.decl_buffer( + data_buf = tvm.tirx.decl_tensor( data.shape, data.dtype, "data_buf", data_alignment=8, layout=None ) out_shape = list(get_const_tuple(data.shape)) @@ -196,11 +196,11 @@ def topk(data, k=1, axis=-1, ret_type="both", is_ascend=False, dtype="int64"): out_bufs = [] if ret_type in ["both", "values"]: out_bufs.append( - tvm.tirx.decl_buffer(out_shape, data.dtype, "value_buf", data_alignment=8, layout=None) + tvm.tirx.decl_tensor(out_shape, data.dtype, "value_buf", data_alignment=8, layout=None) ) if ret_type in ["both", "indices"]: out_bufs.append( - tvm.tirx.decl_buffer(out_shape, dtype, "indices_buf", data_alignment=8, layout=None) + tvm.tirx.decl_tensor(out_shape, dtype, "indices_buf", data_alignment=8, layout=None) ) out_shapes = [out_shape] * len(out_bufs) diff --git a/python/tvm/topi/sparse_reshape.py b/python/tvm/topi/sparse_reshape.py index 43f2cb92b289..9e98581f43c4 100644 --- a/python/tvm/topi/sparse_reshape.py +++ b/python/tvm/topi/sparse_reshape.py @@ -19,7 +19,7 @@ from tvm.script.ir_builder import IRBuilder from tvm.te import div, extern, floordiv, floormod -from tvm.tirx import Cast, decl_buffer +from tvm.tirx import Cast, decl_tensor from tvm.tirx.script import ir_builder as T @@ -90,21 +90,21 @@ def gen_ir( prev_shape_size = prev_shape_ptr.shape[0] new_shape_size = new_shape_ptr.shape[0] - multipliers_buf = T.alloc_buffer( + multipliers_buf = T.alloc_tensor( [prev_shape_size], new_shape_ptr.dtype, scope="local" ) multipliers = multipliers_buf - dividers_buf = T.alloc_buffer([new_shape_size], new_shape_ptr.dtype, scope="local") + dividers_buf = T.alloc_tensor([new_shape_size], new_shape_ptr.dtype, scope="local") dividers = dividers_buf - flattened_indices_buf = T.alloc_buffer( + flattened_indices_buf = T.alloc_tensor( [sparse_indices_ptr.shape[0]], new_shape_ptr.dtype, scope="local" ) flattened_indices = flattened_indices_buf - total_ele_buf = T.alloc_buffer([1], new_shape_ptr.dtype, scope="local") + total_ele_buf = T.alloc_tensor([1], new_shape_ptr.dtype, scope="local") total_ele = total_ele_buf - division_total_ele_buf = T.alloc_buffer([1], new_shape_ptr.dtype, scope="local") + division_total_ele_buf = T.alloc_tensor([1], new_shape_ptr.dtype, scope="local") division_total_ele = division_total_ele_buf - equal_shape_buf = T.alloc_buffer([1], "bool", scope="local") + equal_shape_buf = T.alloc_tensor([1], "bool", scope="local") equal_shape = equal_shape_buf T.buffer_store( @@ -234,7 +234,7 @@ def gen_ir( ) with T.parallel(0, new_sparse_indices_ptr.shape[0]) as i: - current_element_buf = T.alloc_buffer( + current_element_buf = T.alloc_tensor( [1], new_shape_ptr.dtype, scope="local" ) current_element = current_element_buf @@ -267,10 +267,10 @@ def gen_ir( return ib.get() - new_sparse_indices_buf = decl_buffer( + new_sparse_indices_buf = decl_tensor( new_sparse_indices_shape, sparse_indices.dtype, "new_sparse_indices_buf" ) - new_shape_buf = decl_buffer(new_shape_shape, prev_shape.dtype, "new_shape_buf") + new_shape_buf = decl_tensor(new_shape_shape, prev_shape.dtype, "new_shape_buf") return extern( [new_sparse_indices_shape, new_shape_shape], diff --git a/python/tvm/topi/vision/nms.py b/python/tvm/topi/vision/nms.py index 3394fc9b11a9..a33301c7d6a3 100644 --- a/python/tvm/topi/vision/nms.py +++ b/python/tvm/topi/vision/nms.py @@ -122,16 +122,16 @@ def get_valid_counts(data, score_threshold=0, id_index=0, score_index=1): id_index_const = tvm.tirx.const(id_index, "int32") score_index_const = tvm.tirx.const(score_index, "int32") - valid_count_buf = tvm.tirx.decl_buffer((batch_size,), "int32", "valid_count", layout=None) - out_tensor_buf = tvm.tirx.decl_buffer( + valid_count_buf = tvm.tirx.decl_tensor((batch_size,), "int32", "valid_count", layout=None) + out_tensor_buf = tvm.tirx.decl_tensor( (batch_size, num_anchors, box_data_length), data.dtype, "out_tensor", layout=None ) - out_indices_buf = tvm.tirx.decl_buffer( + out_indices_buf = tvm.tirx.decl_tensor( (batch_size, num_anchors), "int32", "out_indices", layout=None ) if is_score_threshold_tensor: - score_thresh_buf = tvm.tirx.decl_buffer( + score_thresh_buf = tvm.tirx.decl_tensor( score_threshold.shape, score_threshold.dtype, "score_threshold", layout=None ) valid_count, out_tensor, out_indices = te.extern( @@ -149,7 +149,7 @@ def get_valid_counts(data, score_threshold=0, id_index=0, score_index=1): dtype=["int32", data.dtype, "int32"], out_buffers=[valid_count_buf, out_tensor_buf, out_indices_buf], in_buffers=[ - tvm.tirx.decl_buffer(data.shape, data.dtype, "data", layout=None), + tvm.tirx.decl_tensor(data.shape, data.dtype, "data", layout=None), score_thresh_buf, ], name="get_valid_counts", @@ -174,7 +174,7 @@ def _ir_with_const_threshold(ins, outs): _ir_with_const_threshold, dtype=["int32", data.dtype, "int32"], out_buffers=[valid_count_buf, out_tensor_buf, out_indices_buf], - in_buffers=[tvm.tirx.decl_buffer(data.shape, data.dtype, "data", layout=None)], + in_buffers=[tvm.tirx.decl_tensor(data.shape, data.dtype, "data", layout=None)], name="get_valid_counts", tag="get_valid_counts", ) @@ -212,7 +212,7 @@ def _classic_nms_ir( with T.parallel(0, batch_size) as i: # Step 1: Reorder data by sorted score - nkeep_buf = T.alloc_buffer((1,), "int32", scope="local") + nkeep_buf = T.alloc_tensor((1,), "int32", scope="local") nkeep_local = nkeep_buf T.buffer_store( nkeep_local, @@ -251,16 +251,16 @@ def _classic_nms_ir( T.buffer_store(out_box_indices, T.int32(-1), (i, j)) # Step 2: Apply NMS - greedy suppression - num_valid_boxes_buf = T.alloc_buffer((1,), "int32", scope="local") + num_valid_boxes_buf = T.alloc_tensor((1,), "int32", scope="local") num_valid_boxes = num_valid_boxes_buf T.buffer_store(num_valid_boxes, T.int32(0), T.buffer_indices(num_valid_boxes, 0)) - best_idx_buf = T.alloc_buffer((1,), "int32", scope="local") + best_idx_buf = T.alloc_tensor((1,), "int32", scope="local") best_idx = best_idx_buf - best_score_buf = T.alloc_buffer((1,), data.dtype, scope="local") + best_score_buf = T.alloc_tensor((1,), data.dtype, scope="local") best_score = best_score_buf - tmp_idx_buf = T.alloc_buffer((1,), "int32", scope="local") + tmp_idx_buf = T.alloc_tensor((1,), "int32", scope="local") tmp_idx = tmp_idx_buf - tmp_val_buf = T.alloc_buffer((1,), data.dtype, scope="local") + tmp_val_buf = T.alloc_tensor((1,), data.dtype, scope="local") tmp_val = tmp_val_buf zero = tvm.tirx.Cast(data.dtype, T.float32(0.0)) @@ -614,7 +614,7 @@ def compute_iou(lhs_idx, rhs_idx): ) if return_indices: - valid_idx_buf = T.alloc_buffer((1,), "int32", scope="local") + valid_idx_buf = T.alloc_tensor((1,), "int32", scope="local") valid_idx = valid_idx_buf T.buffer_store(valid_idx, T.int32(0), T.buffer_indices(valid_idx, 0)) @@ -746,22 +746,22 @@ def non_max_suppression( ) sort_tensor = argsort(score_tensor, valid_count=valid_count, axis=1, is_ascend=False) - data_buf = tvm.tirx.decl_buffer(data.shape, data.dtype, "data", layout=None) - sort_buf = tvm.tirx.decl_buffer( + data_buf = tvm.tirx.decl_tensor(data.shape, data.dtype, "data", layout=None) + sort_buf = tvm.tirx.decl_tensor( sort_tensor.shape, sort_tensor.dtype, "sorted_index", layout=None ) - valid_count_buf = tvm.tirx.decl_buffer( + valid_count_buf = tvm.tirx.decl_tensor( valid_count.shape, valid_count.dtype, "valid_count", layout=None ) - indices_buf = tvm.tirx.decl_buffer(indices.shape, indices.dtype, "indices", layout=None) + indices_buf = tvm.tirx.decl_tensor(indices.shape, indices.dtype, "indices", layout=None) - out_data_buf = tvm.tirx.decl_buffer(data.shape, data.dtype, "out_data", layout=None) - out_box_indices_buf = tvm.tirx.decl_buffer( + out_data_buf = tvm.tirx.decl_tensor(data.shape, data.dtype, "out_data", layout=None) + out_box_indices_buf = tvm.tirx.decl_tensor( (batch_size, num_anchors), "int32", "out_box_indices", layout=None ) if return_indices: - out_valid_box_count_buf = tvm.tirx.decl_buffer( + out_valid_box_count_buf = tvm.tirx.decl_tensor( (batch_size, 1), "int32", "out_valid_box_count", layout=None ) @@ -841,7 +841,7 @@ def non_max_suppression( def _rearrange_out(data, batch_size, num_anchors, box_data_length, score_index): """Move valid boxes (score >= 0) to the top of output.""" - out_buf = tvm.tirx.decl_buffer( + out_buf = tvm.tirx.decl_tensor( (batch_size, num_anchors, box_data_length), data.dtype, "rearranged", layout=None ) @@ -851,7 +851,7 @@ def _rearrange_ir(ins, outs): out = outs[0] with T.parallel(0, batch_size) as i: - valid_idx_buf = T.alloc_buffer((1,), "int32", scope="local") + valid_idx_buf = T.alloc_tensor((1,), "int32", scope="local") valid_idx = valid_idx_buf T.buffer_store(valid_idx, T.int32(0), T.buffer_indices(valid_idx, 0)) @@ -946,7 +946,7 @@ def nms_inner_loop(i, j, nkeep, num_valid_boxes_local): with T.if_(tvm.tirx.all(iou_threshold > te.const(0), valid_count[i] > te.const(0))): with T.then_(): - num_valid_boxes_local_buf = T.alloc_buffer((1,), "int32", scope="local") + num_valid_boxes_local_buf = T.alloc_tensor((1,), "int32", scope="local") num_valid_boxes_local = num_valid_boxes_local_buf T.buffer_store( num_valid_boxes_local, T.int32(0), T.buffer_indices(num_valid_boxes_local, 0) @@ -997,15 +997,15 @@ def searchsorted_ir(scores_buf, score_thresh_buf, valid_count_buf): return ib.get() - scores_buf = tvm.tirx.decl_buffer( + scores_buf = tvm.tirx.decl_tensor( scores.shape, scores.dtype, "scores_buf", data_alignment=8, layout=None ) - searchsorted_buf = tvm.tirx.decl_buffer( + searchsorted_buf = tvm.tirx.decl_tensor( (batch_classes,), "int32", "searchsorted", data_alignment=8, layout=None ) if hasattr(score_threshold, "shape"): - score_thresh_buf = tvm.tirx.decl_buffer( + score_thresh_buf = tvm.tirx.decl_tensor( score_threshold.shape, score_threshold.dtype, "score_thresh_buf", diff --git a/python/tvm/topi/vision/nms_util.py b/python/tvm/topi/vision/nms_util.py index 369382d039a3..c88b2d27ee4e 100644 --- a/python/tvm/topi/vision/nms_util.py +++ b/python/tvm/topi/vision/nms_util.py @@ -67,8 +67,8 @@ def binary_search(y, num_boxes, scores, score_threshold, out): Must be called within an IRBuilder context. """ - lo_buf = T.decl_buffer([1], "int32", scope="local") - hi_buf = T.decl_buffer([1], "int32", scope="local") + lo_buf = T.decl_tensor([1], "int32", scope="local") + hi_buf = T.decl_tensor([1], "int32", scope="local") lo = lo_buf hi = hi_buf T.buffer_store(lo, T.int32(0), T.buffer_indices(lo, 0)) @@ -423,10 +423,10 @@ def run_all_class_nms( num_class = batch_class // batch if return_scores is False: - all_class_num0_buf = tvm.tirx.decl_buffer( + all_class_num0_buf = tvm.tirx.decl_tensor( (batch_class, num_boxes), "int32", "all_class_nms0", data_alignment=8, layout=None ) - all_class_num1_buf = tvm.tirx.decl_buffer( + all_class_num1_buf = tvm.tirx.decl_tensor( (batch_class,), "int32", "all_class_nms1", data_alignment=8, layout=None ) extern_inputs = [boxes, sorted_scores, sorted_indices, valid_count] diff --git a/src/backend/cuda/codegen/codegen_cuda.cc b/src/backend/cuda/codegen/codegen_cuda.cc index 8945d5bce319..9beb526c180b 100644 --- a/src/backend/cuda/codegen/codegen_cuda.cc +++ b/src/backend/cuda/codegen/codegen_cuda.cc @@ -1394,7 +1394,7 @@ void CodeGenCUDA::Dispatch_(const CallNode* op, std::ostream& os) { if (const auto* call = arg.as(); call && call->op.same_as(tirx::builtin::buffer_data()) && call->args.size() == 1) { var_node = call->args[0].as(); - TVM_FFI_ICHECK(var_node && var_node->ty.as()) + TVM_FFI_ICHECK(var_node && var_node->ty.as()) << "print_buffer expects buffer_data to project a BufferVar"; } PrimType dtype_ty = op->ty.as_or_throw(); @@ -1585,12 +1585,12 @@ void CodeGenCUDA::Dispatch_(const AttrStmtNode* op) { void CodeGenCUDA::Dispatch_(const BindNode* op) { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(tirx::builtin::alloc_buffer())) return DispatchAllocBuffer(op, call); + if (call->op.same_as(tirx::builtin::alloc_tensor())) return DispatchAllocTensor(op, call); } CodeGenC::Dispatch_(op); } -void CodeGenCUDA::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenCUDA::DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call) { tvm::Tuple shape = buffer_call->args[0].as_or_throw(); DLDataType dtype = buffer_call->args[1].as_or_throw()->value; ffi::String scope = buffer_call->args[2].as_or_throw()->value; diff --git a/src/backend/cuda/codegen/codegen_cuda.h b/src/backend/cuda/codegen/codegen_cuda.h index d383a67426bc..06c9160e1b5c 100644 --- a/src/backend/cuda/codegen/codegen_cuda.h +++ b/src/backend/cuda/codegen/codegen_cuda.h @@ -79,7 +79,7 @@ class CodeGenCUDA final : public CodeGenC { void Dispatch_(const EvaluateNode* op) final; void Dispatch_(const ReturnNode* op) final; void Dispatch_(const BindNode* op) final; - void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call); + void DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call); void Dispatch_(const AttrStmtNode* op) final; // Target diff --git a/src/backend/cuda/codegen/llvm/codegen_nvptx.cc b/src/backend/cuda/codegen/llvm/codegen_nvptx.cc index 5ab7d0f912a0..5a05b3977859 100644 --- a/src/backend/cuda/codegen/llvm/codegen_nvptx.cc +++ b/src/backend/cuda/codegen/llvm/codegen_nvptx.cc @@ -80,13 +80,13 @@ class CodeGenNVPTX : public CodeGenLLVM { void Dispatch_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return DispatchAllocBuffer(op, call); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return DispatchAllocTensor(op, call); } CodeGenLLVM::Dispatch_(op); } - void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) { + void DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call) { tvm::Tuple shape = buffer_call->args[0].as_or_throw(); DLDataType dtype = buffer_call->args[1].as_or_throw()->value; ffi::String scope = buffer_call->args[2].as_or_throw()->value; diff --git a/src/backend/cuda/transforms/lower_iket.cc b/src/backend/cuda/transforms/lower_iket.cc index 55e2eec39737..87d645b4bb42 100644 --- a/src/backend/cuda/transforms/lower_iket.cc +++ b/src/backend/cuda/transforms/lower_iket.cc @@ -597,7 +597,7 @@ class StripIket : public StmtExprMutator { private: UnchangedOr Mutate_(const BindNode* alloc, InplaceMode inplace_mode) final { if (const auto* call = alloc->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer()) && + call && call->op.same_as(tirx::builtin::alloc_tensor()) && token_buffers_.count(alloc->var.get())) { return Evaluate(0); } diff --git a/src/backend/metal/codegen/codegen_metal.cc b/src/backend/metal/codegen/codegen_metal.cc index 565290d04a44..592264346640 100644 --- a/src/backend/metal/codegen/codegen_metal.cc +++ b/src/backend/metal/codegen/codegen_metal.cc @@ -55,7 +55,7 @@ Var GetSimdgroupBufferVar(const Expr& data) { if (const auto* call = data.as(); call && call->op.same_as(tirx::builtin::buffer_data()) && call->args.size() == 1) { const auto* buffer = call->args[0].as(); - TVM_FFI_ICHECK(buffer && buffer->ty.as()) + TVM_FFI_ICHECK(buffer && buffer->ty.as()) << "Metal simdgroup data operands expect buffer_data to project a BufferVar"; return ffi::GetRef(buffer); } @@ -333,8 +333,8 @@ void CodeGenMetal::PrintStorageScope(const std::string& scope, std::ostream& os) void CodeGenMetal::Dispatch_(const BindNode* op) { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(tirx::builtin::alloc_buffer())) return DispatchAllocBuffer(op, call); - if (call->op.same_as(tirx::builtin::decl_buffer())) return DispatchDeclBuffer(op, call); + if (call->op.same_as(tirx::builtin::alloc_tensor())) return DispatchAllocTensor(op, call); + if (call->op.same_as(tirx::builtin::decl_tensor())) return DispatchDeclTensor(op, call); } // Stateful reads cannot be substituted after the underlying state changes. if (auto prim_value = op->value.as(); @@ -365,7 +365,7 @@ void CodeGenMetal::Dispatch_(const BindNode* op) { stream << "*)" << value << ";\n"; } -void CodeGenMetal::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenMetal::DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call) { tvm::Tuple shape = buffer_call->args[0].as_or_throw(); DLDataType dtype = buffer_call->args[1].as_or_throw()->value; ffi::String scope = buffer_call->args[2].as_or_throw()->value; diff --git a/src/backend/metal/codegen/codegen_metal.h b/src/backend/metal/codegen/codegen_metal.h index d0724eb6bfd5..2a84f9900856 100644 --- a/src/backend/metal/codegen/codegen_metal.h +++ b/src/backend/metal/codegen/codegen_metal.h @@ -54,7 +54,7 @@ class CodeGenMetal final : public CodeGenC { const std::string& value) final; // overload visitor void Dispatch_(const BindNode* op) final; // NOLINT(*) - void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call); + void DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call); void Dispatch_(const prim::SelectNode* op, std::ostream& os) final; // NOLINT(*) void Dispatch_(const prim::BroadcastNode* op, std::ostream& os) final; // NOLINT(*) void Dispatch_(const CallNode* op, std::ostream& os) final; // NOLINT(*) diff --git a/src/backend/opencl/codegen/codegen_opencl.cc b/src/backend/opencl/codegen/codegen_opencl.cc index 76461f6d9c17..f20eb6d7e53b 100644 --- a/src/backend/opencl/codegen/codegen_opencl.cc +++ b/src/backend/opencl/codegen/codegen_opencl.cc @@ -47,8 +47,8 @@ const VarNode* TryUnwrapTextureVar(const Expr& texture) { call && call->op.same_as(tirx::builtin::buffer_data())) { TVM_FFI_ICHECK_EQ(call->args.size(), 1U); const auto* buffer = call->args[0].as(); - TVM_FFI_ICHECK(buffer && buffer->ty.as()) - << "buffer_data expects a Var with BufferType"; + TVM_FFI_ICHECK(buffer && buffer->ty.as()) + << "buffer_data expects a Var with TensorType"; return buffer; } return nullptr; @@ -94,7 +94,7 @@ class InferTextureAccess : public StmtExprVisitor { } ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::decl_buffer())) { + call && call->op.same_as(tirx::builtin::decl_tensor())) { if (const VarNode* source = TryUnwrapTextureVar(call->args[0])) { auto it = buffer_data_map_.find(source); buffer_data_map_[op->var.get()] = it == buffer_data_map_.end() ? source : it->second; @@ -450,12 +450,12 @@ std::string CodeGenOpenCL::CastTo(std::string value, const PrimType& target) { void CodeGenOpenCL::Dispatch_(const BindNode* op) { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(tirx::builtin::alloc_buffer())) return DispatchAllocBuffer(op, call); + if (call->op.same_as(tirx::builtin::alloc_tensor())) return DispatchAllocTensor(op, call); } CodeGenC::Dispatch_(op); } -void CodeGenOpenCL::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenOpenCL::DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call) { tvm::Tuple shape = buffer_call->args[0].as_or_throw(); DLDataType dtype = buffer_call->args[1].as_or_throw()->value; BufferVar buffer = op->var.as_or_throw(); @@ -467,7 +467,7 @@ void CodeGenOpenCL::DispatchAllocBuffer(const BindNode* op, const CallNode* buff constant_size *= dim_imm->value.as().value(); } allocation_size_.insert({buffer.get(), constant_size * PrimType(dtype).lanes()}); - CodeGenC::DispatchAllocBuffer(op, buffer_call); + CodeGenC::DispatchAllocTensor(op, buffer_call); } void CodeGenOpenCL::Dispatch_(const CallNode* op, std::ostream& os) { diff --git a/src/backend/opencl/codegen/codegen_opencl.h b/src/backend/opencl/codegen/codegen_opencl.h index 18a507c28678..97c19463dc36 100644 --- a/src/backend/opencl/codegen/codegen_opencl.h +++ b/src/backend/opencl/codegen/codegen_opencl.h @@ -65,7 +65,7 @@ class CodeGenOpenCL final : public CodeGenC { // overload visitor void Dispatch_(const BindNode* op) final; - void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call); + void DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call); void Dispatch_(const prim::BroadcastNode* op, std::ostream& os) final; // NOLINT(*) void Dispatch_(const prim::RampNode* op, std::ostream& os) final; // NOLINT(*) void Dispatch_(const CallNode* op, std::ostream& os) final; // NOLINT(*) diff --git a/src/backend/rocm/codegen/llvm/codegen_amdgpu.cc b/src/backend/rocm/codegen/llvm/codegen_amdgpu.cc index 9d03861ab9ef..5856f60e4464 100644 --- a/src/backend/rocm/codegen/llvm/codegen_amdgpu.cc +++ b/src/backend/rocm/codegen/llvm/codegen_amdgpu.cc @@ -98,13 +98,13 @@ class CodeGenAMDGPU : public CodeGenLLVM { void Dispatch_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return DispatchAllocBuffer(op, call); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return DispatchAllocTensor(op, call); } CodeGenLLVM::Dispatch_(op); } - void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) { + void DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call) { tvm::Tuple shape = buffer_call->args[0].as_or_throw(); DLDataType dtype = buffer_call->args[1].as_or_throw()->value; ffi::String scope = buffer_call->args[2].as_or_throw()->value; diff --git a/src/backend/trn/codegen/codegen_trn.cc b/src/backend/trn/codegen/codegen_trn.cc index a2078498c9f9..6ceef277e263 100644 --- a/src/backend/trn/codegen/codegen_trn.cc +++ b/src/backend/trn/codegen/codegen_trn.cc @@ -216,13 +216,13 @@ std::string CodeGenTrainium::GetStorageScopeStr(const std::string& scope) { // void CodeGenTrainium::Dispatch_(const BindNode* op) { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(tirx::builtin::alloc_buffer())) return DispatchAllocBuffer(op, call); - if (call->op.same_as(tirx::builtin::decl_buffer())) return DispatchDeclBuffer(op, call); + if (call->op.same_as(tirx::builtin::alloc_tensor())) return DispatchAllocTensor(op, call); + if (call->op.same_as(tirx::builtin::decl_tensor())) return DispatchDeclTensor(op, call); } CodeGenC::Dispatch_(op); } -void CodeGenTrainium::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenTrainium::DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call) { tvm::Tuple shape = buffer_call->args[0].as_or_throw(); DLDataType dtype = buffer_call->args[1].as_or_throw()->value; ffi::String scope = buffer_call->args[2].as_or_throw()->value; @@ -627,7 +627,7 @@ void CodeGenTrainium::Dispatch_(const prim::FloorModNode* op, std::ostream& os) os << PrintExpr(op->a) << " % " << PrintExpr(op->b); } -void CodeGenTrainium::DispatchDeclBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenTrainium::DispatchDeclTensor(const BindNode* op, const CallNode* buffer_call) { Expr data = buffer_call->args[0]; tvm::Tuple shape = buffer_call->args[1].as_or_throw(); DLDataType dtype = buffer_call->args[2].as_or_throw()->value; @@ -641,13 +641,13 @@ void CodeGenTrainium::DispatchDeclBuffer(const BindNode* op, const CallNode* buf call && call->op.same_as(tirx::builtin::buffer_data()) && call->args.size() == 1) { data_var = call->args[0].as(); } - TVM_FFI_ICHECK(data_var) << "Trainium codegen expects DeclBuffer data to be a buffer variable"; + TVM_FFI_ICHECK(data_var) << "Trainium codegen expects DeclTensor data to be a buffer variable"; if (data_var->ty.as()) { buffer_idmap_[buffer] = GetVarID(data_var); buffer_data_varmap_[buffer] = data_var; return; } - TVM_FFI_ICHECK(data_var->ty.as()); + TVM_FFI_ICHECK(data_var->ty.as()); BufferVar source_buffer = ffi::GetRef(data_var).as_or_throw(); auto source_it = buffer_data_varmap_.find(source_buffer); TVM_FFI_ICHECK(source_it != buffer_data_varmap_.end()) diff --git a/src/backend/trn/codegen/codegen_trn.h b/src/backend/trn/codegen/codegen_trn.h index 422cc4258024..196b560bc527 100644 --- a/src/backend/trn/codegen/codegen_trn.h +++ b/src/backend/trn/codegen/codegen_trn.h @@ -59,7 +59,7 @@ class CodeGenTrainium final : public CodeGenC { void Dispatch_(const VarNode* op, std::ostream& os) final; // NOLINT(*) void PrintType(const PrimType& t, std::ostream& os) final; // NOLINT(*) void Dispatch_(const BindNode* op) final; - void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call); + void DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call); void Dispatch_(const AttrStmtNode* op) final; // NOLINT(*) void Dispatch_(const ForNode* op) final; // NOLINT(*) void Dispatch_(const BufferStoreNode* op) final; // NOLINT(*)= @@ -71,7 +71,7 @@ class CodeGenTrainium final : public CodeGenC { void Dispatch_(const prim::CastNode* op, std::ostream& os) final; // NOLINT(*) void Dispatch_(const prim::FloorDivNode* op, std::ostream& os) final; // NOLINT(*) void Dispatch_(const prim::FloorModNode* op, std::ostream& os) final; // NOLINT(*) - void DispatchDeclBuffer(const BindNode* op, const CallNode* buffer_call); + void DispatchDeclTensor(const BindNode* op, const CallNode* buffer_call); void Dispatch_(const IfThenElseNode* op) final; // NOLINT(*) void Dispatch_(const prim::AndNode* op, std::ostream& os) final; // NOLINT(*) void Dispatch_(const prim::OrNode* op, std::ostream& os) final; // NOLINT(*) diff --git a/src/backend/trn/transform/lower_trainium_layout.cc b/src/backend/trn/transform/lower_trainium_layout.cc index 087fea2cb841..813d361e777a 100644 --- a/src/backend/trn/transform/lower_trainium_layout.cc +++ b/src/backend/trn/transform/lower_trainium_layout.cc @@ -68,7 +68,7 @@ class TrainiumLayoutApplier : public tirx::IRMutatorWithAnalyzer { } if (buffer.value()->layout.has_value()) { BufferVar flattened = storage_lower->GetFlattenedBuffer(buffer.value()); - auto type = CopyBufferType(buffer.value()); + auto type = CopyTensorType(buffer.value()); type->layout = std::nullopt; BufferVar source = RebuildBufferVar(buffer.value(), std::move(type)); param_flattened_buffers.emplace_back(flattened, source); @@ -80,7 +80,7 @@ class TrainiumLayoutApplier : public tirx::IRMutatorWithAnalyzer { auto new_stmt = storage_lower->Mutate(stmt, InplaceMode::kDisallow).ValueOrUnchanged(stmt); for (const auto& [buf, source] : param_flattened_buffers) { new_stmt = - SeqStmt::Flatten(Bind(buf, Call(buf.type(), tirx::builtin::decl_buffer(), + SeqStmt::Flatten(Bind(buf, Call(buf.type(), tirx::builtin::decl_tensor(), {source.data(), tvm::Tuple(buf->shape), DataTypeImm(buf->dtype->dtype), StringImm(buf.scope())}, {})), @@ -110,7 +110,7 @@ class TrainiumLayoutApplier : public tirx::IRMutatorWithAnalyzer { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { + call && call->op.same_as(tirx::builtin::alloc_tensor())) { BufferVar original_buffer = op->var.as_or_throw(); if (!original_buffer->layout.has_value()) { return ffi::Unchanged(); @@ -120,7 +120,7 @@ class TrainiumLayoutApplier : public tirx::IRMutatorWithAnalyzer { return ffi::Unchanged(); } return Bind(buffer.var(), - Call(buffer.type(), tirx::builtin::alloc_buffer(), + Call(buffer.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(buffer->shape, call->args[0]->span), DataTypeImm(buffer->dtype->dtype, call->args[1]->span), StringImm(buffer.scope(), call->args[2]->span)}, @@ -128,7 +128,7 @@ class TrainiumLayoutApplier : public tirx::IRMutatorWithAnalyzer { op->span); } if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::decl_buffer())) { + call && call->op.same_as(tirx::builtin::decl_tensor())) { BufferVar original_buffer = op->var.as_or_throw(); Expr original_data = call->args[0]; auto data_update = Mutate(original_data, inplace_mode); @@ -139,7 +139,7 @@ class TrainiumLayoutApplier : public tirx::IRMutatorWithAnalyzer { return ffi::Unchanged(); } return Bind(buffer, - Call(buffer.type(), tirx::builtin::decl_buffer(), + Call(buffer.type(), tirx::builtin::decl_tensor(), {std::move(data), tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())}, call->attrs, call->ty_args, call->span), @@ -155,7 +155,7 @@ class TrainiumLayoutApplier : public tirx::IRMutatorWithAnalyzer { } auto trn_layout = buf->layout.as(); BufferVar flattened; - ffi::ObjectPtr type; + ffi::ObjectPtr type; if (IsTrainiumLayout(trn_layout)) { ffi::Array new_shape = buf.scope() == "trn.psum" ? ffi::Array{trn_layout->GetSpan(ffi::String("Bank")), @@ -164,7 +164,7 @@ class TrainiumLayoutApplier : public tirx::IRMutatorWithAnalyzer { : ffi::Array{trn_layout->GetSize(ffi::String("P")), trn_layout->GetSpan(ffi::String("F"))}; flattened = buf; - type = CopyBufferType(flattened); + type = CopyTensorType(flattened); type->shape = new_shape; type->strides = {}; } else if (is_alloc) { @@ -188,16 +188,16 @@ class TrainiumLayoutApplier : public tirx::IRMutatorWithAnalyzer { } } flattened = buf; - type = CopyBufferType(flattened); + type = CopyTensorType(flattened); type->shape = {ana->Simplify(mem_span)}; type->strides = {}; } else { flattened = buf.GetFlattenedBuffer(); - type = CopyBufferType(flattened); + type = CopyTensorType(flattened); } } else { flattened = buf.GetFlattenedBuffer(); - type = CopyBufferType(flattened); + type = CopyTensorType(flattened); } if (flattened->dtype->dtype == DLDataType{kDLBool, 8, 1}) { type->dtype = PrimType::Int(8); diff --git a/src/backend/vulkan/codegen/codegen_spirv.cc b/src/backend/vulkan/codegen/codegen_spirv.cc index 0e4871567733..c41d4412fbbf 100644 --- a/src/backend/vulkan/codegen/codegen_spirv.cc +++ b/src/backend/vulkan/codegen/codegen_spirv.cc @@ -875,7 +875,7 @@ void CodeGenSPIRV::Dispatch_(const IfThenElseNode* op) { builder_->StartLabel(merge_label); } -void CodeGenSPIRV::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenSPIRV::DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call) { tvm::Tuple shape = buffer_call->args[0].as_or_throw(); DLDataType dtype = buffer_call->args[1].as_or_throw()->value; ffi::String scope = buffer_call->args[2].as_or_throw()->value; @@ -945,7 +945,7 @@ void CodeGenSPIRV::DispatchAllocBuffer(const BindNode* op, const CallNode* buffe } } -void CodeGenSPIRV::DispatchDeclBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenSPIRV::DispatchDeclTensor(const BindNode* op, const CallNode* buffer_call) { Expr data = buffer_call->args[0]; DLDataType dtype = buffer_call->args[2].as_or_throw()->value; BufferVar buffer = op->var.as_or_throw(); @@ -1002,8 +1002,8 @@ void CodeGenSPIRV::Dispatch_(const AssertStmtNode* op) { void CodeGenSPIRV::Dispatch_(const BindNode* op) { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(tirx::builtin::alloc_buffer())) return DispatchAllocBuffer(op, call); - if (call->op.same_as(tirx::builtin::decl_buffer())) return DispatchDeclBuffer(op, call); + if (call->op.same_as(tirx::builtin::alloc_tensor())) return DispatchAllocTensor(op, call); + if (call->op.same_as(tirx::builtin::decl_tensor())) return DispatchDeclTensor(op, call); } TVM_FFI_ICHECK(!var_map_.count(op->var.get())); if (auto prim_type = op->var->ty.as()) { diff --git a/src/backend/vulkan/codegen/codegen_spirv.h b/src/backend/vulkan/codegen/codegen_spirv.h index 831ccd128c68..31d2b015d0e9 100644 --- a/src/backend/vulkan/codegen/codegen_spirv.h +++ b/src/backend/vulkan/codegen/codegen_spirv.h @@ -118,8 +118,8 @@ class CodeGenSPIRV : public tirx::ExprFunctor, void Dispatch_(const ForNode* op) override; void Dispatch_(const WhileNode* op) override; void Dispatch_(const IfThenElseNode* op) override; - void DispatchDeclBuffer(const BindNode* op, const CallNode* buffer_call); - void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call); + void DispatchDeclTensor(const BindNode* op, const CallNode* buffer_call); + void DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call); void Dispatch_(const AttrStmtNode* op) override; void Dispatch_(const AssertStmtNode* op) override; void Dispatch_(const BindNode* op) override; @@ -141,7 +141,7 @@ class CodeGenSPIRV : public tirx::ExprFunctor, /*! \brief Whether the element type of the buffer is known. * * This value is determined based on the type_annotation of the - * buffer variable (alloc_buffer binding) or of the parameter (shader + * buffer variable (alloc_tensor binding) or of the parameter (shader * arguments). */ bool element_type_known{false}; @@ -149,7 +149,7 @@ class CodeGenSPIRV : public tirx::ExprFunctor, /*! \brief The known element type of the buffer. * * This value is determined based on the type_annotation of the - * buffer variable (alloc_buffer binding) or of the parameter (shader + * buffer variable (alloc_tensor binding) or of the parameter (shader * arguments). */ PrimType element_type{PrimType::Void()}; diff --git a/src/backend/webgpu/codegen/codegen_webgpu.cc b/src/backend/webgpu/codegen/codegen_webgpu.cc index 1a505f966fd6..5cce127f7b98 100644 --- a/src/backend/webgpu/codegen/codegen_webgpu.cc +++ b/src/backend/webgpu/codegen/codegen_webgpu.cc @@ -125,7 +125,7 @@ class WebGPUWorkgroupInfoCollector : public StmtExprVisitor { ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::decl_buffer())) { + call && call->op.same_as(tirx::builtin::decl_tensor())) { if (auto source = GetBufferDataVar(call->args[0])) { buffer_aliases_.insert_or_assign(op->var.get(), ResolveBuffer(source.value())); return std::nullopt; @@ -657,8 +657,8 @@ void CodeGenWebGPU::Dispatch_(const TensorLoadNode* op, std::ostream& os) { // void CodeGenWebGPU::Dispatch_(const BindNode* op) { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(tirx::builtin::alloc_buffer())) return DispatchAllocBuffer(op, call); - if (call->op.same_as(tirx::builtin::decl_buffer())) return DispatchDeclBuffer(op, call); + if (call->op.same_as(tirx::builtin::alloc_tensor())) return DispatchAllocTensor(op, call); + if (call->op.same_as(tirx::builtin::decl_tensor())) return DispatchDeclTensor(op, call); } // Stateful reads cannot be substituted after the underlying state changes. if (auto prim_value = op->value.as(); @@ -739,7 +739,7 @@ void CodeGenWebGPU::Dispatch_(const BufferStoreNode* op) { } } -void CodeGenWebGPU::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenWebGPU::DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call) { tvm::Tuple shape = buffer_call->args[0].as_or_throw(); DLDataType dtype = buffer_call->args[1].as_or_throw()->value; ffi::String scope = buffer_call->args[2].as_or_throw()->value; diff --git a/src/backend/webgpu/codegen/codegen_webgpu.h b/src/backend/webgpu/codegen/codegen_webgpu.h index ed662374796c..6501a9ad8fc8 100644 --- a/src/backend/webgpu/codegen/codegen_webgpu.h +++ b/src/backend/webgpu/codegen/codegen_webgpu.h @@ -82,7 +82,7 @@ class CodeGenWebGPU final : public CodeGenC { void Dispatch_(const BindNode* op) final; void Dispatch_(const BufferStoreNode* op) final; void Dispatch_(const ForNode* op) final; - void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call); + void DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call); void Dispatch_(const AssertStmtNode* op) final; void Dispatch_(const WhileNode* op) final; void Dispatch_(const BreakNode* op) final; diff --git a/src/ir/expr.cc b/src/ir/expr.cc index 6b883640e786..77a05553287d 100644 --- a/src/ir/expr.cc +++ b/src/ir/expr.cc @@ -436,7 +436,7 @@ TVM_FFI_INLINE ffi::Expected> VarVisit( // A PrimType carries only a dtype, so it has nothing to visit. Broad callbacks do not see this // skipped field; dynamically typed Vars still descend through the Type value. if (!self->ty.as()) { - // Only Simple is clamped: Pattern co-introduces type fields such as BufferType shape + // Only Simple is clamped: Pattern co-introduces type fields such as TensorType shape // variables, so that ambient region must continue through the dynamic type. if (visitor->def_region_kind() == kTVMFFIDefRegionKindSimple) { TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->WithDefRegionKind( @@ -467,7 +467,7 @@ TVM_FFI_INLINE ffi::Expected> VarMutate( // A PrimType carries only a dtype, so it has nothing to substitute. Broad callbacks do not see // this skipped field; dynamically typed Vars still descend through the Type value. if (!self->ty.as()) { - // Pattern co-introduces type fields such as BufferType shape variables; Simple does not. + // Pattern co-introduces type fields such as TensorType shape variables; Simple does not. ffi::Expected> mapped_ty_result = mutator->def_region_kind() == kTVMFFIDefRegionKindSimple ? mutator->WithDefRegionKind(kTVMFFIDefRegionKindNone, @@ -510,7 +510,7 @@ TVM_FFI_INLINE ffi::Expected> VarMaybeInplaceMutate( // A PrimType carries only a dtype, so it has nothing to substitute. Broad callbacks do not see // this skipped field; dynamically typed Vars still descend through the Type value. if (!self->ty.as()) { - // Pattern co-introduces type fields such as BufferType shape variables; Simple does not. + // Pattern co-introduces type fields such as TensorType shape variables; Simple does not. ffi::Expected> mapped_ty_result = mutator->def_region_kind() == kTVMFFIDefRegionKindSimple ? mutator->WithDefRegionKind( diff --git a/src/relax/backend/vm/vm_shape_lower.cc b/src/relax/backend/vm/vm_shape_lower.cc index 47ce4ce39c23..f736514239c4 100644 --- a/src/relax/backend/vm/vm_shape_lower.cc +++ b/src/relax/backend/vm/vm_shape_lower.cc @@ -269,7 +269,7 @@ class PrimExprSlotCollector : public ExprVisitor, public TypeVisitor { * \code * * @T.prim_func - * def shape_func(H: T.Buffer([3], "int64")): + * def shape_func(H: T.Tensor([3], "int64")): * H[1] = H[2] + 1 * * \endcode @@ -715,7 +715,7 @@ class VMShapeLowerMutator TVM_FFI_ICHECK_GT(heap_size_->value, 0); // construct a PrimFunc that compute the shape. ffi::Array buffer_shape{heap_size_}; - tirx::BufferVar buffer = tirx::decl_buffer(buffer_shape, PrimType(ShapeDType()), "H", "global"); + tirx::BufferVar buffer = tirx::decl_tensor(buffer_shape, PrimType(ShapeDType()), "H", "global"); ffi::Map var_map; for (const auto& [expr, slot] : slot_map_) { diff --git a/src/relax/distributed/axis_group_graph.cc b/src/relax/distributed/axis_group_graph.cc index 80cc1d5b0690..afdc7a7e3dd2 100644 --- a/src/relax/distributed/axis_group_graph.cc +++ b/src/relax/distributed/axis_group_graph.cc @@ -359,7 +359,7 @@ void BuildAxisGraphCallTIR(const Var& output_var, const Call& call, const tirx:: ffi::Array input_list = call->args[1].as_or_throw()->fields; input_list.push_back(output_var); for (int i = 0; i < static_cast(input_list.size()); i++) { - if (func->params[i]->ty.as()) { + if (func->params[i]->ty.as()) { input_var_to_relax_expr.Set(func->params[i], input_list[i]); } } diff --git a/src/relax/distributed/transform/lower_global_view_to_local_view.cc b/src/relax/distributed/transform/lower_global_view_to_local_view.cc index e1f562f05d43..397a7bb5666c 100644 --- a/src/relax/distributed/transform/lower_global_view_to_local_view.cc +++ b/src/relax/distributed/transform/lower_global_view_to_local_view.cc @@ -132,7 +132,7 @@ class DistributedBufferCompactor : public s_tir::StmtExprMutator { ffi::Array new_params; ffi::Map replace_buffer_map; for (const Var& param : prim_func->params) { - if (!param->ty.as()) { + if (!param->ty.as()) { new_params.push_back(param); continue; } @@ -163,7 +163,7 @@ class DistributedBufferCompactor : public s_tir::StmtExprMutator { std::unordered_set visited; for (int i = 0, j = 0; i < static_cast(prim_func->params.size()); i++) { Var param_var = prim_func->params[i]; - if (!param_var->ty.as()) { + if (!param_var->ty.as()) { continue; } BufferVar param_buffer = param_var.as_or_throw(); @@ -250,7 +250,7 @@ class DistributedBufferCompactor : public s_tir::StmtExprMutator { shape.push_back(buffer->shape[i]); } } - BufferType new_type(buffer->storage_scope, buffer->dtype, std::move(shape), buffer->strides, + TensorType new_type(buffer->storage_scope, buffer->dtype, std::move(shape), buffer->strides, buffer->elem_offset, buffer->data_alignment, buffer->offset_factor, buffer->layout, buffer->allocated_addr); return BufferVar(buffer.name(), std::move(new_type), buffer.span()); @@ -396,7 +396,7 @@ class LowerTIRToLocalView : public ExprMutator { for (size_t i = 0; i < args.size(); ++i) { const Expr& arg = args[i]; const tirx::Var& param = prim_func->params[i]; - if (param->ty.as()) { + if (param->ty.as()) { const auto* ty = GetTypeAs(arg); TVM_FFI_CHECK(ty, TypeError) << "Expected buffer parameter " << param << " to receive a distributed tensor, but " diff --git a/src/relax/op/op.cc b/src/relax/op/op.cc index 5b4185fe90b9..9c15ce1dd44b 100644 --- a/src/relax/op/op.cc +++ b/src/relax/op/op.cc @@ -290,8 +290,8 @@ TVM_FFI_STATIC_INIT_BLOCK() { * * For dynamic shapes, it is not always possible to infer the output * of a TIR PrimFunc from its inputs. For example, a PrimFunc that - * accepts input buffer `T.Buffer([16], "float32")` and output buffer - * `T.Buffer([M, N], "float32")` infers the values of `M` and `N` from + * accepts input buffer `T.Tensor([16], "float32")` and output buffer + * `T.Tensor([M, N], "float32")` infers the values of `M` and `N` from * the shape of the provided output buffer. * * If the arguments provided are not compatible with the PrimFunc's diff --git a/src/relax/op/tensor/inspect.cc b/src/relax/op/tensor/inspect.cc index ca60f3e135a9..76888d663076 100644 --- a/src/relax/op/tensor/inspect.cc +++ b/src/relax/op/tensor/inspect.cc @@ -265,7 +265,7 @@ Expr LegalizeTensorShape(const BlockBuilder& bb, const Call& call) { tirx::Var ndim("ndim", PrimType::Int(32)); tirx::BufferVar shape_buffer = - tirx::decl_buffer({ndim.as_or_throw()}, field_ty, "shape"); + tirx::decl_tensor({ndim.as_or_throw()}, field_ty, "shape"); tirx::Var extent("extent", field_ty); @@ -284,7 +284,7 @@ Expr LegalizeTensorShape(const BlockBuilder& bb, const Call& call) { {StringImm("Specified axis may not be larger than the tensor's dimensionality")}), tirx::Bind( shape_buffer, - tvm::Call(shape_buffer.type(), tvm::tirx::builtin::decl_buffer(), + tvm::Call(shape_buffer.type(), tvm::tirx::builtin::decl_tensor(), {tvm::Call( shape_buffer.DataPointerType(), tirx::builtin::tvm_struct_get(), {dlpack_handle, IntImm::Int32(0), diff --git a/src/relax/transform/fuse_tir.cc b/src/relax/transform/fuse_tir.cc index 72eb56004da6..d6d23f9fea6e 100644 --- a/src/relax/transform/fuse_tir.cc +++ b/src/relax/transform/fuse_tir.cc @@ -794,7 +794,7 @@ class FusedTIRConstructor : public ExprVisitor { return unique_name; }; // Update buffer with new symbolic shape according to the ty - tirx::BufferType new_type(buffer->storage_scope, buffer->dtype, output_shapes[i], + tirx::TensorType new_type(buffer->storage_scope, buffer->dtype, output_shapes[i], buffer->strides, buffer->elem_offset, buffer->data_alignment, buffer->offset_factor, buffer->layout, buffer->allocated_addr); tirx::BufferVar new_buffer(unify_name_hints(), std::move(new_type), buffer.span()); @@ -833,10 +833,10 @@ class FusedTIRConstructor : public ExprVisitor { PrimType dtype = tensor->dtype.value(); tirx::BufferVar buffer; if (tir_buffer_param.has_value()) { - buffer = tirx::decl_buffer(shape_expr->values, dtype, name_hint, + buffer = tirx::decl_tensor(shape_expr->values, dtype, name_hint, tir_buffer_param.value().scope()); } else { - buffer = tirx::decl_buffer(shape_expr->values, dtype, name_hint); + buffer = tirx::decl_tensor(shape_expr->values, dtype, name_hint); } out->push_back(std::move(buffer)); diff --git a/src/relax/transform/specialize_primfunc_based_on_callsite.cc b/src/relax/transform/specialize_primfunc_based_on_callsite.cc index 81b36a44bc04..40c8136a24f8 100644 --- a/src/relax/transform/specialize_primfunc_based_on_callsite.cc +++ b/src/relax/transform/specialize_primfunc_based_on_callsite.cc @@ -100,7 +100,7 @@ class SpecializeTIRCallArgs : ExprMutator { name = std::string({static_cast('A' + i)}); } - const BufferVar& buffer = tirx::decl_buffer(GetShapeFromTensorType(tensor_ty), + const BufferVar& buffer = tirx::decl_tensor(GetShapeFromTensorType(tensor_ty), tensor_ty->dtype.value(), name, scope); param_map.Set(pfunc->params[i], buffer); } @@ -112,7 +112,7 @@ class SpecializeTIRCallArgs : ExprMutator { scope = ty->vdevice.value()->memory_scope; } const BufferVar& buffer = - tirx::decl_buffer(GetShapeFromTensorType(ty), ty->dtype.value(), "ret_val", scope); + tirx::decl_tensor(GetShapeFromTensorType(ty), ty->dtype.value(), "ret_val", scope); param_map.Set(pfunc->params[pfunc->params.size() - 1], buffer); } else { TVM_FFI_ICHECK(out_ty->IsInstance()) @@ -133,7 +133,7 @@ class SpecializeTIRCallArgs : ExprMutator { scope = ty->vdevice.value()->memory_scope; } - const BufferVar& buffer = tirx::decl_buffer(GetShapeFromTensorType(ty), ty->dtype.value(), + const BufferVar& buffer = tirx::decl_tensor(GetShapeFromTensorType(ty), ty->dtype.value(), "ret_val_" + std::to_string(index), scope); param_map.Set(pfunc->params[args.size() + index], buffer); index++; diff --git a/src/relax/transform/split_call_tir_by_pattern.cc b/src/relax/transform/split_call_tir_by_pattern.cc index 395669b19216..2956fb3b0864 100644 --- a/src/relax/transform/split_call_tir_by_pattern.cc +++ b/src/relax/transform/split_call_tir_by_pattern.cc @@ -411,7 +411,7 @@ class TIRPatternMatcher { ffi::Array pattern_symbolic_vars; int buffer_count = 0; while (buffer_count < static_cast(pattern_func->params.size()) && - pattern_func->params[buffer_count]->ty.as()) { + pattern_func->params[buffer_count]->ty.as()) { ++buffer_count; } for (int i = buffer_count; i < static_cast(pattern_func->params.size()); i++) { diff --git a/src/s_tir/analysis/calculate_allocated_memory.cc b/src/s_tir/analysis/calculate_allocated_memory.cc index ca52a59f4aed..303791d23e76 100644 --- a/src/s_tir/analysis/calculate_allocated_memory.cc +++ b/src/s_tir/analysis/calculate_allocated_memory.cc @@ -50,7 +50,7 @@ std::string GetStorageScope(const Var& var) { /*! * \brief Allocation calculator for buffer allocation bindings. */ -class AllocBufferCalculator : public StmtExprVisitor { +class AllocTensorCalculator : public StmtExprVisitor { public: using StmtExprVisitor::Visit_; @@ -66,13 +66,13 @@ class AllocBufferCalculator : public StmtExprVisitor { private: ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return DispatchAllocBuffer(op, call); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return DispatchAllocTensor(op, call); } return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchAllocBuffer(const BindNode* op, const CallNode* call) { + ffi::Optional DispatchAllocTensor(const BindNode* op, const CallNode* call) { tvm::Tuple shape = call->args[0].as_or_throw(); DLDataType dtype = call->args[1].as_or_throw()->value; ffi::String scope = call->args[2].as_or_throw()->value; @@ -120,7 +120,7 @@ class AllocBufferCalculator : public StmtExprVisitor { tvm::ffi::Map > CalculateAllocatedBytes( const PrimFunc& func) { tvm::ffi::Map > results; - auto alloc_buffer_result = ffi::make_object()->operator()(func); + auto alloc_buffer_result = ffi::make_object()->operator()(func); results.Set("main", alloc_buffer_result); return results; } @@ -132,7 +132,7 @@ tvm::ffi::Map > CalculateAlloca if (auto prim_func = kv.second.as()) { ffi::String func_name = kv.first->name_hint; auto alloc_buffer_result = - ffi::make_object()->operator()(prim_func.value()); + ffi::make_object()->operator()(prim_func.value()); results.Set(func_name, alloc_buffer_result); } } diff --git a/src/s_tir/analysis/domain_touched.cc b/src/s_tir/analysis/domain_touched.cc index cf45f7f521ae..727b15f54539 100644 --- a/src/s_tir/analysis/domain_touched.cc +++ b/src/s_tir/analysis/domain_touched.cc @@ -147,7 +147,7 @@ ffi::Map> DomainTouchedAccessMap(const Pri auto buffer_access_map = visitor->GetAccessedBufferRegions(); ffi::Map> ret; for (auto& var : func->params) { - if (!var->ty.as()) { + if (!var->ty.as()) { continue; } BufferVar buffer = var.as_or_throw(); diff --git a/src/s_tir/analysis/is_pure_function.cc b/src/s_tir/analysis/is_pure_function.cc index 05a8314a4aee..2cbc34e9a6f9 100644 --- a/src/s_tir/analysis/is_pure_function.cc +++ b/src/s_tir/analysis/is_pure_function.cc @@ -49,13 +49,13 @@ class PurityChecker : TIRVisitorWithPath { void Dispatch_(const BindNode* op, ffi::reflection::AccessPath path) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return DispatchAllocBuffer(op, call, path); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return DispatchAllocTensor(op, call, path); } return TIRVisitorWithPath::Dispatch_(op, path); } - void DispatchAllocBuffer(const BindNode* op, const CallNode* call, + void DispatchAllocTensor(const BindNode* op, const CallNode* call, ffi::reflection::AccessPath path) { internal_allocations_.insert(op->var); allocation_calls_.insert(call); diff --git a/src/s_tir/analysis/sblock_access_region_detector.cc b/src/s_tir/analysis/sblock_access_region_detector.cc index 8386e1e7e96b..90919950d8b1 100644 --- a/src/s_tir/analysis/sblock_access_region_detector.cc +++ b/src/s_tir/analysis/sblock_access_region_detector.cc @@ -235,8 +235,8 @@ ffi::Optional BlockReadWriteDetector::Visit_(const IfThenElseNod ffi::Optional BlockReadWriteDetector::Visit_(const BindNode* op) { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::decl_buffer())) { - // A DeclBuffer data expression defines the alias source. It is not an + call && call->op.same_as(tirx::builtin::decl_tensor())) { + // A DeclTensor data expression defines the alias source. It is not an // opaque buffer access by the containing block. return WithDefRegionKind(kTVMFFIDefRegionKindSimple, [&]() { return Visit(op->var); }); } diff --git a/src/s_tir/analysis/verify_gpu_code.cc b/src/s_tir/analysis/verify_gpu_code.cc index b3e0adf297db..30f31ebdda7b 100644 --- a/src/s_tir/analysis/verify_gpu_code.cc +++ b/src/s_tir/analysis/verify_gpu_code.cc @@ -70,13 +70,13 @@ class GPUCodeVerifier : public StmtExprVisitor { ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return DispatchAllocBuffer(op, call); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return DispatchAllocTensor(op, call); } return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchAllocBuffer(const BindNode* op, const CallNode* call) { + ffi::Optional DispatchAllocTensor(const BindNode* op, const CallNode* call) { TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(StmtExprVisitor::Visit_(op)); tvm::Tuple shape = call->args[0].as_or_throw(); DLDataType dtype = call->args[1].as_or_throw()->value; diff --git a/src/s_tir/backend/adreno/inject_texture_alloc.cc b/src/s_tir/backend/adreno/inject_texture_alloc.cc index cfb66295710d..78d051dbcf52 100644 --- a/src/s_tir/backend/adreno/inject_texture_alloc.cc +++ b/src/s_tir/backend/adreno/inject_texture_alloc.cc @@ -63,20 +63,20 @@ class TextureAllocInjector : public s_tir::IRMutatorWithAnalyzer { private: UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return Mutate_AllocBuffer(op, call, inplace_mode); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return Mutate_AllocTensor(op, call, inplace_mode); } return StmtExprMutator::Mutate_(op, inplace_mode); } - UnchangedOr Mutate_AllocBuffer(const BindNode* op, const CallNode* call, + UnchangedOr Mutate_AllocTensor(const BindNode* op, const CallNode* call, InplaceMode inplace_mode) { Stmt stmt = StmtExprMutator::Mutate_(op, inplace_mode).ValueOrUnchanged(ffi::GetRef(op)); ffi::String scope = call->args[2].as_or_throw()->value; if (IsTextureStorage(scope)) { op = stmt.as(); if (const auto* call = op ? op->value.as() : nullptr; - !call || !call->op.same_as(tirx::builtin::alloc_buffer())) { + !call || !call->op.same_as(tirx::builtin::alloc_tensor())) { TVM_FFI_THROW(InternalError) << "Expected an allocation binding after buffer mutation"; } const auto* allocation = op->value.as(); @@ -100,7 +100,7 @@ class TextureAllocInjector : public s_tir::IRMutatorWithAnalyzer { {texture.width, texture.height, texture.depth})); args.push_back(IntImm::Int64(channel_size)); stmt = Bind(op->var.as_or_throw(), - Call(op->var.as_or_throw().type(), tirx::builtin::decl_buffer(), + Call(op->var.as_or_throw().type(), tirx::builtin::decl_tensor(), {Call(op->var.as_or_throw().DataPointerType(), tirx::builtin::nd_mem_alloc_with_scope(), args), tvm::Tuple(op->var.as_or_throw()->shape), diff --git a/src/s_tir/ir/data_type_rewriter.cc b/src/s_tir/ir/data_type_rewriter.cc index 09d83aec6e7d..b3d857ee9529 100644 --- a/src/s_tir/ir/data_type_rewriter.cc +++ b/src/s_tir/ir/data_type_rewriter.cc @@ -168,7 +168,7 @@ ffi::Map IndexDataTypeNormalizer::VisitBlockAnnotations( if (obj == nullptr) { return obj; } - if (auto var = obj.as(); var && var.value()->ty.as()) { + if (auto var = obj.as(); var && var.value()->ty.as()) { BufferVar buffer = var.value().as_or_throw(); if (BufferVar new_buffer = this->Mutate(buffer, InplaceMode::kDisallow) .as_or_throw>() diff --git a/src/s_tir/meta_schedule/postproc/disallow_async_strided_mem_copy.cc b/src/s_tir/meta_schedule/postproc/disallow_async_strided_mem_copy.cc index 94310d55df26..86c4be96bee9 100644 --- a/src/s_tir/meta_schedule/postproc/disallow_async_strided_mem_copy.cc +++ b/src/s_tir/meta_schedule/postproc/disallow_async_strided_mem_copy.cc @@ -79,11 +79,11 @@ struct AsyncStridedMemCopyFinder : public StmtExprVisitor { } // get store buffer; assert it exists and is contiguous given it uses a single index - auto bufferstore = bufferstorenode->buffer.as(); + auto bufferstore = bufferstorenode->buffer.as(); // get load buffer; assert it exists and is contiguous given it uses a single index BufferVar load_buffer = bufferloadnode->source.as_or_throw(); - auto bufferload = load_buffer.as(); + auto bufferload = load_buffer.as(); if (!bufferstore || !bufferload) { TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(StmtExprVisitor::Visit_(attrStmt)); diff --git a/src/s_tir/schedule/analysis/analysis.cc b/src/s_tir/schedule/analysis/analysis.cc index 4114fda91f73..d0365d7ad81e 100644 --- a/src/s_tir/schedule/analysis/analysis.cc +++ b/src/s_tir/schedule/analysis/analysis.cc @@ -1325,8 +1325,8 @@ std::pair, bool> GetBufferDefiningSite(const StmtSRef& b continue; } // Try to find the buffer in `allloc_buffers` - for (const BufferVar& alloc_buffer : block->alloc_buffers) { - if (buffer.same_as(alloc_buffer)) { + for (const BufferVar& alloc_tensor : block->alloc_buffers) { + if (buffer.same_as(alloc_tensor)) { return {ffi::GetRef(defining_site_sref), true}; } } diff --git a/src/s_tir/schedule/analysis/reducer.cc b/src/s_tir/schedule/analysis/reducer.cc index 7d1fddc0ea6b..dc99866030d2 100644 --- a/src/s_tir/schedule/analysis/reducer.cc +++ b/src/s_tir/schedule/analysis/reducer.cc @@ -595,12 +595,12 @@ bool ReductionIterNotIndexOutputBuffer(const SBlock& block) { return ffi::WalkResult::Advance(); }; auto visit_alloc = [&](const tirx::Bind& alloc) -> ffi::Expected { - // Inline AllocBuffer statements (e.g. `T.local_scalar(...)` expansions) + // Inline AllocTensor statements (e.g. `T.local_scalar(...)` expansions) // declare buffer-local scratch storage inside the block body; treat them // the same as block->alloc_buffers entries for the "write-without-signature" // check below. if (const auto* call = alloc->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { + call && call->op.same_as(tirx::builtin::alloc_tensor())) { buffer_allocated.insert(alloc->var.get()); } return ffi::WalkResult::Advance(); diff --git a/src/s_tir/schedule/ir_comparator.cc b/src/s_tir/schedule/ir_comparator.cc index d5f78b816aac..40834b44eb8d 100644 --- a/src/s_tir/schedule/ir_comparator.cc +++ b/src/s_tir/schedule/ir_comparator.cc @@ -514,7 +514,7 @@ bool TensorizeComparator::CompareBuffer(const BufferVar& lhs, const BufferVar& r equal = (*it).second.same_as(lhs); } else { // Remap the buffer variable definition without recursively comparing its - // BufferType. Tensorization intentionally matches a region of a larger + // TensorType. Tensorization intentionally matches a region of a larger // workload buffer against the intrinsic's smaller descriptor buffer. auto data_it = equal_map_.find(lhs.var()); if (data_it != equal_map_.end()) { diff --git a/src/s_tir/schedule/primitive.h b/src/s_tir/schedule/primitive.h index 89c058f9ba70..99a718a6215c 100644 --- a/src/s_tir/schedule/primitive.h +++ b/src/s_tir/schedule/primitive.h @@ -701,7 +701,7 @@ TVM_DLL void PadEinsum(ScheduleState self, const StmtSRef& block_sref, * appears in the block's ancestor loops as `rolling axis`, fold and circularize the buffer along * the rolling dimension, append block predicate to avoid recomputing overlapping elements. * It requires: - * 1) The buffer to be an intermediate buffer defined via `alloc_buffer`. + * 1) The buffer to be an intermediate buffer defined via `alloc_tensor`. * 2) The LCA of the producer and consumer of the buffer is a for loop, typically, * the producer and consumer of the buffer are cascaded through compute_at. * 3) The access region of the buffer has at least one dimension that contains diff --git a/src/s_tir/schedule/primitive/cache_index.cc b/src/s_tir/schedule/primitive/cache_index.cc index 88fc3cbf6e97..f720b2d93c92 100644 --- a/src/s_tir/schedule/primitive/cache_index.cc +++ b/src/s_tir/schedule/primitive/cache_index.cc @@ -283,7 +283,7 @@ ffi::Array MakeIndexCacheStage(IndexInfo* info, const ffi::String& stora sym::EvalSet(info->var_binding.at(it), sym::AsIntSet(info->range_map)).max() + 1); } info->cache_buffer.push_back(BufferVar( - index_buffer_name, BufferType(storage_scope, data_ty, buffer_shape, {1}, {0}, 0, 0))); + index_buffer_name, TensorType(storage_scope, data_ty, buffer_shape, {1}, {0}, 0, 0))); // Create loop vars and block vars' binding_value std::vector loop_vars; diff --git a/src/s_tir/schedule/primitive/cache_read_write.cc b/src/s_tir/schedule/primitive/cache_read_write.cc index 1dabc9f01ea2..a79903d90d78 100644 --- a/src/s_tir/schedule/primitive/cache_read_write.cc +++ b/src/s_tir/schedule/primitive/cache_read_write.cc @@ -1080,7 +1080,7 @@ class CacheReadRewriter : public StmtExprMutator { if (block == scope_sref_->stmt) { // If so, put buffer allocation on the parent scope ffi::ObjectPtr n = ffi::make_object(*stmt.as()); - // In cache_inplace case, alloc_buffer may be already exits. + // In cache_inplace case, alloc_tensor may be already exits. if (info_->alloc.has_value()) { n->alloc_buffers.push_back(info_->alloc.value()); stmt = SBlock(n); @@ -1403,7 +1403,7 @@ class CacheWriteRewriter : public StmtExprMutator { // Put buffer allocation on the parent scope if (block == scope_sref_->stmt) { ffi::ObjectPtr n = ffi::make_object(*stmt.as()); - // In cache_inplace case, alloc_buffer may be already exits. + // In cache_inplace case, alloc_tensor may be already exits. if (info_->alloc.has_value()) { n->alloc_buffers.push_back(info_->alloc.value()); stmt = SBlock(n); @@ -1638,7 +1638,7 @@ class ReindexCacheWriteRewriter : public CacheWriteRewriter { */ BufferVar CreateReindexBuffer(const BufferVar& buffer, const ffi::Array& block_iters, const std::unordered_set& covered) { - ffi::ObjectPtr new_buffer = CopyBufferType(buffer); + ffi::ObjectPtr new_buffer = CopyTensorType(buffer); std::vector new_shape; std::vector new_strides; for (const auto& iter : block_iters) { @@ -2031,7 +2031,7 @@ StmtSRef CacheRead(ScheduleState self, const StmtSRef& block_sref, int read_buff info.cache_region = cache_region; info.write_buffer = WithScope(read_buffer, storage_scope); if (!cache_full_region) { - auto write_buffer = CopyBufferType(info.write_buffer); + auto write_buffer = CopyTensorType(info.write_buffer); std::vector shape; for (auto cache_range : info.cache_region->region) { shape.push_back(cache_range->extent); @@ -2121,7 +2121,7 @@ StmtSRef CacheWrite(ScheduleState self, const StmtSRef& block_sref, int write_bu info.cache_region = cache_region; info.read_buffer = WithScope(write_buffer, storage_scope); if (!cache_full_region) { - auto read_buffer_type = CopyBufferType(info.read_buffer); + auto read_buffer_type = CopyTensorType(info.read_buffer); std::vector shape; for (auto cache_range : info.cache_region->region) { shape.push_back(cache_range->extent); @@ -2268,7 +2268,7 @@ void CollectReindexCacheStageInfoAndCreateBuffer( } // Create new buffer - ffi::ObjectPtr new_buffer = CopyBufferType(old_buffer); + ffi::ObjectPtr new_buffer = CopyTensorType(old_buffer); new_buffer->storage_scope = storage_scope; new_buffer->shape = new_shape; BufferVar rebuilt = diff --git a/src/s_tir/schedule/primitive/compute_inline.cc b/src/s_tir/schedule/primitive/compute_inline.cc index 60fcf4112a6d..17159f149fb3 100644 --- a/src/s_tir/schedule/primitive/compute_inline.cc +++ b/src/s_tir/schedule/primitive/compute_inline.cc @@ -411,7 +411,7 @@ class BaseInliner : public StmtExprMutator { /*! * \brief Update the following block signature: - * 1) T.alloc_buffer, if the block is scope root + * 1) T.alloc_tensor, if the block is scope root * 2) T.reads, if the block is not scope root * 3) T.writes, if the block is not scope root * \param block The block to be updated @@ -423,9 +423,9 @@ class BaseInliner : public StmtExprMutator { ffi::Array alloc_buffers; if (is_scope_root) { alloc_buffers.reserve(block->alloc_buffers.size()); - for (const BufferVar& alloc_buffer : block->alloc_buffers) { - if (!alloc_buffer.same_as(inlined_buffer_)) { - alloc_buffers.push_back(alloc_buffer); + for (const BufferVar& alloc_tensor : block->alloc_buffers) { + if (!alloc_tensor.same_as(inlined_buffer_)) { + alloc_buffers.push_back(alloc_tensor); } } } else { diff --git a/src/s_tir/schedule/primitive/layout_transformation.cc b/src/s_tir/schedule/primitive/layout_transformation.cc index 5099f435c1bb..2e5ff6f948f0 100644 --- a/src/s_tir/schedule/primitive/layout_transformation.cc +++ b/src/s_tir/schedule/primitive/layout_transformation.cc @@ -1292,7 +1292,7 @@ void TransformLayout(ScheduleState self, const StmtSRef& block_sref, int buffer_ } // Step 2: Infer the shape of the new buffer - auto new_buffer_type = CopyBufferType(old_buffer); + auto new_buffer_type = CopyTensorType(old_buffer); new_buffer_type->shape = index_map->MapShape(old_buffer->shape, analyzer); BufferVar new_buffer = RebuildBufferVar(old_buffer, std::move(new_buffer_type)); diff --git a/src/s_tir/schedule/primitive/pad_einsum.cc b/src/s_tir/schedule/primitive/pad_einsum.cc index 6cafdc3536db..899004f2a5e2 100644 --- a/src/s_tir/schedule/primitive/pad_einsum.cc +++ b/src/s_tir/schedule/primitive/pad_einsum.cc @@ -156,7 +156,7 @@ struct BufferPadding { shape.push_back(buffer_region->source.as_or_throw()->shape[i]); } } - result.padded_buffer = decl_buffer(shape, result.buffer->dtype, result.buffer.name() + "_pad", + result.padded_buffer = decl_tensor(shape, result.buffer->dtype, result.buffer.name() + "_pad", result.buffer.scope()); return result; } diff --git a/src/s_tir/schedule/primitive/reduction.cc b/src/s_tir/schedule/primitive/reduction.cc index 33388020dc5d..569a2f8532b0 100644 --- a/src/s_tir/schedule/primitive/reduction.cc +++ b/src/s_tir/schedule/primitive/reduction.cc @@ -749,7 +749,7 @@ ffi::Array CreateRFactorBuffers(const ffi::Array& buf_st ffi::Array rf_shape = buffer->shape; rf_shape.insert(rf_shape.begin() + factor_axis, rf_loop->extent); - ffi::ObjectPtr n = CopyBufferType(buffer); + ffi::ObjectPtr n = CopyTensorType(buffer); n->shape = rf_shape; rf_buffers.push_back(RebuildBufferVar(buffer, std::move(n), buffer.name() + ".rf")); } diff --git a/src/s_tir/schedule/primitive/rolling_buffer.cc b/src/s_tir/schedule/primitive/rolling_buffer.cc index e1a96a657ed2..cdf2dd4f8041 100644 --- a/src/s_tir/schedule/primitive/rolling_buffer.cc +++ b/src/s_tir/schedule/primitive/rolling_buffer.cc @@ -260,7 +260,7 @@ class RollingBufferInfoCollector { } ffi::Array new_shape = buffer->shape; new_shape.Set(roll_axis, region[roll_axis]->extent); - auto new_buffer_type = CopyBufferType(buffer); + auto new_buffer_type = CopyTensorType(buffer); new_buffer_type->shape = new_shape; BufferVar new_buffer = RebuildBufferVar(buffer, std::move(new_buffer_type)); diff --git a/src/s_tir/schedule/transform.cc b/src/s_tir/schedule/transform.cc index 6f7e27af083e..082fb79f7c98 100644 --- a/src/s_tir/schedule/transform.cc +++ b/src/s_tir/schedule/transform.cc @@ -45,14 +45,14 @@ SBlock WithAnnotation(const SBlockNode* block, const ffi::String& attr_key, /******** Buffer Related ********/ BufferVar WithScope(const BufferVar& buffer, const ffi::String& scope) { - BufferType new_type(scope, buffer->dtype, buffer->shape, buffer->strides, buffer->elem_offset, + TensorType new_type(scope, buffer->dtype, buffer->shape, buffer->strides, buffer->elem_offset, buffer->data_alignment, buffer->offset_factor, buffer->layout, buffer->allocated_addr); return BufferVar(buffer.name() + "_" + scope, new_type, buffer.span()); } BufferVar WithDType(const BufferVar& buffer, PrimType dtype) { - BufferType new_type(buffer->storage_scope, dtype, buffer->shape, buffer->strides, + TensorType new_type(buffer->storage_scope, dtype, buffer->shape, buffer->strides, buffer->elem_offset, buffer->data_alignment, buffer->offset_factor, buffer->layout, buffer->allocated_addr); return BufferVar(buffer.name(), new_type, buffer.span()); diff --git a/src/s_tir/script/ir_builder/ir.cc b/src/s_tir/script/ir_builder/ir.cc index a54c9251a194..c3256273e2c0 100644 --- a/src/s_tir/script/ir_builder/ir.cc +++ b/src/s_tir/script/ir_builder/ir.cc @@ -195,7 +195,7 @@ BufferVar SBlockAllocBuffer(ffi::Array shape, PrimType dtype, ffi::Opt if (scope == "global" || scope == "shared" || scope == "shared.dyn" || scope == "local") { TVM_FFI_ICHECK(allocated_addr.empty()) << "ValueError: For `" << scope - << "` scope, Ts.alloc_buffer does not accept `allocated_addr`"; + << "` scope, Ts.alloc_tensor does not accept `allocated_addr`"; } ffi::Optional opt_elem_offset = elem_offset.defined() ? ffi::Optional(elem_offset) : std::nullopt; @@ -205,8 +205,8 @@ BufferVar SBlockAllocBuffer(ffi::Array shape, PrimType dtype, ffi::Opt auto opt_func_frame = builder->FindFrame(); if (opt_func_frame.has_value()) { TVM_FFI_CHECK(opt_func_frame.value().as() != nullptr, ValueError) - << "ValueError: `Ts.alloc_buffer()` is only for s_tir PrimFuncs. " - "Use `T.alloc_buffer()` inside default (tirx) PrimFuncs."; + << "ValueError: `Ts.alloc_tensor()` is only for s_tir PrimFuncs. " + "Use `T.alloc_tensor()` inside default (tirx) PrimFuncs."; } // Walk up the frame stack: attach to the innermost enclosing s_tir::SBlock (lifting diff --git a/src/s_tir/script/ir_builder/script_complete.cc b/src/s_tir/script/ir_builder/script_complete.cc index 5fa423dbecc0..a91d6b3e556e 100644 --- a/src/s_tir/script/ir_builder/script_complete.cc +++ b/src/s_tir/script/ir_builder/script_complete.cc @@ -64,8 +64,8 @@ class ScriptCompleter : public s_tir::StmtExprMutator { UnchangedOr Mutate_(const s_tir::SBlockNode* op, InplaceMode inplace_mode) final { // Buffers allocated in the block can be accessed by its body. - for (const auto& alloc_buffer : op->alloc_buffers) { - buffer_var_map_->Set(alloc_buffer.var(), alloc_buffer); + for (const auto& alloc_tensor : op->alloc_buffers) { + buffer_var_map_->Set(alloc_tensor.var(), alloc_tensor); } for (const auto& match_buffer : op->match_buffers) { const BufferVar& target_buffer = match_buffer->buffer; @@ -80,8 +80,8 @@ class ScriptCompleter : public s_tir::StmtExprMutator { this->is_root_block_ = is_root_block; // Remove buffers allocated inside block to detect its access region - for (const auto& alloc_buffer : op->alloc_buffers) { - buffer_var_map_->erase(alloc_buffer.var()); + for (const auto& alloc_tensor : op->alloc_buffers) { + buffer_var_map_->erase(alloc_tensor.var()); } for (const auto& match_buffer : op->match_buffers) { const BufferVar& target_buffer = match_buffer->buffer; diff --git a/src/s_tir/script/printer/stmt.cc b/src/s_tir/script/printer/stmt.cc index a57326ac57b6..bd2b9f3e4ba2 100644 --- a/src/s_tir/script/printer/stmt.cc +++ b/src/s_tir/script/printer/stmt.cc @@ -95,9 +95,9 @@ ffi::Array SBlockBody(DocTranslatorObj* d, const s_tir::SBlockNode* blo } for (const tirx::BufferVar& buffer : block->alloc_buffers) { CallDoc rhs = d->Translate(buffer.var()->ty).value().as_or_throw(); - TVM_FFI_CHECK(rhs->callee.as_or_throw()->name == "Buffer", TypeError) - << "Ts.sblock_alloc_buffer cannot reconstruct this nonrepresentable BufferType"; - const auto* buffer_type = buffer.var()->ty.as(); + TVM_FFI_CHECK(rhs->callee.as_or_throw()->name == "Tensor", TypeError) + << "Ts.sblock_alloc_buffer cannot reconstruct this nonrepresentable TensorType"; + const auto* buffer_type = buffer.var()->ty.as(); TVM_FFI_CHECK( buffer_type->allocated_addr.empty() || (buffer_type->storage_scope != "global" && buffer_type->storage_scope != "shared" && @@ -171,8 +171,8 @@ ffi::Optional MatchBufferRegionDocTranslate(DocTranslatorObj* d, ffi::A << "printer statement-only node cannot fulfill a destination"; ExprDoc source = d->Translate(match->source).value(); CallDoc rhs = d->Translate(match->buffer.var()->ty).value().as_or_throw(); - TVM_FFI_CHECK(rhs->callee.as_or_throw()->name == "Buffer", TypeError) - << "Ts.match_buffer cannot reconstruct this nonrepresentable BufferType"; + TVM_FFI_CHECK(rhs->callee.as_or_throw()->name == "Tensor", TypeError) + << "Ts.match_buffer cannot reconstruct this nonrepresentable TensorType"; rhs->callee = NamespaceDoc("s_tir")->Attr("match_buffer"); rhs->args.insert(rhs->args.begin(), source); IdDoc lhs = VarDoc(d, match->buffer); diff --git a/src/s_tir/tensor_intrin.cc b/src/s_tir/tensor_intrin.cc index e381448198d5..d2c21d6c0b4d 100644 --- a/src/s_tir/tensor_intrin.cc +++ b/src/s_tir/tensor_intrin.cc @@ -26,8 +26,8 @@ namespace tvm { namespace s_tir { -using tirx::BufferTypeNode; using tirx::PrimFunc; +using tirx::TensorTypeNode; TVM_FFI_STATIC_INIT_BLOCK() { TensorIntrinNode::RegisterReflection(); } @@ -47,7 +47,7 @@ TensorIntrin::TensorIntrin(PrimFunc desc, PrimFunc impl) { << "The number of parameters of the description and the implementation of the " "tensor intrinsic doesn't match."; auto is_handle = [](const Var& param) { - return param->ty.as() != nullptr || param->ty.as() != nullptr; + return param->ty.as() != nullptr || param->ty.as() != nullptr; }; for (size_t i = 0; i < desc->params.size(); i++) { TVM_FFI_CHECK(is_handle(desc->params[i]), ValueError) diff --git a/src/s_tir/transform/bound_checker.cc b/src/s_tir/transform/bound_checker.cc index 89df4a57cf07..dd45276eeaea 100644 --- a/src/s_tir/transform/bound_checker.cc +++ b/src/s_tir/transform/bound_checker.cc @@ -81,13 +81,13 @@ class BoundChecker : public StmtExprMutator { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return Mutate_AllocBuffer(op, call, inplace_mode); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return Mutate_AllocTensor(op, call, inplace_mode); } return StmtExprMutator::Mutate_(op, inplace_mode); } - UnchangedOr Mutate_AllocBuffer(const BindNode* op, const CallNode* call, + UnchangedOr Mutate_AllocTensor(const BindNode* op, const CallNode* call, InplaceMode inplace_mode) { if (UpdateIsNeeded(op->var.as_or_throw().var())) { tvm::Tuple shape = call->args[0].as_or_throw(); diff --git a/src/s_tir/transform/compact_buffer_region.cc b/src/s_tir/transform/compact_buffer_region.cc index 516c149dc06f..87755cdee44c 100644 --- a/src/s_tir/transform/compact_buffer_region.cc +++ b/src/s_tir/transform/compact_buffer_region.cc @@ -126,8 +126,8 @@ class BufferAccessRegionCollector : public StmtExprVisitor { // collect buffer access regions region_collector->Visit(f->body); - // Compact any remaining flat AllocBuffer nodes at function scope - region_collector->CompactPendingFlatAllocBuffers(); + // Compact any remaining flat AllocTensor nodes at function scope + region_collector->CompactPendingFlatAllocTensors(); return std::move(region_collector->buffer_access_region_); } @@ -194,8 +194,8 @@ class BufferAccessRegionCollector : public StmtExprVisitor { dom_map_.emplace(op->loop_var.get(), sym::IntSet::FromRange(loop_range)); size_t n_pending_before = pending_flat_alloc_buffers_.size(); TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(StmtExprVisitor::Visit_(op)); - // Compact flat AllocBuffers defined inside this For scope - CompactPendingFlatAllocBuffers(n_pending_before); + // Compact flat AllocTensors defined inside this For scope + CompactPendingFlatAllocTensors(n_pending_before); dom_map_.erase(op->loop_var.get()); ancestor_iters_.pop_back(); return std::nullopt; @@ -203,11 +203,11 @@ class BufferAccessRegionCollector : public StmtExprVisitor { ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return DispatchAllocBuffer(op); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return DispatchAllocTensor(op); } if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::decl_buffer())) + call && call->op.same_as(tirx::builtin::decl_tensor())) return StmtExprVisitor::Visit_(op); TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(StmtExprVisitor::Visit(op->value)); if (auto value = op->value.as(); value && sym::IsIndexTypedExpr(value.value())) { @@ -346,8 +346,8 @@ class BufferAccessRegionCollector : public StmtExprVisitor { return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchAllocBuffer(const BindNode* op) { - // AllocBuffer is flat: register the buffer def and track for post-scope compaction. + ffi::Optional DispatchAllocTensor(const BindNode* op) { + // AllocTensor is flat: register the buffer def and track for post-scope compaction. RecordBufferDefinition(op->var.as_or_throw().var()); pending_flat_alloc_buffers_.push_back(op->var.as_or_throw()); return StmtExprVisitor::Visit_(op); @@ -366,7 +366,7 @@ class BufferAccessRegionCollector : public StmtExprVisitor { dom_map_.emplace(iter->var.get(), sym::IntSet::FromRange(dom)); size_t n_pending_before = pending_flat_alloc_buffers_.size(); TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(StmtExprVisitor::Visit_(op)); - CompactPendingFlatAllocBuffers(n_pending_before); + CompactPendingFlatAllocTensors(n_pending_before); dom_map_.erase(iter->var.get()); ancestor_iters_.pop_back(); return std::nullopt; @@ -524,10 +524,10 @@ class BufferAccessRegionCollector : public StmtExprVisitor { } /*! - * \brief Compact pending flat AllocBuffer nodes registered since position n_before. + * \brief Compact pending flat AllocTensor nodes registered since position n_before. * Call SimplifyAndNarrowBufferRegionFromNDIntSet for each, then remove them. */ - void CompactPendingFlatAllocBuffers(size_t n_before = 0) { + void CompactPendingFlatAllocTensors(size_t n_before = 0) { for (size_t i = n_before; i < pending_flat_alloc_buffers_.size(); ++i) { const BufferVar& buf = pending_flat_alloc_buffers_[i]; auto it = relaxed_accesses_.find(buf); @@ -541,7 +541,7 @@ class BufferAccessRegionCollector : public StmtExprVisitor { /**************** Class members ****************/ /*! \brief Only collect accessed region within original buffer shape bound. */ bool collect_inbound_{true}; - /*! \brief Pending flat AllocBuffer nodes to compact when leaving scope. */ + /*! \brief Pending flat AllocTensor nodes to compact when leaving scope. */ std::vector pending_flat_alloc_buffers_; /*! \brief The iteration scopes from the current node up to the root. */ @@ -646,7 +646,7 @@ class BufferCompactor : public StmtExprMutator { RewriteBufferRegions(&n->writes); RewriteMatchBuffers(&n->match_buffers); n->alloc_buffers = - op->alloc_buffers.Map([this](const BufferVar& buf) { return RewriteAllocBuffer(buf); }); + op->alloc_buffers.Map([this](const BufferVar& buf) { return RewriteAllocTensor(buf); }); // Recursively rewrite the body after installing the allocation remaps. return StmtExprMutator::Mutate_(block.get(), block.unique() ? inplace_mode : InplaceMode::kDisallow) @@ -655,13 +655,13 @@ class BufferCompactor : public StmtExprMutator { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { const auto* call = op->value.as(); - if (!call || (!call->op.same_as(tirx::builtin::alloc_buffer()) && - !call->op.same_as(tirx::builtin::decl_buffer()))) { + if (!call || (!call->op.same_as(tirx::builtin::alloc_tensor()) && + !call->op.same_as(tirx::builtin::decl_tensor()))) { return StmtExprMutator::Mutate_(op, inplace_mode); } BufferVar buffer = op->var.as_or_throw(); - BufferVar new_buffer = RewriteAllocBuffer(buffer); - bool is_alloc = call->op.same_as(tirx::builtin::alloc_buffer()); + BufferVar new_buffer = RewriteAllocTensor(buffer); + bool is_alloc = call->op.same_as(tirx::builtin::alloc_tensor()); if (new_buffer.same_as(buffer) || (is_alloc && PrimType(call->args[1].as_or_throw()->value) != new_buffer->dtype)) { @@ -679,7 +679,7 @@ class BufferCompactor : public StmtExprMutator { return StmtExprMutator::Mutate_(binding.get(), inplace_mode).ValueOrUnchanged(binding); } - BufferVar RewriteAllocBuffer(const BufferVar& buffer) { + BufferVar RewriteAllocTensor(const BufferVar& buffer) { auto it = buffer_info_.find(buffer.var()); if (it != buffer_info_.end()) { const BufferVar& new_buffer = it->second.new_buffer; @@ -809,7 +809,7 @@ Stmt BufferCompactorCompact( // prepare new buffer ffi::Array shape = region.Map([](const Range& range) { return range->extent; }); ffi::Array strides = CalcStrides(alloc_info, shape); - ffi::ObjectPtr n = CopyBufferType(buffer); + ffi::ObjectPtr n = CopyTensorType(buffer); n->shape = std::move(shape); n->strides = std::move(strides); alloc_info.new_buffer = RebuildBufferVar(buffer, std::move(n)); diff --git a/src/s_tir/transform/hoist_expression.cc b/src/s_tir/transform/hoist_expression.cc index c2f78150ca02..8cb883f68bce 100644 --- a/src/s_tir/transform/hoist_expression.cc +++ b/src/s_tir/transform/hoist_expression.cc @@ -349,8 +349,8 @@ class HoistInfoCollector : public StmtExprVisitor { if (!bind) { non_bind_count++; } else if (const auto* call = bind->value.as(); - call && (call->op.same_as(tirx::builtin::alloc_buffer()) || - call->op.same_as(tirx::builtin::decl_buffer()))) { + call && (call->op.same_as(tirx::builtin::alloc_tensor()) || + call->op.same_as(tirx::builtin::decl_tensor()))) { non_bind_count++; } } diff --git a/src/s_tir/transform/inject_double_buffer.cc b/src/s_tir/transform/inject_double_buffer.cc index d3d61f5ca044..c1db8eba78f4 100644 --- a/src/s_tir/transform/inject_double_buffer.cc +++ b/src/s_tir/transform/inject_double_buffer.cc @@ -176,13 +176,13 @@ class DoubleBufferInjector : public StmtExprMutator { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return Mutate_AllocBuffer(op, call, inplace_mode); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return Mutate_AllocTensor(op, call, inplace_mode); } return StmtExprMutator::Mutate_(op, inplace_mode); } - UnchangedOr Mutate_AllocBuffer(const BindNode* op, const CallNode* call, + UnchangedOr Mutate_AllocTensor(const BindNode* op, const CallNode* call, InplaceMode inplace_mode) { const VarNode* buf = op->var.as_or_throw().get(); auto it = dbuffer_info_.find(buf); @@ -195,11 +195,11 @@ class DoubleBufferInjector : public StmtExprMutator { << "Has FlattenBuffer been run?"; entry.stride = shape->fields[0].as_or_throw(); - // In flat IR, AllocBuffer appears before its usage in the SeqStmt, + // In flat IR, AllocTensor appears before its usage in the SeqStmt, // so entry.loop may not be set yet. Defer double-buffer allocation // processing to be handled in VisitStmt_(ForNode*). pending_dbuffer_allocs_[buf] = ffi::GetRef(op); - // Remove the original AllocBuffer (will be re-emitted in ForNode visitor) + // Remove the original AllocTensor (will be re-emitted in ForNode visitor) return Evaluate(0); } else { return StmtExprMutator::Mutate_(op, inplace_mode); @@ -221,7 +221,7 @@ class DoubleBufferInjector : public StmtExprMutator { auto& alloc_nest = loop_allocs_[entry.loop]; const auto* call = alloc->value.as(); alloc_nest.emplace_back(Bind(new_buf.var(), - Call(new_buf.type(), tirx::builtin::alloc_buffer(), + Call(new_buf.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(new_buf->shape, call->args[0]->span), DataTypeImm(new_buf->dtype->dtype, call->args[1]->span), StringImm(new_buf.scope(), call->args[2]->span)}, @@ -352,7 +352,7 @@ class DoubleBufferInjector : public StmtExprMutator { // Stride gives the distance between the two halves of the // double-buffer, not the stride of the buffer's index. - auto type = CopyBufferType(buf); + auto type = CopyTensorType(buf); type->shape = {buf->shape[0] + stride}; buf = RebuildBufferVar(buf, std::move(type)); @@ -441,7 +441,7 @@ class DoubleBufferInjector : public StmtExprMutator { // The allocation size of the buffer std::unordered_map dbuffer_info_; // The updated BufferVar objects - // Pending double-buffer AllocBuffer nodes (deferred from flat AllocBuffer visit) + // Pending double-buffer AllocTensor nodes (deferred from flat AllocTensor visit) std::unordered_map pending_dbuffer_allocs_; }; diff --git a/src/s_tir/transform/inject_software_pipeline.cc b/src/s_tir/transform/inject_software_pipeline.cc index 40da4541bb67..34ce611d7a2d 100644 --- a/src/s_tir/transform/inject_software_pipeline.cc +++ b/src/s_tir/transform/inject_software_pipeline.cc @@ -285,14 +285,14 @@ class PipelineBodyRewriter : public StmtExprMutator { } UnchangedOr Mutate_(const SBlockNode* op, InplaceMode inplace_mode) final { - for (const BufferVar& alloc_buffer : op->alloc_buffers) { - buffer_data_to_buffer_.Set(alloc_buffer.var(), alloc_buffer); + for (const BufferVar& alloc_tensor : op->alloc_buffers) { + buffer_data_to_buffer_.Set(alloc_tensor.var(), alloc_tensor); } SBlock block = StmtExprMutator::Mutate_(op, inplace_mode) .ValueOrUnchanged(ffi::GetRef(op)) .as_or_throw(); - for (const BufferVar& alloc_buffer : op->alloc_buffers) { - buffer_data_to_buffer_.erase(alloc_buffer.var()); + for (const BufferVar& alloc_tensor : op->alloc_buffers) { + buffer_data_to_buffer_.erase(alloc_tensor.var()); } return block; } @@ -392,7 +392,7 @@ class PipelineRewriter : public StmtExprMutator { for (const BufferVar& buffer : pipeline_allocs_) { int num_versions = ComputeBufferVersions(buffer, infos.at(buffer)); if (num_versions > 1) { - buffer_remap_.Set(buffer, RewriteAllocBuffer(buffer, num_versions)); + buffer_remap_.Set(buffer, RewriteAllocTensor(buffer, num_versions)); } } for (const auto& [_, remapped] : buffer_remap_) { @@ -577,8 +577,8 @@ class PipelineRewriter : public StmtExprMutator { * \param num_versions The number of versions to keep. * \return The resized buffer. */ - BufferVar RewriteAllocBuffer(const BufferVar& buffer, int num_versions) { - ffi::ObjectPtr new_buffer = CopyBufferType(buffer); + BufferVar RewriteAllocTensor(const BufferVar& buffer, int num_versions) { + ffi::ObjectPtr new_buffer = CopyTensorType(buffer); new_buffer->shape.insert(new_buffer->shape.begin(), PrimExpr(num_versions)); if (new_buffer->strides.size()) { TVM_FFI_ICHECK(new_buffer->strides.size() + 1 == new_buffer->shape.size()); @@ -1182,7 +1182,7 @@ class PipelineInjector : public StmtExprMutator { if (const auto* realize = for_node->body.as()) { const auto& block = realize->block; for (const auto& buffer : block->alloc_buffers) { - TVM_FFI_ICHECK(buffer->IsInstance()); + TVM_FFI_ICHECK(buffer->IsInstance()); buffer_data_to_buffer_.Set(buffer.var(), buffer); } pipeline_body = block->body; @@ -1284,14 +1284,14 @@ class PipelineInjector : public StmtExprMutator { * \param alloc_buffers The buffer allocations to be added. */ void AddAllocBuffers(SBlockNode* n, const ffi::Array alloc_buffers) { - for (const BufferVar& alloc_buffer : alloc_buffers) { - n->alloc_buffers.push_back(alloc_buffer); + for (const BufferVar& alloc_tensor : alloc_buffers) { + n->alloc_buffers.push_back(alloc_tensor); Region region; - region.reserve(alloc_buffer->shape.size()); - for (const PrimExpr& dim : alloc_buffer->shape) { + region.reserve(alloc_tensor->shape.size()); + for (const PrimExpr& dim : alloc_tensor->shape) { region.push_back(Range::FromMinExtent(0, dim)); } - n->writes.push_back(BufferRegion(alloc_buffer, region)); + n->writes.push_back(BufferRegion(alloc_tensor, region)); } } diff --git a/src/s_tir/transform/inject_virtual_thread.cc b/src/s_tir/transform/inject_virtual_thread.cc index ddcfc2e036b1..c0cf41779908 100644 --- a/src/s_tir/transform/inject_virtual_thread.cc +++ b/src/s_tir/transform/inject_virtual_thread.cc @@ -152,11 +152,11 @@ class VarTouchedAnalysis : public StmtExprVisitor { } ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return DispatchAllocBuffer(op, call); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return DispatchAllocTensor(op, call); } if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::decl_buffer())) + call && call->op.same_as(tirx::builtin::decl_tensor())) return StmtExprVisitor::Visit_(op); expr_touched_->Reset(false); expr_touched_->Visit(op->value); @@ -189,7 +189,7 @@ class VarTouchedAnalysis : public StmtExprVisitor { } return std::nullopt; } - ffi::Optional DispatchAllocBuffer(const BindNode* op, const CallNode* call) { + ffi::Optional DispatchAllocTensor(const BindNode* op, const CallNode* call) { expr_touched_->Reset(false); tvm::Tuple shape = call->args[0].as_or_throw(); for (const Expr& extent : shape->fields) { @@ -345,7 +345,7 @@ class VTInjector : public s_tir::IRMutatorWithAnalyzer { PrimExpr stride = it->second / prim::MakeConst(offset.ty(), dtype.lanes()); offset = RewriteIndex(offset, stride); Expr data = - buffer.value()->ty.as() + buffer.value()->ty.as() ? GetRemappedBuffer(buffer.value().as_or_throw(), it->second).data() : op->args[1]; @@ -420,7 +420,7 @@ class VTInjector : public s_tir::IRMutatorWithAnalyzer { TVM_FFI_ICHECK_EQ(buf->shape.size(), 1) << "Expected buffers being rewritten to already be flattened."; - auto writer = CopyBufferType(buf); + auto writer = CopyTensorType(buf); writer->shape = {buf->shape[0] * alloc_extent}; buf = RebuildBufferVar(buf, std::move(writer)); @@ -449,11 +449,11 @@ class VTInjector : public s_tir::IRMutatorWithAnalyzer { // Bind UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return Mutate_AllocBuffer(op, call, inplace_mode); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return Mutate_AllocTensor(op, call, inplace_mode); } if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::decl_buffer())) + call && call->op.same_as(tirx::builtin::decl_tensor())) return StmtExprMutator::Mutate_(op, inplace_mode); auto value_result = this->Mutate(op->value, inplace_mode); bool value_unchanged = value_result.UnchangedOrSameAs(op->value); @@ -591,8 +591,8 @@ class VTInjector : public s_tir::IRMutatorWithAnalyzer { return SeqStmt(new_seq); } // Allocate - // AllocBuffer - UnchangedOr Mutate_AllocBuffer(const BindNode* op, const CallNode* call, + // AllocTensor + UnchangedOr Mutate_AllocTensor(const BindNode* op, const CallNode* call, InplaceMode inplace_mode) { tvm::Tuple shape = call->args[0].as_or_throw(); ffi::Array original_shape = shape->fields.as_or_throw>(); @@ -618,12 +618,12 @@ class VTInjector : public s_tir::IRMutatorWithAnalyzer { if (new_shape.same_as(original_shape)) { return ffi::Unchanged(); } else { - auto type = CopyBufferType(op->var.as_or_throw()); + auto type = CopyTensorType(op->var.as_or_throw()); type->shape = new_shape; BufferVar new_buffer = RebuildBufferVar(op->var.as_or_throw(), std::move(type)); VarRemapSet(op->var.as_or_throw(), new_buffer); return Bind(new_buffer.var(), - Call(new_buffer.type(), tirx::builtin::alloc_buffer(), + Call(new_buffer.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(new_buffer->shape, call->args[0]->span), DataTypeImm(new_buffer->dtype->dtype, call->args[1]->span), StringImm(new_buffer.scope(), call->args[2]->span)}, diff --git a/src/s_tir/transform/ir_utils.cc b/src/s_tir/transform/ir_utils.cc index e10c08eecf2e..928edf1c03b4 100644 --- a/src/s_tir/transform/ir_utils.cc +++ b/src/s_tir/transform/ir_utils.cc @@ -158,16 +158,16 @@ class StorageAlignCollector : public StmtExprVisitor { return StmtExprVisitor::Visit_(op); } - /*! \brief AllocBuffer: check for buffer_dim_align annotations. */ + /*! \brief AllocTensor: check for buffer_dim_align annotations. */ ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return DispatchAllocBuffer(op, call); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return DispatchAllocTensor(op, call); } return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchAllocBuffer(const BindNode* op, const CallNode* call) { + ffi::Optional DispatchAllocTensor(const BindNode* op, const CallNode* call) { DictAttrs annotations = call->attrs.as_or_throw(); auto it = annotations->dict.find(attr::buffer_dim_align); if (it != annotations->dict.end()) { diff --git a/src/s_tir/transform/lower_cross_thread_reduction.cc b/src/s_tir/transform/lower_cross_thread_reduction.cc index fbac8846816b..c64d4da116a1 100644 --- a/src/s_tir/transform/lower_cross_thread_reduction.cc +++ b/src/s_tir/transform/lower_cross_thread_reduction.cc @@ -151,7 +151,7 @@ ffi::Array MakeScratchpads(const ffi::Array& reduction_buf for (const BufferVar& buffer : reduction_buffers) { ffi::String name = is_cross_thread_buffer ? "cross" : "in"; name = name + "_thread_" + buffer.name(); - new_buffers.push_back(BufferVar(name, BufferType(/*storage_scope=*/"local", + new_buffers.push_back(BufferVar(name, TensorType(/*storage_scope=*/"local", /*dtype=*/buffer->dtype, /*shape=*/{IntImm::Int32(1)}, /*strides=*/{IntImm::Int32(1)}, diff --git a/src/s_tir/transform/lower_match_buffer.cc b/src/s_tir/transform/lower_match_buffer.cc index 98d9d6f88d8f..452228784968 100644 --- a/src/s_tir/transform/lower_match_buffer.cc +++ b/src/s_tir/transform/lower_match_buffer.cc @@ -104,7 +104,7 @@ class MatchBufferLower : public StmtExprMutator { if ((op->op.same_as(tirx::builtin::masked_load()) || op->op.same_as(tirx::builtin::masked_store())) && !op->args.empty()) { - if (auto var = op->args[0].as(); var && var.value()->ty.as()) { + if (auto var = op->args[0].as(); var && var.value()->ty.as()) { BufferVar buffer = var.value().as_or_throw(); TVM_FFI_ICHECK(!match_buffers_.count(buffer)) << "Predicated buffer access is not currently supported in lower match buffer pass."; @@ -112,7 +112,7 @@ class MatchBufferLower : public StmtExprMutator { } if (op->op.same_as(tirx::builtin::buffer_data()) && op->args.size() == 1) { if (auto var = op->args[0].as(); - var.has_value() && var.value()->ty.as()) { + var.has_value() && var.value()->ty.as()) { auto it = match_buffers_.find(var.value().as_or_throw()); if (it != match_buffers_.end()) { return (*it).second->source.as_or_throw().data(); diff --git a/src/s_tir/transform/lower_opaque_block.cc b/src/s_tir/transform/lower_opaque_block.cc index f2bf16c0b21e..85bc60cceb37 100644 --- a/src/s_tir/transform/lower_opaque_block.cc +++ b/src/s_tir/transform/lower_opaque_block.cc @@ -81,7 +81,7 @@ class OpaqueBlockLower : public StmtExprMutator { IntImm::Int32(buffer->data_alignment)); allocate_annotations.Set(tirx::attr::buffer_allocated_addr, buffer->allocated_addr); body = SeqStmt::Flatten( - Bind(buffer.var(), Call(buffer.type(), tirx::builtin::alloc_buffer(), + Bind(buffer.var(), Call(buffer.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())}, DictAttrs(allocate_annotations))), diff --git a/src/s_tir/transform/lower_vtcm_alloc.cc b/src/s_tir/transform/lower_vtcm_alloc.cc index bd2ce3891586..d34f19dc3f02 100644 --- a/src/s_tir/transform/lower_vtcm_alloc.cc +++ b/src/s_tir/transform/lower_vtcm_alloc.cc @@ -41,13 +41,13 @@ class VtcmAllocator : public StmtExprMutator { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return Mutate_AllocBuffer(op, call, inplace_mode); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return Mutate_AllocTensor(op, call, inplace_mode); } return StmtExprMutator::Mutate_(op, inplace_mode); } - UnchangedOr Mutate_AllocBuffer(const BindNode* op, const CallNode* call, + UnchangedOr Mutate_AllocTensor(const BindNode* op, const CallNode* call, InplaceMode inplace_mode) { ffi::String scope = call->args[2].as_or_throw()->value; if (IsVtcmStorage(scope)) { @@ -60,7 +60,7 @@ class VtcmAllocator : public StmtExprMutator { BufferVar buffer = op->var.as_or_throw(); return Bind( buffer, - Call(buffer.type(), tirx::builtin::decl_buffer(), + Call(buffer.type(), tirx::builtin::decl_tensor(), {Call(buffer.DataPointerType(), tirx::builtin::nd_mem_alloc_with_scope(), args), tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())}, diff --git a/src/s_tir/transform/manifest_shared_memory_local_stage.cc b/src/s_tir/transform/manifest_shared_memory_local_stage.cc index 957ff4d985e1..dd9a7edb3c4a 100644 --- a/src/s_tir/transform/manifest_shared_memory_local_stage.cc +++ b/src/s_tir/transform/manifest_shared_memory_local_stage.cc @@ -170,7 +170,7 @@ class IntermediateStageRewriter { new_buffer_shape.push_back(relaxed_loop->extent); } BufferVar new_buffer = WithScope(buffer, "local"); - ffi::ObjectPtr type = CopyBufferType(new_buffer); + ffi::ObjectPtr type = CopyTensorType(new_buffer); type->shape = new_buffer_shape; new_buffer = RebuildBufferVar(new_buffer, std::move(type)); return {new_buffer, buffer_indices}; diff --git a/src/s_tir/transform/memhammer_intermediate_stage.cc b/src/s_tir/transform/memhammer_intermediate_stage.cc index baf6082cf617..5fa34049cb14 100644 --- a/src/s_tir/transform/memhammer_intermediate_stage.cc +++ b/src/s_tir/transform/memhammer_intermediate_stage.cc @@ -234,14 +234,14 @@ class BufferLoadReplacer : public StmtExprMutator { * \param storage_scope the storage scope of the new cache * \param compute_location the compute location. * \param outer_loops the outer loops of this stmt - * \param alloc_buffer the new cache block + * \param alloc_tensor the new cache block * \return a pair. The first is the stmt after transformation. * The second is the SeqStmt that contains 2 stages (one original and another inserted). */ std::pair InsertCacheStage(Stmt stmt, bool is_write_cache, ffi::String storage_scope, ffi::Optional compute_location, const ffi::Array& outer_loops, - BufferVar* alloc_buffer) { + BufferVar* alloc_tensor) { Stmt body = stmt; std::vector loops; std::vector loops_under_compute_location; @@ -383,10 +383,10 @@ std::pair InsertCacheStage(Stmt stmt, bool is_write_cache, ffi::S } else { new_buffer = WithScope(buf_store->buffer, storage_scope); } - ffi::ObjectPtr buffer_type = CopyBufferType(new_buffer); + ffi::ObjectPtr buffer_type = CopyTensorType(new_buffer); buffer_type->shape = new_shape; new_buffer = RebuildBufferVar(new_buffer, std::move(buffer_type)); - *alloc_buffer = new_buffer; + *alloc_tensor = new_buffer; Stmt generate_body; if (is_write_cache) { @@ -472,7 +472,7 @@ Stmt CreateLocalStage::Rewrite(const Stmt& stmt, const ConstraintSet& constraint constraints.outer_loops, &cache_buffer) .first; if (cache_buffer.defined()) { - output->alloc_buffer.push_back(cache_buffer); + output->alloc_tensor.push_back(cache_buffer); } return after_caching; } diff --git a/src/s_tir/transform/memhammer_lower_auto_copy.cc b/src/s_tir/transform/memhammer_lower_auto_copy.cc index e448ffec417e..3f554480b35f 100644 --- a/src/s_tir/transform/memhammer_lower_auto_copy.cc +++ b/src/s_tir/transform/memhammer_lower_auto_copy.cc @@ -174,7 +174,7 @@ class AutoPadder { reverse_strides.push_back(stride); } // Step 3. create the new padded buffer - ffi::ObjectPtr b = CopyBufferType(buffer); + ffi::ObjectPtr b = CopyTensorType(buffer); ffi::Array strides; for (int i = static_cast(reverse_strides.size()) - 1; i >= 0; i--) { strides.push_back(reverse_strides[i]); @@ -765,7 +765,7 @@ class AutoCopyMutator : public StmtExprMutator { for (RewriteRule* rule : rules) { n->body = rule->Apply(std::move(n->body), constraints, &outputs); } - for (const BufferVar& buffer : outputs.alloc_buffer) { + for (const BufferVar& buffer : outputs.alloc_tensor) { n->alloc_buffers.push_back(buffer); } for (const auto& p : outputs.padding_min) { diff --git a/src/s_tir/transform/memhammer_rewrite_rule.h b/src/s_tir/transform/memhammer_rewrite_rule.h index 224376ba6f19..017fa32db6cf 100644 --- a/src/s_tir/transform/memhammer_rewrite_rule.h +++ b/src/s_tir/transform/memhammer_rewrite_rule.h @@ -76,7 +76,7 @@ struct ConstraintSet { /*! \brief The set containing all possible outputs of a rewrite rule */ struct OutputSet { /*! \brief New buffers allocated after rewrite */ - ffi::Array alloc_buffer; + ffi::Array alloc_tensor; /*! \brief The minimal padding size of a buffer in base 2 logarithm */ ffi::Map padding_min; }; @@ -246,14 +246,14 @@ class WmmaToShared : public RewriteRule { * \param storage_scope the storage scope of the new cache * \param compute_location the compute location. * \param outer_loops the outer loops of this stmt - * \param alloc_buffer the new cache block + * \param alloc_tensor the new cache block * \return a pair. The first is the stmt after transformation. * The second is the SeqStmt that contains 2 stages (one original and another inserted). */ std::pair InsertCacheStage(Stmt stmt, bool is_write_cache, ffi::String storage_scope, ffi::Optional compute_location, const ffi::Array& outer_loops, - BufferVar* alloc_buffer); + BufferVar* alloc_tensor); } // namespace s_tir } // namespace tvm diff --git a/src/s_tir/transform/memhammer_tensorcore_rewrite.cc b/src/s_tir/transform/memhammer_tensorcore_rewrite.cc index f20eb39e192d..932f8f656d37 100644 --- a/src/s_tir/transform/memhammer_tensorcore_rewrite.cc +++ b/src/s_tir/transform/memhammer_tensorcore_rewrite.cc @@ -143,7 +143,7 @@ Stmt RewriteWmmaLoad(Stmt stmt) { BufferVar tgt_buffer = buf_store->buffer; std::string layout = tgt_buffer.scope() == "wmma.matrix_a" ? "row_major" : "col_major"; BufferVar new_src_buffer( - /*name=*/"src", BufferType(/*storage_scope=*/src_buffer.scope(), + /*name=*/"src", TensorType(/*storage_scope=*/src_buffer.scope(), /*dtype=*/dtype, /*shape=*/{IntImm::Int32(16), IntImm::Int32(16)}, /*strides=*/{PrimVar("s1", int32_ty), PrimVar("s0", int32_ty)}, @@ -151,7 +151,7 @@ Stmt RewriteWmmaLoad(Stmt stmt) { /*data_alignment=*/64, /*offset_factor=*/16)); BufferVar new_tgt_buffer( - /*name=*/"tgt", BufferType(/*storage_scope=*/tgt_buffer.scope(), + /*name=*/"tgt", TensorType(/*storage_scope=*/tgt_buffer.scope(), /*dtype=*/dtype, /*shape=*/{IntImm::Int32(16), IntImm::Int32(16)}, /*strides=*/{}, @@ -253,10 +253,10 @@ Stmt RewriteWmmaStore(Stmt stmt) { const PrimType& dtype = dtype_ty; BufferVar new_src_buffer( - "src", BufferType(src_buffer.scope(), dtype, {IntImm::Int32(16), IntImm::Int32(16)}, {}, + "src", TensorType(src_buffer.scope(), dtype, {IntImm::Int32(16), IntImm::Int32(16)}, {}, PrimVar("src_elem_offset", int32_ty), 64, 16)); BufferVar new_tgt_buffer( - "tgt", BufferType(tgt_buffer.scope(), dtype, {IntImm::Int32(16), IntImm::Int32(16)}, + "tgt", TensorType(tgt_buffer.scope(), dtype, {IntImm::Int32(16), IntImm::Int32(16)}, {PrimVar("s1", int32_ty), PrimVar("s0", int32_ty)}, PrimVar("tgt_elem_offset", int32_ty), 64, 16)); @@ -356,7 +356,7 @@ Stmt WmmaToGlobal::Rewrite(const Stmt& stmt, const ConstraintSet& constraints, // Step 1. add a shared memory cache std::tie(body, seq) = InsertCacheStage(std::move(body), true, "shared.dyn", compute_location, constraints.outer_loops, &cache_buffer); - output->alloc_buffer.push_back(cache_buffer); + output->alloc_tensor.push_back(cache_buffer); output->padding_min.Set(cache_buffer, 8); // Step 2. do coalesced rewrite and tensor core rewrite respectively for 2 parts auto rewriter = ffi::make_object(seq.get(), constraints); @@ -480,10 +480,10 @@ Stmt RewriteMmaStore(Stmt stmt) { PrimType dtype_ty = src_buffer->dtype; const PrimType& dtype = dtype_ty; BufferVar new_src_buffer( - "src", BufferType(src_buffer.scope(), dtype, {IntImm::Int32(8), IntImm::Int32(8)}, {}, + "src", TensorType(src_buffer.scope(), dtype, {IntImm::Int32(8), IntImm::Int32(8)}, {}, PrimVar("src_elem_offset", int32_ty), 64, 8)); BufferVar new_tgt_buffer( - "tgt", BufferType(tgt_buffer.scope(), dtype, {IntImm::Int32(8), IntImm::Int32(8)}, + "tgt", TensorType(tgt_buffer.scope(), dtype, {IntImm::Int32(8), IntImm::Int32(8)}, {PrimVar("s1", int32_ty), PrimVar("s0", int32_ty)}, PrimVar("tgt_elem_offset", int32_ty), 64, 8)); @@ -573,7 +573,7 @@ Stmt MmaToGlobal::Rewrite(const Stmt& stmt, const ConstraintSet& constraints, // Step 1. add a shared memory cache std::tie(body, seq) = InsertCacheStage(std::move(body), true, "shared.dyn", compute_location, constraints.outer_loops, &cache_buffer); - output->alloc_buffer.push_back(cache_buffer); + output->alloc_tensor.push_back(cache_buffer); output->padding_min.Set(cache_buffer, 8); // Step 2. do coalesced rewrite and tensor core rewrite respectively for 2 parts auto rewriter = ffi::make_object(seq.get(), constraints); diff --git a/src/s_tir/transform/merge_shared_memory_allocations.cc b/src/s_tir/transform/merge_shared_memory_allocations.cc index 445eaaaa17c5..7cfb26e268c5 100644 --- a/src/s_tir/transform/merge_shared_memory_allocations.cc +++ b/src/s_tir/transform/merge_shared_memory_allocations.cc @@ -115,13 +115,13 @@ class AllocateCollector : public StmtExprVisitor { ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return DispatchAllocBuffer(op, call); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return DispatchAllocTensor(op, call); } return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchAllocBuffer(const BindNode* op, const CallNode* call) { + ffi::Optional DispatchAllocTensor(const BindNode* op, const CallNode* call) { ffi::String scope = call->args[2].as_or_throw()->value; StorageScope storage_scope = StorageScope::Create(scope); if (is_dynamic_ && storage_scope.rank == runtime::StorageRank::kShared && @@ -150,9 +150,9 @@ class AllocateCollector : public StmtExprVisitor { // before_scope -> scope_body -> after_scope // // This pass tries to detect last point that we need to keep memory -// alive under the same scope as AllocBuffer. -// The storage need to be kept alive between AllocBuffer and last access. -// The free point is only inserted at the same scope of AllocBuffer. +// alive under the same scope as AllocTensor. +// The storage need to be kept alive between AllocTensor and last access. +// The free point is only inserted at the same scope of AllocTensor. // class SharedMemLinearAccessPatternFinder final : public StmtExprVisitor { public: @@ -178,7 +178,7 @@ class SharedMemLinearAccessPatternFinder final : public StmtExprVisitor { BufferVar buffer; }; - ffi::Optional DispatchAllocBuffer(const BindNode* op) { + ffi::Optional DispatchAllocTensor(const BindNode* op) { size_t level = scope_.size(); const VarNode* buf = op->var.as_or_throw().get(); alloc_info_[buf].buffer = op->var.as_or_throw(); @@ -226,17 +226,17 @@ class SharedMemLinearAccessPatternFinder final : public StmtExprVisitor { ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return DispatchAllocBuffer(op); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return DispatchAllocTensor(op); } if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::decl_buffer())) { - return DispatchDeclBuffer(op, call); + call && call->op.same_as(tirx::builtin::decl_tensor())) { + return DispatchDeclTensor(op, call); } return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchDeclBuffer(const BindNode* op, const CallNode* call) { + ffi::Optional DispatchDeclTensor(const BindNode* op, const CallNode* call) { if (auto source = GetBufferDataVar(call->args[0])) { const VarNode* allocation = ResolveAlias(source.value().get()); if (alloc_info_.count(allocation)) { @@ -435,7 +435,7 @@ class SharedMemoryRewriter : public StmtExprMutator { * Same name string is fine — Var identity is by pointer, not name. */ BufferVar MakeMergedBuffer(PrimExpr size) { - return decl_buffer({std::move(size)}, PrimType::UInt(8), + return decl_tensor({std::move(size)}, PrimType::UInt(8), is_dynamic_ ? "buf_dyn_shmem" : "buf_shmem", is_dynamic_ ? "shared.dyn" : "shared"); } @@ -481,7 +481,7 @@ class SharedMemoryRewriter : public StmtExprMutator { // The uint8 merged allocation intentionally supplies storage for // typed views; target codegen emits the required pointer cast. visited_body = SeqStmt::Flatten( - Bind(remapped, Call(remapped.type(), tirx::builtin::decl_buffer(), + Bind(remapped, Call(remapped.type(), tirx::builtin::decl_tensor(), {scope.merged_buffer.data(), tvm::Tuple(remapped->shape), DataTypeImm(remapped->dtype->dtype), StringImm(remapped.scope())}, {})), @@ -495,13 +495,13 @@ class SharedMemoryRewriter : public StmtExprMutator { return AttrStmt(op->node, op->attr_key, op->value, visited_body, op->span); } - // 7. Wrap with the merged-buffer AllocBuffer. + // 7. Wrap with the merged-buffer AllocTensor. ffi::Map annotations; if (scope.has_volatile_alloc) { annotations.Set(tirx::attr::kVolatile, true); } Stmt alloc_stmt = Bind(scope.merged_buffer.var(), - Call(scope.merged_buffer.type(), tirx::builtin::alloc_buffer(), + Call(scope.merged_buffer.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(scope.merged_buffer->shape), DataTypeImm(scope.merged_buffer->dtype->dtype), StringImm(scope.merged_buffer.scope())}, @@ -516,7 +516,7 @@ class SharedMemoryRewriter : public StmtExprMutator { return StmtExprMutator::Mutate_(op, inplace_mode); } - UnchangedOr Mutate_AllocBuffer(const BindNode* op, const CallNode* call, + UnchangedOr Mutate_AllocTensor(const BindNode* op, const CallNode* call, InplaceMode inplace_mode) { ffi::String scope = call->args[2].as_or_throw()->value; StorageScope storage_scope = StorageScope::Create(scope); @@ -539,17 +539,17 @@ class SharedMemoryRewriter : public StmtExprMutator { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { - return Mutate_AllocBuffer(op, call, inplace_mode); + call && call->op.same_as(tirx::builtin::alloc_tensor())) { + return Mutate_AllocTensor(op, call, inplace_mode); } if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::decl_buffer())) { - return Mutate_DeclBuffer(op, call, inplace_mode); + call && call->op.same_as(tirx::builtin::decl_tensor())) { + return Mutate_DeclTensor(op, call, inplace_mode); } return StmtExprMutator::Mutate_(op, inplace_mode); } - UnchangedOr Mutate_DeclBuffer(const BindNode* op, const CallNode* call, + UnchangedOr Mutate_DeclTensor(const BindNode* op, const CallNode* call, InplaceMode inplace_mode) { if (!scope_stack_.empty()) { if (auto source = GetBufferDataVar(call->args[0])) { @@ -568,7 +568,7 @@ class SharedMemoryRewriter : public StmtExprMutator { !new_buf.same_as(node->var)) { const auto* new_call = node->value.as(); return Bind(new_buf, - Call(new_buf.type(), tirx::builtin::decl_buffer(), + Call(new_buf.type(), tirx::builtin::decl_tensor(), {new_call->args[0], tvm::Tuple(new_buf->shape), DataTypeImm(new_buf->dtype->dtype), StringImm(new_buf.scope())}, new_call->attrs, new_call->ty_args, new_call->span), @@ -641,7 +641,7 @@ class SharedMemoryRewriter : public StmtExprMutator { << "MergeSharedMemoryAllocations expects flat memory buffers, " << "and is to be run after " << "FlattenBuffer"; - buffer = RebuildBufferVar(buffer, CopyBufferType(buffer)); + buffer = RebuildBufferVar(buffer, CopyTensorType(buffer)); } scope.buffer_remap[key] = buffer; @@ -659,7 +659,7 @@ class SharedMemoryRewriter : public StmtExprMutator { return StmtExprMutator::Mutate_(op, inplace_mode); } Var buffer = buffer_opt.value(); - bool is_shared = buffer->ty.as() + bool is_shared = buffer->ty.as() ? IsAppropriateSharedMemory(buffer.as_or_throw()) : IsAppropriateSharedMemory(buffer); if (!is_shared || scope_stack_.empty() || @@ -667,7 +667,7 @@ class SharedMemoryRewriter : public StmtExprMutator { return StmtExprMutator::Mutate_(op, inplace_mode); } PrimExpr extra_offset = GetBufferOffset(buffer, dtype); - Expr merged_data = buffer->ty.as() + Expr merged_data = buffer->ty.as() ? GetUpdatedBuffer(buffer.as_or_throw()).data() : scope_stack_.back().merged_buffer.data(); @@ -684,7 +684,7 @@ class SharedMemoryRewriter : public StmtExprMutator { Var buffer = buffer_opt.value(); DLDataType dtype; bool is_shared; - if (buffer->ty.as()) { + if (buffer->ty.as()) { BufferVar typed_buffer = buffer.as_or_throw(); dtype = typed_buffer->dtype->dtype; is_shared = IsAppropriateSharedMemory(typed_buffer); @@ -702,7 +702,7 @@ class SharedMemoryRewriter : public StmtExprMutator { } PrimExpr extra_offset = GetBufferOffset(buffer, dtype); PrimExpr offset = Mutate(op->args[1]).ValueOrUnchanged(op->args[1]).as_or_throw(); - if (buffer->ty.as()) { + if (buffer->ty.as()) { Expr merged_data = GetUpdatedBuffer(buffer.as_or_throw()).data(); ffi::Array args = op->args; args.Set(0, merged_data); diff --git a/src/s_tir/transform/storage_access.cc b/src/s_tir/transform/storage_access.cc index 1bcfc7befd12..989e1a72fe08 100644 --- a/src/s_tir/transform/storage_access.cc +++ b/src/s_tir/transform/storage_access.cc @@ -116,7 +116,7 @@ ffi::Optional StorageAccessVisitor::Visit_(const EvaluateNode* o ffi::Optional StorageAccessVisitor::Visit_(const BindNode* op) { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::decl_buffer())) { + call && call->op.same_as(tirx::builtin::decl_tensor())) { if (auto source = GetBufferDataVar(call->args[0])) { buffer_aliases_.insert_or_assign(op->var.as_or_throw().get(), ResolveBuffer(source.value())); @@ -124,7 +124,7 @@ ffi::Optional StorageAccessVisitor::Visit_(const BindNode* op) { return StmtExprVisitor::Visit_(op); } if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) + call && call->op.same_as(tirx::builtin::alloc_tensor())) return StmtExprVisitor::Visit_(op); allow_append_ = true; TVM_FFI_ICHECK_EQ(curr_stmt_.access.size(), 0U); @@ -358,7 +358,7 @@ ffi::Optional StorageAccessVisitor::Visit_(const CallNode* op) { } StorageScope StorageAccessVisitor::GetScope(Var buffer_var) const { - if (auto buffer_type = buffer_var->ty.as()) { + if (auto buffer_type = buffer_var->ty.as()) { return StorageScope::Create(buffer_type.value()->storage_scope); } if (buffer_var->ty.as()) { diff --git a/src/s_tir/transform/tensorcore_infer_fragment.cc b/src/s_tir/transform/tensorcore_infer_fragment.cc index a2faffa6c869..785d0178b26b 100644 --- a/src/s_tir/transform/tensorcore_infer_fragment.cc +++ b/src/s_tir/transform/tensorcore_infer_fragment.cc @@ -208,7 +208,7 @@ class InferFragmenter : public s_tir::StmtExprMutator { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); - call && call->op.same_as(tirx::builtin::alloc_buffer())) { + call && call->op.same_as(tirx::builtin::alloc_tensor())) { auto it = fragment_getter.fragments.find(op->var.get()); if (it == fragment_getter.fragments.end()) return ffi::Unchanged(); const FragmentInfo& info = it->second; diff --git a/src/s_tir/transform/transform_mma_buffer_layout.cc b/src/s_tir/transform/transform_mma_buffer_layout.cc index 2838e9089d46..7f53c0d3527f 100644 --- a/src/s_tir/transform/transform_mma_buffer_layout.cc +++ b/src/s_tir/transform/transform_mma_buffer_layout.cc @@ -76,7 +76,7 @@ class MmaBufferLayoutTransformer : public StmtExprMutator { {IntImm::Int32(dim0->value / 16), IntImm::Int32(dim1->value / 8), 2, 2}); BufferVar new_buffer = - decl_buffer(std::move(new_shape), buffer->dtype, buffer.name(), "local"); + decl_tensor(std::move(new_shape), buffer->dtype, buffer.name(), "local"); VarRemapSet(buffer, new_buffer); return new_buffer; @@ -97,7 +97,7 @@ class MmaBufferLayoutTransformer : public StmtExprMutator { {IntImm::Int32(dim0->value / 32), IntImm::Int32(dim1->value / 8), 4, 2}); BufferVar new_buffer = - decl_buffer(std::move(new_shape), buffer->dtype, buffer.name(), "local"); + decl_tensor(std::move(new_shape), buffer->dtype, buffer.name(), "local"); VarRemapSet(buffer, new_buffer); return new_buffer; @@ -118,7 +118,7 @@ class MmaBufferLayoutTransformer : public StmtExprMutator { {IntImm::Int32(dim0->value / 8), IntImm::Int32(dim1->value / 32), 1, 8}); BufferVar new_buffer = - decl_buffer(std::move(new_shape), buffer->dtype, buffer.name(), "local"); + decl_tensor(std::move(new_shape), buffer->dtype, buffer.name(), "local"); VarRemapSet(buffer, new_buffer); return new_buffer; } diff --git a/src/script/printer/script_printer.cc b/src/script/printer/script_printer.cc index a2fb3728e330..6a712759318f 100644 --- a/src/script/printer/script_printer.cc +++ b/src/script/printer/script_printer.cc @@ -173,7 +173,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { RegisterScriptRepr(); RegisterScriptRepr(); RegisterScriptRepr(); - RegisterScriptRepr(); + RegisterScriptRepr(); RegisterScriptRepr(); RegisterScriptRepr(); RegisterScriptRepr(); diff --git a/src/target/llvm/codegen_cpu.cc b/src/target/llvm/codegen_cpu.cc index 4c779daece71..fc26ab70e7b0 100644 --- a/src/target/llvm/codegen_cpu.cc +++ b/src/target/llvm/codegen_cpu.cc @@ -574,7 +574,7 @@ void CodeGenCPU::CreateComputeScope(const AttrStmtNode* op) { llvm::Argument* v = &(*it); const Var& var = vargs[idx]; var_map_[var.get()] = v; - if ((var->ty.as() || var->ty.as()) && + if ((var->ty.as() || var->ty.as()) && !alias_var_set_.count(var.get())) { // set non alias. fcompute->addParamAttr(idx, llvm::Attribute::NoAlias); @@ -594,7 +594,7 @@ void CodeGenCPU::CreateComputeScope(const AttrStmtNode* op) { function_ = fcompute; ffi::Array debug_param_types = vargs.Map([](const Var& var) -> Type { - if (const auto* buffer_type = var->ty.as()) { + if (const auto* buffer_type = var->ty.as()) { // Compute-scope captures use their physical LLVM pointer values. return buffer_type->DataPointerType(); } diff --git a/src/target/llvm/codegen_llvm.cc b/src/target/llvm/codegen_llvm.cc index 00c7cbc9bc43..d337b20d36dc 100644 --- a/src/target/llvm/codegen_llvm.cc +++ b/src/target/llvm/codegen_llvm.cc @@ -272,8 +272,8 @@ llvm::Function* CodeGenLLVM::DeclareFunctionInternal(const GlobalVar& gvar, cons } for (const Var& param : func->params) { - TVM_FFI_ICHECK(!param->ty.as()) - << "Cannot codegen BufferType-annotated parameter " << param << "; please lower it first"; + TVM_FFI_ICHECK(!param->ty.as()) + << "Cannot codegen TensorType-annotated parameter " << param << "; please lower it first"; } std::vector param_types; @@ -2208,7 +2208,7 @@ void CodeGenLLVM::Dispatch_(const IfThenElseNode* op) { builder_->SetInsertPoint(end_block); } -void CodeGenLLVM::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenLLVM::DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call) { tvm::Tuple shape = buffer_call->args[0].as_or_throw(); DLDataType dtype = buffer_call->args[1].as_or_throw()->value; BufferVar buffer = op->var.as_or_throw(); @@ -2289,8 +2289,8 @@ void CodeGenLLVM::Dispatch_(const AssertStmtNode* op) { void CodeGenLLVM::Dispatch_(const BindNode* op) { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(tirx::builtin::alloc_buffer())) return DispatchAllocBuffer(op, call); - if (call->op.same_as(tirx::builtin::decl_buffer())) return DispatchDeclBuffer(op, call); + if (call->op.same_as(tirx::builtin::alloc_tensor())) return DispatchAllocTensor(op, call); + if (call->op.same_as(tirx::builtin::decl_tensor())) return DispatchDeclTensor(op, call); } EmitDebugLocation(op); const VarNode* v = op->var.get(); @@ -2337,7 +2337,7 @@ void CodeGenLLVM::Dispatch_(const SeqStmtNode* op) { } } -void CodeGenLLVM::DispatchDeclBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenLLVM::DispatchDeclTensor(const BindNode* op, const CallNode* buffer_call) { Expr data = buffer_call->args[0]; DLDataType dtype = buffer_call->args[2].as_or_throw()->value; ffi::String scope = buffer_call->args[3].as_or_throw()->value; @@ -2463,9 +2463,9 @@ void CodeGenLLVM::AddDebugInformation(llvm::Value* llvm_value, const Var& tir_va if (!di_subprogram_) return; Type debug_type = tir_var->ty; - if (const auto* buffer_type = debug_type.as()) { + if (const auto* buffer_type = debug_type.as()) { // A BufferVar is a compiler-side identity. Its LLVM value is the physical - // data pointer installed by AllocBuffer or DeclBuffer. + // data pointer installed by AllocTensor or DeclTensor. debug_type = buffer_type->DataPointerType(); } auto dbg_dtype = GetDebugType(debug_type); diff --git a/src/target/llvm/codegen_llvm.h b/src/target/llvm/codegen_llvm.h index b2687a2c07eb..d09423420eb1 100644 --- a/src/target/llvm/codegen_llvm.h +++ b/src/target/llvm/codegen_llvm.h @@ -244,13 +244,13 @@ class CodeGenLLVM : public tirx::ExprFunctor, void Dispatch_(const BreakNode* op) override; void Dispatch_(const ContinueNode* op) override; void Dispatch_(const IfThenElseNode* op) override; - void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call); + void DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call); void Dispatch_(const AttrStmtNode* op) override; void Dispatch_(const AssertStmtNode* op) override; void Dispatch_(const BindNode* op) override; void Dispatch_(const SeqStmtNode* op) override; void Dispatch_(const EvaluateNode* op) override; - void DispatchDeclBuffer(const BindNode* op, const CallNode* buffer_call); + void DispatchDeclTensor(const BindNode* op, const CallNode* buffer_call); // Get constant string llvm::Constant* GetConstString(const std::string& str); @@ -563,7 +563,7 @@ class CodeGenLLVM : public tirx::ExprFunctor, std::unordered_map alloc_storage_info_; // The definition of local variable. std::unordered_map var_map_; - // Canonical physical storage identity for DeclBuffer aliases. + // Canonical physical storage identity for DeclTensor aliases. std::unordered_map buffer_physical_root_; // global strings std::unordered_map str_map_; diff --git a/src/target/source/codegen_c.cc b/src/target/source/codegen_c.cc index 6b5a02c96ffd..c0535ff552b9 100644 --- a/src/target/source/codegen_c.cc +++ b/src/target/source/codegen_c.cc @@ -711,8 +711,8 @@ void CodeGenC::Dispatch_(const CallNode* op, std::ostream& os) { // NOLINT(*) if (op->op.same_as(tirx::builtin::buffer_data())) { TVM_FFI_ICHECK_EQ(op->args.size(), 1U); const auto* buffer = op->args[0].as(); - TVM_FFI_ICHECK(buffer && buffer->ty.as()) - << "buffer_data expects a Var with BufferType"; + TVM_FFI_ICHECK(buffer && buffer->ty.as()) + << "buffer_data expects a Var with TensorType"; os << GetVarID(buffer); } else if (op->op.same_as(builtin_call_extern_) || op->op.same_as(builtin_call_pure_extern_)) { TVM_FFI_ICHECK_GE(op->args.size(), 1U); @@ -924,7 +924,7 @@ void CodeGenC::PrintVecBinaryOp(const std::string& op, const PrimType& t, PrimEx } } -void CodeGenC::DispatchDeclBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenC::DispatchDeclTensor(const BindNode* op, const CallNode* buffer_call) { Expr data = buffer_call->args[0]; DLDataType dtype = buffer_call->args[2].as_or_throw()->value; ffi::String scope = buffer_call->args[3].as_or_throw()->value; @@ -1231,8 +1231,8 @@ void CodeGenC::Dispatch_(const prim::SelectNode* op, std::ostream& os) { // NOL void CodeGenC::Dispatch_(const BindNode* op) { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(tirx::builtin::alloc_buffer())) return DispatchAllocBuffer(op, call); - if (call->op.same_as(tirx::builtin::decl_buffer())) return DispatchDeclBuffer(op, call); + if (call->op.same_as(tirx::builtin::alloc_tensor())) return DispatchAllocTensor(op, call); + if (call->op.same_as(tirx::builtin::decl_tensor())) return DispatchDeclTensor(op, call); } RegisterHandleTypeFromPointer(op->var, &op->value); std::string value = PrintExpr(op->value); @@ -1254,7 +1254,7 @@ void CodeGenC::Dispatch_(const BindNode* op) { } } -void CodeGenC::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) { +void CodeGenC::DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call) { tvm::Tuple shape = buffer_call->args[0].as_or_throw(); DLDataType dtype = buffer_call->args[1].as_or_throw()->value; ffi::String scope = buffer_call->args[2].as_or_throw()->value; diff --git a/src/target/source/codegen_c.h b/src/target/source/codegen_c.h index 7d89e5f831f2..6c5d40a8a3e9 100644 --- a/src/target/source/codegen_c.h +++ b/src/target/source/codegen_c.h @@ -211,12 +211,12 @@ class CodeGenC : public tirx::ExprFunctor, void Dispatch_(const BreakNode* op) override; void Dispatch_(const ContinueNode* op) override; void Dispatch_(const IfThenElseNode* op) override; - void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call); + void DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call); void Dispatch_(const AttrStmtNode* op) override; void Dispatch_(const AssertStmtNode* op) override; void Dispatch_(const EvaluateNode* op) override; void Dispatch_(const SeqStmtNode* op) override; - void DispatchDeclBuffer(const BindNode* op, const CallNode* buffer_call); + void DispatchDeclTensor(const BindNode* op, const CallNode* buffer_call); /*! * \brief Print expr representing the thread tag diff --git a/src/te/operation/create_primfunc.cc b/src/te/operation/create_primfunc.cc index 8d4bfb556344..c20b58639673 100644 --- a/src/te/operation/create_primfunc.cc +++ b/src/te/operation/create_primfunc.cc @@ -332,10 +332,10 @@ ffi::Array GenerateOutputBuffers(const te::ComputeOp& compute_op, Cre // Step 2. Prepare buffers for compute outputs // - Declare buffers // - Update `op2buffers` - // - Add the non-argument tensors to `alloc_buffer` of the root block + // - Add the non-argument tensors to `alloc_tensor` of the root block ffi::Array buffers; for (const te::Tensor& tensor : tensors) { - BufferVar buffer = decl_buffer(tensor->shape, tensor->dtype, tensor->GetNameHint(), "global"); + BufferVar buffer = decl_tensor(tensor->shape, tensor->dtype, tensor->GetNameHint(), "global"); info->tensor2buffers[tensor] = buffer; buffers.push_back(buffer); if (!info->IsArg(tensor)) { @@ -727,7 +727,7 @@ Stmt GenerateStmtFromExternOp(const te::ExternOp& extern_op, CreateFuncInfo* inf if (auto offset_var = placeholder->elem_offset.as()) { var_map[offset_var.value().get()] = zero_offset; } - ffi::ObjectPtr type = CopyBufferType(output_buffer); + ffi::ObjectPtr type = CopyTensorType(output_buffer); type->elem_offset = zero_offset; output_buffer = RebuildBufferVar(output_buffer, std::move(type)); input_buffer_map[placeholder.get()] = output_buffer; @@ -815,7 +815,7 @@ void RewriteStageToBlock(const te::Operation& op, CreateFuncInfo* info, // buffer declaration recorded in the tensor2buffer binds map if (info->tensor2buffers.count(tensor) == 0) { const BufferVar& buffer = - decl_buffer(placeholder->shape, placeholder->dtype, placeholder->name, "global"); + decl_tensor(placeholder->shape, placeholder->dtype, placeholder->name, "global"); info->tensor2buffers[tensor] = buffer; } } else if (auto compute_op = op.as()) { diff --git a/src/tirx/analysis/var_use_def_analysis.cc b/src/tirx/analysis/var_use_def_analysis.cc index 80e24ceee634..858c80d292c5 100644 --- a/src/tirx/analysis/var_use_def_analysis.cc +++ b/src/tirx/analysis/var_use_def_analysis.cc @@ -82,7 +82,7 @@ ffi::Optional VarUseDefAnalyzer::Visit_(const prim::LetNode* op) ffi::Optional VarUseDefAnalyzer::Visit_(const VarNode* op) { Var var = ffi::GetRef(op); - if (var->ty.as()) { + if (var->ty.as()) { BufferVar buffer = var.as_or_throw(); if (def_region_kind() == kTVMFFIDefRegionKindSimple) { bool is_first_buffer_definition = !buffer_def_count_.count(op); @@ -127,7 +127,7 @@ void VarUseDefAnalyzer::HandleUse(const Var& var) { void VarUseDefAnalyzer::HandleDef(const BufferVar& buf) { auto ptr = buf.get(); - // Some lowering pipelines may duplicate identical DeclBuffer nodes that + // Some lowering pipelines may duplicate identical DeclTensor nodes that // reference the same BufferVar object. Treat repeated definition of the same // buffer object as idempotent. if (buffer_def_count_.count(ptr)) { diff --git a/src/tirx/analysis/verify_memory.cc b/src/tirx/analysis/verify_memory.cc index 5cf8a23d23d6..790e693df24f 100644 --- a/src/tirx/analysis/verify_memory.cc +++ b/src/tirx/analysis/verify_memory.cc @@ -111,7 +111,7 @@ class MemoryAccessVerifier final : public StmtExprVisitor { bool IsFromFunctionArgs(const VarNode* var) const { const VarNode* V = var; for (const Var& param : func_->params) { - if (param->ty.as() && V == param.get()) return true; + if (param->ty.as() && V == param.get()) return true; } while (true) { diff --git a/src/tirx/analysis/verify_well_formed.h b/src/tirx/analysis/verify_well_formed.h index e78d143e9394..d097e2c77a50 100644 --- a/src/tirx/analysis/verify_well_formed.h +++ b/src/tirx/analysis/verify_well_formed.h @@ -121,8 +121,8 @@ class UndefinedVarVerifier : public Verifier, /*! \brief Verify that buffers with a declaration are not used outside their declared scope. * * When a buffer is declared via one of the following sites: - * - BufferType-annotated PrimFunc parameters - * - DeclBuffer statement + * - TensorType-annotated PrimFunc parameters + * - DeclTensor statement * - Dialect-specific definitions exposed by PathVisitor * * it must not appear in a BufferLoad, BufferStore, or BufferRegion outside that declaration's @@ -155,7 +155,7 @@ class UndefinedBufferVerifier : public Verifierty.as()) return; + if (!var->ty.as()) return; auto buffer = var.as_or_throw(); auto active_def = currently_defined_.find(buffer); if (active_def != currently_defined_.end()) { @@ -177,7 +177,7 @@ class UndefinedBufferVerifier : public Verifier, ffi::Optional, ffi::Optional>, PrimExpr>>; -BufferVar RebuildBufferVarFromType(const BufferVar& buffer, BufferType type, +BufferVar RebuildBufferVarFromType(const BufferVar& buffer, TensorType type, ffi::String name_suffix = "") { return BufferVar(buffer.name() + name_suffix, std::move(type), buffer.span()); } @@ -61,7 +61,7 @@ ffi::ObjectRef RealizeBufferSubscript( slice, Span span) { BufferVar buffer = value.as_or_throw(); - BufferType buffer_ty = buffer.type(); + TensorType buffer_ty = buffer.type(); TVM_FFI_CHECK_LE(slice.size(), buffer_ty->shape.size(), IndexError) << "Too many indices for a " << buffer_ty->shape.size() << "-dimensional buffer"; @@ -203,7 +203,7 @@ TensorRegion BufferRegionFromPoint(BufferVar buffer, ffi::Array indice TVM_FFI_STATIC_INIT_BLOCK() { namespace refl = tvm::ffi::reflection; - refl::TypeAttrDef().def("__subscript_expr_realize__", RealizeBufferSubscript); + refl::TypeAttrDef().def("__subscript_expr_realize__", RealizeBufferSubscript); refl::TypeAttrDef().def("__subscript_expr_realize__", RealizeBufferRegionSubscript); } @@ -215,9 +215,9 @@ ffi::Array SimplifyArray(sym::AnalyzerObj* ana, ffi::Array a return array; } -BufferVar decl_buffer(ffi::Array shape, PrimType dtype, ffi::String name, +BufferVar decl_tensor(ffi::Array shape, PrimType dtype, ffi::String name, ffi::String storage_scope, Span span) { - return BufferVar(name, BufferType(storage_scope, dtype, shape, {}, PrimExpr(), 0, 0), span); + return BufferVar(name, TensorType(storage_scope, dtype, shape, {}, PrimExpr(), 0, 0), span); } // Split the given expression w.r.t the add operator @@ -419,10 +419,10 @@ ffi::Array BufferVar::OffsetOf(ffi::Array input_indices) con // The buffer offset in convention of number of elements of // original data ignoring number of lanes. // We also perform optimization to simplify the indexing expression. -ffi::Array BufferTypeNode::ElemOffset(ffi::Array input_indices, +ffi::Array TensorTypeNode::ElemOffset(ffi::Array input_indices, bool inner) const { TVM_FFI_ICHECK_EQ(shape.size(), input_indices.size()) - << "BufferType is " << shape.size() << "-dimensional, cannot be indexed with the " + << "TensorType is " << shape.size() << "-dimensional, cannot be indexed with the " << input_indices.size() << "-dimensional indices provided."; if (strides.size()) { @@ -453,7 +453,7 @@ ffi::Array BufferTypeNode::ElemOffset(ffi::Array input_indic return SimplifyArray(ana.get(), {output_index}); } -inline ffi::Array BufferOffset(const BufferTypeNode* n, ffi::Array index, +inline ffi::Array BufferOffset(const TensorTypeNode* n, ffi::Array index, PrimType dtype) { ffi::Array offsets = n->ElemOffset(index); // If the BufferVar has element type with more than one lane, scale to @@ -501,14 +501,14 @@ BufferVar BufferVar::GetFlattenedBuffer() const { // (see test_tir_transform_flatten_buffer). Reset to the default layout // for the new shape so the buffer stays internally consistent. return RebuildBufferVarFromType( - *this, BufferType(self->storage_scope, self->dtype, output_shape, {}, self->elem_offset, + *this, TensorType(self->storage_scope, self->dtype, output_shape, {}, self->elem_offset, self->data_alignment, self->offset_factor, TileLayoutNode::DefaultLayout(output_shape), self->allocated_addr)); } } PrimExpr BufferVar::vload(ffi::Array begin, PrimType value_dtype) const { - const BufferTypeNode* n = operator->(); + const TensorTypeNode* n = operator->(); TVM_FFI_ICHECK(n != nullptr); PrimType buffer_dtype(n->dtype); int value_lanes = @@ -532,7 +532,7 @@ PrimExpr BufferVar::vload(ffi::Array begin, PrimType value_dtype) cons } Stmt BufferVar::vstore(ffi::Array begin, PrimExpr value) const { - const BufferTypeNode* n = operator->(); + const TensorTypeNode* n = operator->(); TVM_FFI_ICHECK(n != nullptr); PrimType value_dtype = value.ty(); PrimType buffer_dtype(n->dtype); @@ -561,7 +561,7 @@ ffi::String BufferVar::scope() const { return (*this)->storage_scope; } BufferVar BufferVar::MakeStrideView() const { if ((*this)->strides.size() != 0) return *this; if ((*this)->shape.size() == 0) return *this; - const BufferTypeNode* self = operator->(); + const TensorTypeNode* self = operator->(); TVM_FFI_ICHECK(self != nullptr); PrimExpr acc = IntImm(PrimType(self->DefaultIndexType()), 1); std::vector temp; @@ -574,13 +574,13 @@ BufferVar BufferVar::MakeStrideView() const { strides.push_back(temp[i - 1]); } return RebuildBufferVarFromType( - *this, BufferType(self->storage_scope, self->dtype, self->shape, std::move(strides), + *this, TensorType(self->storage_scope, self->dtype, self->shape, std::move(strides), self->elem_offset, self->data_alignment, self->offset_factor, self->layout, self->allocated_addr)); } BufferVar BufferVar::MakeSlice(ffi::Array begins, ffi::Array extents) const { - const BufferTypeNode* n = operator->(); + const TensorTypeNode* n = operator->(); TVM_FFI_ICHECK(n != nullptr); sym::Analyzer ana; begins = SimplifyArray(ana.get(), begins); @@ -607,14 +607,14 @@ BufferVar BufferVar::MakeSlice(ffi::Array begins, ffi::Array } return RebuildBufferVarFromType( *this, - BufferType(n->storage_scope, n->dtype, extents, strides, elem_offset[0], n->data_alignment, 0, + TensorType(n->storage_scope, n->dtype, extents, strides, elem_offset[0], n->data_alignment, 0, TileLayoutNode::DefaultLayout(extents)), "_slice"); } Expr BufferVar::access_ptr(int access_mask, PointerType ptr_type, int content_lanes, PrimExpr offset, ffi::Optional input_extent) const { - const BufferTypeNode* self = operator->(); + const TensorTypeNode* self = operator->(); TVM_FFI_ICHECK(self != nullptr); // An access pointer addresses the same allocation as the buffer data. The // requested type controls its pointee, while the buffer controls its address @@ -648,7 +648,7 @@ Expr BufferVar::access_ptr(int access_mask, PointerType ptr_type, int content_la return Call(ptr_type, tirx::builtin::tvm_access_ptr(), acc_args); } -BufferVar::BufferVar(ffi::String name, BufferType type, Span span) +BufferVar::BufferVar(ffi::String name, TensorType type, Span span) : Var(Var(std::move(name), std::move(type), std::move(span))) {} Expr BufferVar::data() const { return Call(DataPointerType(), builtin::buffer_data(), {var()}); } @@ -664,13 +664,13 @@ tirx::BufferVar BufferWithOffsetAlignment(ffi::Array shape, PrimType d } return tirx::BufferVar( - name, BufferType(memory_scope, dtype, shape, {}, elem_offset, data_alignment, offset_factor)); + name, TensorType(memory_scope, dtype, shape, {}, elem_offset, data_alignment, offset_factor)); } BufferVar BufferVar::with_allocated_addr(ffi::Array allocated_addr) const { const auto* self = operator->(); return RebuildBufferVarFromType( - *this, BufferType(self->storage_scope, self->dtype, self->shape, self->strides, + *this, TensorType(self->storage_scope, self->dtype, self->shape, self->strides, self->elem_offset, self->data_alignment, self->offset_factor, self->layout, std::move(allocated_addr))); } @@ -678,7 +678,7 @@ BufferVar BufferVar::with_allocated_addr(ffi::Array allocated_addr) co BufferVar BufferVar::with_dtype(PrimType dtype) const { const auto* self = operator->(); return RebuildBufferVarFromType( - *this, BufferType(self->storage_scope, std::move(dtype), self->shape, self->strides, + *this, TensorType(self->storage_scope, std::move(dtype), self->shape, self->strides, self->elem_offset, self->data_alignment, self->offset_factor, self->layout, self->allocated_addr)); } @@ -694,7 +694,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { namespace refl = tvm::ffi::reflection; refl::GlobalDef() .def("tirx.BufferVar", - [](ffi::String name, BufferType type, Span span) { + [](ffi::String name, TensorType type, Span span) { return BufferVar(std::move(name), std::move(type), std::move(span)); }) .def_method( diff --git a/src/tirx/ir/data_type_rewriter.cc b/src/tirx/ir/data_type_rewriter.cc index 7c8422d310d6..68b55cf1bba1 100644 --- a/src/tirx/ir/data_type_rewriter.cc +++ b/src/tirx/ir/data_type_rewriter.cc @@ -359,7 +359,7 @@ UnchangedOr IndexDataTypeRewriter::Mutate_(const AttrStmtNode* op, Inplace UnchangedOr IndexDataTypeRewriter::Mutate(ffi::AnyView value, InplaceMode inplace_mode) { bool is_enabled = is_enabled_; - if (value.as()) is_enabled_ = true; + if (value.as()) is_enabled_ = true; auto result = DataTypeLegalizer::Mutate(value, inplace_mode); is_enabled_ = is_enabled; return result; diff --git a/src/tirx/ir/function.cc b/src/tirx/ir/function.cc index 662ca9918170..b46dbb876c75 100644 --- a/src/tirx/ir/function.cc +++ b/src/tirx/ir/function.cc @@ -41,7 +41,7 @@ tvm::Type InferType(const PrimFunc& prim_func) { ffi::Array params; for (const auto& param : prim_func->params) { tvm::Type param_ty = [&]() -> tvm::Type { - if (param->ty.as()) { + if (param->ty.as()) { BufferVar buf = param.as_or_throw(); relax::ShapeExpr shape( buf->shape.Map([](PrimExpr dim) { return cast(PrimType::Int(64), dim); })); diff --git a/src/tirx/ir/ir_mutator_with_analyzer.cc b/src/tirx/ir/ir_mutator_with_analyzer.cc index b9eccbdf947a..86f174eadbfc 100644 --- a/src/tirx/ir/ir_mutator_with_analyzer.cc +++ b/src/tirx/ir/ir_mutator_with_analyzer.cc @@ -52,7 +52,7 @@ using sym::detail::EnterConstraintFacts; void IRMutatorWithAnalyzer::MarkBufferParamShapes(const tirx::PrimFunc& func) { // Mark all symbolic buffer-parameter shape values as positive. for (const tirx::Var& param : func->params) { - if (!param->ty.as()) { + if (!param->ty.as()) { continue; } tirx::BufferVar buffer = param.as_or_throw(); diff --git a/src/tirx/ir/specialize.cc b/src/tirx/ir/specialize.cc index fddfa2b256ea..45840920d489 100644 --- a/src/tirx/ir/specialize.cc +++ b/src/tirx/ir/specialize.cc @@ -158,10 +158,10 @@ class PrimFuncSpecializer : public StmtExprMutator { private: ffi::Optional Visit_(const VarNode* op) final { - if (op->ty.as()) { + if (op->ty.as()) { if (def_region_kind() == kTVMFFIDefRegionKindSimple) { const BufferVar buffer = GetBufferVar(op); - specializer_->MutateAllocBuffer(buffer); + specializer_->MutateAllocTensor(buffer); } else { specializer_->ValidateBufferUse(GetBufferVar(op)); } @@ -172,10 +172,10 @@ class PrimFuncSpecializer : public StmtExprMutator { ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); call && - (call->op.same_as(builtin::alloc_buffer()) || call->op.same_as(builtin::decl_buffer()))) { + (call->op.same_as(builtin::alloc_tensor()) || call->op.same_as(builtin::decl_tensor()))) { TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->WithDefRegionKind( kTVMFFIDefRegionKindSimple, [&]() { return this->Visit(op->var); })); - if (call->op.same_as(builtin::decl_buffer())) return Visit(call->args[0]); + if (call->op.same_as(builtin::decl_tensor())) return Visit(call->args[0]); return std::nullopt; } return StmtExprVisitor::Visit_(op); @@ -295,7 +295,7 @@ class PrimFuncSpecializer : public StmtExprMutator { buffer->strides.same_as(strides) && !layout_changed && !storage_scope_changed) { return buffer; } else { - auto n = CopyBufferType(buffer); + auto n = CopyTensorType(buffer); n->elem_offset = std::move(elem_offset); n->shape = std::move(shape); n->strides = std::move(strides); @@ -309,7 +309,7 @@ class PrimFuncSpecializer : public StmtExprMutator { } } - void MutateAllocBuffer(const BufferVar& alloc_buf) { + void MutateAllocTensor(const BufferVar& alloc_buf) { TVM_FFI_ICHECK(defined_buffers_.insert(alloc_buf.get()).second) << "Multiple points of definition found for buffer " << alloc_buf; VarRemapSet(alloc_buf, MutateBuffer(alloc_buf)); @@ -327,9 +327,9 @@ class PrimFuncSpecializer : public StmtExprMutator { << "mutation must occur at the buffer's point of definition " << "(see discussion on https://github.com/apache/tvm/pull/14565 for more details). " << "Please add a definition for this buffer, " - << "either as a BufferType-annotated PrimFunc parameter, " + << "either as a TensorType-annotated PrimFunc parameter, " << "in a block's buffer allocations, " - << "or in a DeclBuffer statement."; + << "or in a DeclTensor statement."; } /*! \brief Definition identities used only to validate declaration order. */ @@ -349,14 +349,14 @@ class PrimFuncSpecializer : public StmtExprMutator { * \param var_map The var mapping to be updated. * \note This function will match target buffer's shape, strides and element_offset * For example, we define a buffer in PrimFunc: - * A: T.Buffer([m, n]) + * A: T.Tensor([m, n]) * - * Then we match it with a buffer B = tirx.decl_buffer((8, 16)) + * Then we match it with a buffer B = tirx.decl_tensor((8, 16)) * * It means we have two var mappings here: m = 8 and n = 16 * * If the buffer signature is not a Var, the mapping will fail. - * e.g. A: T.Buffer([m * 2, n + 1]) + * e.g. A: T.Tensor([m * 2, n + 1]) */ void UpdateSpecializeVarMap(const PrimFunc& func, const Var& param, const BufferVar& specific_buf, VarMap* var_map) { @@ -365,7 +365,7 @@ void UpdateSpecializeVarMap(const PrimFunc& func, const Var& param, const Buffer auto opt_buffer = param.as(); TVM_FFI_CHECK(opt_buffer, ValueError) - << "specialize expects param to have a BufferType annotation"; + << "specialize expects param to have a TensorType annotation"; const BufferVar& buf_to_specialize = opt_buffer.value(); // build var mapping using specific_buf's parameters @@ -440,7 +440,7 @@ void UpdateSpecializeVarMap(const PrimFunc& func, const Var& param, const Expr& << "Specialize expects param to be in PrimFunc's params"; // Specialize a scalar parameter rather than a buffer parameter. TVM_FFI_CHECK(!param.as(), ValueError) - << "Specialize expects param to not have a BufferType annotation"; + << "Specialize expects param to not have a TensorType annotation"; // build var mapping using specific_expr (*var_map)[param] = specific_expr; } diff --git a/src/tirx/ir/stmt.cc b/src/tirx/ir/stmt.cc index 441a9a8e7a0a..d0fb82f3118f 100644 --- a/src/tirx/ir/stmt.cc +++ b/src/tirx/ir/stmt.cc @@ -1022,7 +1022,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { // Evaluate Evaluate::Evaluate(Expr value, Span span) { TVM_FFI_ICHECK(value.defined()); - TVM_FFI_ICHECK(!(value->IsInstance() && value->ty.as())) + TVM_FFI_ICHECK(!(value->IsInstance() && value->ty.as())) << "A buffer variable cannot be used as a scalar Evaluate value; " << "use buffer.data to evaluate its physical pointer"; diff --git a/src/tirx/ir/tir_visitor_with_path.cc b/src/tirx/ir/tir_visitor_with_path.cc index 23929c381cea..a1844b053699 100644 --- a/src/tirx/ir/tir_visitor_with_path.cc +++ b/src/tirx/ir/tir_visitor_with_path.cc @@ -76,14 +76,14 @@ void TIRVisitorWithPath::Visit(const IRModule& mod, AccessPath path) { } void TIRVisitorWithPath::Visit(const PrimFunc& func, AccessPath path) { - // BufferType metadata may introduce symbolic dimensions. Define those + // TensorType metadata may introduce symbolic dimensions. Define those // symbols before entering the buffer parameter itself. std::vector> context; auto ppath = path->Attr("params"); for (size_t i = 0; i < func->params.size(); i++) { const Var& param = func->params[i]; - if (!param->ty.as()) { + if (!param->ty.as()) { context.push_back(WithDef(param, ppath->ArrayItem(i))); } } diff --git a/src/tirx/ir/type.cc b/src/tirx/ir/type.cc index 5b56bcbac66f..d27479c50740 100644 --- a/src/tirx/ir/type.cc +++ b/src/tirx/ir/type.cc @@ -34,7 +34,7 @@ namespace tvm::tirx { -bool BufferTypeNode::IsScalar(bool alloc_or_decl) const { +bool TensorTypeNode::IsScalar(bool alloc_or_decl) const { // TODO(@bohan): logical scope is not considered return shape.size() == 1 && tvm::prim::is_one(shape[0]) && strides.empty() && (!alloc_or_decl || tvm::prim::is_zero(elem_offset)) && data_alignment == 64 && @@ -42,7 +42,7 @@ bool BufferTypeNode::IsScalar(bool alloc_or_decl) const { ffi::StructuralEqual()(layout.value(), TileLayoutNode::DefaultLayout({1})); } -std::optional BufferTypeNode::ConstantAllocationSize() const { +std::optional TensorTypeNode::ConstantAllocationSize() const { int64_t result = 1; for (const PrimExpr& extent : shape) { const auto* size = extent.as(); @@ -71,11 +71,11 @@ TVM_FFI_INLINE ffi::Expected> TensorMapTypeMaybeInpla return ffi::Unchanged(); } -TVM_FFI_INLINE ffi::Expected> BufferTypeVisit( +TVM_FFI_INLINE ffi::Expected> TensorTypeVisit( ffi::StructuralVisitorObj* visitor, ffi::AnyView value) noexcept { // skips: storage_scope, data_alignment, offset_factor - const BufferTypeNode* self = - ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + const TensorTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->dtype)); TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->shape)); // Empty strides denote the common compact layout. Broad callbacks do not see the empty @@ -93,11 +93,11 @@ TVM_FFI_INLINE ffi::Expected> BufferTypeVisit return std::nullopt; } -TVM_FFI_INLINE ffi::Expected> BufferTypeMutate( +TVM_FFI_INLINE ffi::Expected> TensorTypeMutate( ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { // skips: storage_scope, data_alignment, offset_factor - const BufferTypeNode* self = - ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); + const TensorTypeNode* self = + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_dtype, mutator->MutateExpected(self->dtype)); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_shape, @@ -130,7 +130,7 @@ TVM_FFI_INLINE ffi::Expected> BufferTypeMutate( mapped_allocated_addr.UnchangedOrSameAs(self->allocated_addr)) { return ffi::Unchanged(); } - ffi::ObjectPtr copy = ffi::make_object(*self); + ffi::ObjectPtr copy = ffi::make_object(*self); copy->dtype = std::move(mapped_dtype).ValueOrUnchanged(std::move(copy->dtype)); copy->shape = std::move(mapped_shape).ValueOrUnchanged(std::move(copy->shape)); copy->strides = std::move(mapped_strides).ValueOrUnchanged(std::move(copy->strides)); @@ -141,11 +141,11 @@ TVM_FFI_INLINE ffi::Expected> BufferTypeMutate( return ffi::Any(std::move(copy)); } -TVM_FFI_INLINE ffi::Expected> BufferTypeMaybeInplaceMutate( +TVM_FFI_INLINE ffi::Expected> TensorTypeMaybeInplaceMutate( ffi::StructuralMutatorObj* mutator, ffi::AnyView value) noexcept { // skips: storage_scope, data_alignment, offset_factor - BufferTypeNode* self = const_cast( - ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); + TensorTypeNode* self = const_cast( + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(value)); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr, mapped_dtype, mutator->MutateExpected(self->dtype, ffi::InplaceMode::kAllow)); TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr>, mapped_shape, @@ -210,12 +210,12 @@ TVM_FFI_INLINE ffi::Expected> BufferRegionTypeMaybeIn } // namespace -BufferType::BufferType(ffi::String storage_scope, PrimType dtype, ffi::Array shape, +TensorType::TensorType(ffi::String storage_scope, PrimType dtype, ffi::Array shape, ffi::Array strides, PrimExpr elem_offset, int data_alignment, int offset_factor, ffi::Optional layout, ffi::Array allocated_addr, Span span) : Type(ffi::UnsafeInit{}) { - auto n = ffi::make_object(); + auto n = ffi::make_object(); n->dtype = std::move(dtype); n->storage_scope = storage_scope.empty() ? ffi::String("global") : std::move(storage_scope); n->shape = std::move(shape); @@ -235,21 +235,21 @@ BufferType::BufferType(ffi::String storage_scope, PrimType dtype, ffi::Array() + TensorTypeNode::RegisterReflection(); + refl::TypeAttrDef() .attr(refl::type_attr::kStructuralVisit, - ffi::FStructuralVisit::FromNative<&BufferTypeVisit>()) + ffi::FStructuralVisit::FromNative<&TensorTypeVisit>()) .attr(refl::type_attr::kStructuralMutate, - ffi::FStructuralMutate::FromNative<&BufferTypeMutate>()) + ffi::FStructuralMutate::FromNative<&TensorTypeMutate>()) .attr(refl::type_attr::kStructuralMaybeInplaceMutate, - ffi::FStructuralMutate::FromNative<&BufferTypeMaybeInplaceMutate>()); + ffi::FStructuralMutate::FromNative<&TensorTypeMaybeInplaceMutate>()); refl::GlobalDef().def( - "tirx.BufferType", + "tirx.TensorType", [](ffi::String storage_scope, PrimType dtype, ffi::Array shape, ffi::Array strides, PrimExpr elem_offset, int data_alignment, int offset_factor, ffi::Optional layout, ffi::Array allocated_addr, Span span) { - return BufferType(std::move(storage_scope), std::move(dtype), std::move(shape), + return TensorType(std::move(storage_scope), std::move(dtype), std::move(shape), std::move(strides), std::move(elem_offset), data_alignment, offset_factor, std::move(layout), std::move(allocated_addr), std::move(span)); }); diff --git a/src/tirx/op/builtin.cc b/src/tirx/op/builtin.cc index eaa0465f58cf..cd9d4b09315d 100644 --- a/src/tirx/op/builtin.cc +++ b/src/tirx/op/builtin.cc @@ -94,34 +94,34 @@ ffi::Expected InferTypeBuffer(const CallNode* call) noexcept try { tvm::Tuple shape = call->args[shape_index].as_or_throw(); DLDataType dtype = call->args[shape_index + 1].as_or_throw()->value; ffi::String scope = call->args[shape_index + 2].as_or_throw()->value; - auto original = call->ty.as_or_throw(); + auto original = call->ty.as_or_throw(); if (ffi::StructuralEqual()(shape->fields, original->shape) && dtype == original->dtype->dtype && scope == original->storage_scope) { return original; } - auto inferred = ffi::make_object(*original.get()); + auto inferred = ffi::make_object(*original.get()); inferred->shape = shape->fields.Map([](const Expr& extent) { return extent.as_or_throw(); }); inferred->dtype = PrimType(dtype); inferred->storage_scope = scope; - return BufferType(std::move(inferred)); + return TensorType(std::move(inferred)); } catch (const ffi::Error& error) { return ffi::Unexpected(error); } catch (const std::exception& error) { return ffi::Unexpected(ffi::Error("InternalError", error.what(), "")); } -ffi::Expected ValidateDeclBuffer(const CallNode* call) noexcept try { +ffi::Expected ValidateDeclTensor(const CallNode* call) noexcept try { TVM_FFI_CHECK_EQ(call->args.size(), 4U, ValueError); - auto buffer = call->ty.as_or_throw(); + auto buffer = call->ty.as_or_throw(); ffi::String scope = call->args[3].as_or_throw()->value; if (scope == "tmem") { TVM_FFI_CHECK_EQ(buffer->allocated_addr.size(), 1U, ValueError) - << "For `tmem` scope, decl_buffer requires exactly one `allocated_addr` PrimExpr"; + << "For `tmem` scope, decl_tensor requires exactly one `allocated_addr` PrimExpr"; } else if (scope.empty() || scope == "global" || scope == "shared" || scope == "shared.dyn" || scope == "local") { TVM_FFI_CHECK(buffer->allocated_addr.empty(), ValueError) - << "For `" << scope << "` scope, decl_buffer does not accept `allocated_addr`"; + << "For `" << scope << "` scope, decl_tensor does not accept `allocated_addr`"; } return {}; } catch (const ffi::Error& error) { @@ -248,8 +248,8 @@ TVM_DEFINE_CACHED_OP_GETTER(get_active_lane_mask, "tirx.get_active_lane_mask") TVM_DEFINE_CACHED_OP_GETTER(masked_load, "tirx.masked_load") TVM_DEFINE_CACHED_OP_GETTER(masked_store, "tirx.masked_store") TVM_DEFINE_CACHED_OP_GETTER(ignore_loop_partition, "tirx.ignore_loop_partition") -TVM_DEFINE_CACHED_OP_GETTER(alloc_buffer, "tirx.alloc_buffer") -TVM_DEFINE_CACHED_OP_GETTER(decl_buffer, "tirx.decl_buffer") +TVM_DEFINE_CACHED_OP_GETTER(alloc_tensor, "tirx.alloc_tensor") +TVM_DEFINE_CACHED_OP_GETTER(decl_tensor, "tirx.decl_tensor") TVM_DEFINE_CACHED_OP_GETTER(buffer_offset, "tirx.buffer_offset") TVM_DEFINE_CACHED_OP_GETTER(buffer_data, "tirx.buffer_data") TVM_DEFINE_CACHED_OP_GETTER(print_buffer, "tirx.print_buffer") @@ -759,7 +759,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { .set_attr("TScriptDtypePrintLocation", static_cast(ScriptDtypePrintLocation::kNone)); - OpDef("tirx.alloc_buffer") + OpDef("tirx.alloc_tensor") .set_attr("TIRxOpCategory", ffi::String("builtin")) .set_attr("FInferType", FInferType::FromNative<&InferTypeBuffer<0>>()) .add_arg("shape", "The tuple of buffer extents.") @@ -767,11 +767,11 @@ TVM_FFI_STATIC_INIT_BLOCK() { .add_arg("scope", "The storage scope.") .set_attr("TCallEffectKind", static_cast(CallEffectKind::kOpaque)); - OpDef("tirx.decl_buffer") + OpDef("tirx.decl_tensor") .set_attr("TIRxOpCategory", ffi::String("builtin")) .set_attr("FInferType", FInferType::FromNative<&InferTypeBuffer<1>>()) .set_validator(ffi::reflection::NativeFunctionView::FromNative<&ValidateDeclBuffer>()) + const CallNode*)>::FromNative<&ValidateDeclTensor>()) .add_arg("data", "The existing data pointer.") .add_arg("shape", "The tuple of buffer extents.") .add_arg("dtype", "The buffer data type.") diff --git a/src/tirx/script/ir_builder/frame.cc b/src/tirx/script/ir_builder/frame.cc index 8033260df055..be80accbfae5 100644 --- a/src/tirx/script/ir_builder/frame.cc +++ b/src/tirx/script/ir_builder/frame.cc @@ -48,7 +48,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { IfFrameNode::RegisterReflection(); ThenFrameNode::RegisterReflection(); ElseFrameNode::RegisterReflection(); - DeclBufferFrameNode::RegisterReflection(); + DeclTensorFrameNode::RegisterReflection(); } namespace { @@ -276,12 +276,12 @@ void ElseFrameNode::ExitWithScope() { FindIfFrame("T.else_")->else_stmts = stmts; } -void DeclBufferFrameNode::ExitWithScope() { +void DeclTensorFrameNode::ExitWithScope() { TIRFrameNode::ExitWithScope(); if (allocated) { AddToParent(tvm::tirx::SeqStmt::Flatten( tvm::tirx::Bind(buffer, - tvm::Call(buffer.type(), tvm::tirx::builtin::decl_buffer(), + tvm::Call(buffer.type(), tvm::tirx::builtin::decl_tensor(), {data, tvm::Tuple(buffer->shape), tvm::DataTypeImm(buffer->dtype->dtype), tvm::StringImm(buffer.scope())}, @@ -290,11 +290,11 @@ void DeclBufferFrameNode::ExitWithScope() { AsStmt(stmts)), source_span); } else { - // data is undefined in `decl_buffer(...)`, lower to `alloc_buffer(...)`. + // data is undefined in `decl_tensor(...)`, lower to `alloc_tensor(...)`. AddToParent( tvm::tirx::SeqStmt::Flatten( tvm::tirx::Bind(buffer.var(), - Call(buffer.type(), tvm::tirx::builtin::alloc_buffer(), + Call(buffer.type(), tvm::tirx::builtin::alloc_tensor(), {tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())}, DictAttrs(), {}, source_span), diff --git a/src/tirx/script/ir_builder/ir.cc b/src/tirx/script/ir_builder/ir.cc index a1220b76e227..e3e791be2e90 100644 --- a/src/tirx/script/ir_builder/ir.cc +++ b/src/tirx/script/ir_builder/ir.cc @@ -52,7 +52,7 @@ using tvm::tirx::Layout; namespace { -tvm::tirx::BufferType BufferTypeDecl(ffi::Array shape, PrimType dtype, +tvm::tirx::TensorType TensorTypeDecl(ffi::Array shape, PrimType dtype, ffi::Optional data, ffi::Optional> strides, ffi::Optional elem_offset, ffi::String storage_scope, @@ -70,7 +70,7 @@ tvm::tirx::BufferType BufferTypeDecl(ffi::Array shape, PrimType dtype, PrimType shape_dtype = shape.empty() ? PrimType::Int(32) : shape[0].ty(); elem_offset = tvm::PrimVar("elem_offset", shape_dtype); } - return tvm::tirx::BufferType( + return tvm::tirx::TensorType( storage_scope, dtype, shape, strides.value_or(ffi::Array()), elem_offset.value_or(PrimExpr()), align, offset_factor, layout, allocated_addr); } @@ -83,7 +83,7 @@ BufferVar BufferDecl(ffi::Array shape, PrimType dtype, ffi::String buf int offset_factor, ffi::Optional layout, ffi::Array allocated_addr) { return BufferVar(buffer_name, - BufferTypeDecl(shape, dtype, data, strides, elem_offset, storage_scope, align, + TensorTypeDecl(shape, dtype, data, strides, elem_offset, storage_scope, align, offset_factor, layout, allocated_addr)); } @@ -609,7 +609,7 @@ tvm::tirx::Stmt BufferStore(BufferVar buffer, PrimExpr value, ffi::Array shape, PrimType dtype, ffi::String buffer_name, +DeclTensorFrame DeclTensor(ffi::Array shape, PrimType dtype, ffi::String buffer_name, ffi::Optional data, ffi::Optional> strides, ffi::Optional elem_offset, ffi::String storage_scope, int align, int offset_factor, ffi::Optional layout, @@ -619,18 +619,18 @@ DeclBufferFrame DeclBuffer(ffi::Array shape, PrimType dtype, ffi::Stri scope = "global"; } - // Enforce rules for T.decl_buffer based on storage scope + // Enforce rules for T.decl_tensor based on storage scope ffi::Array allocated_addr_arr; if (scope == "tmem") { TVM_FFI_ICHECK(!data.has_value()) - << "ValueError: For `tmem` scope, T.decl_buffer accepts only `allocated_addr`"; + << "ValueError: For `tmem` scope, T.decl_tensor accepts only `allocated_addr`"; TVM_FFI_ICHECK(allocated_addr.has_value()) - << "ValueError: For `tmem` scope, T.decl_buffer requires `allocated_addr` (PrimExpr)"; + << "ValueError: For `tmem` scope, T.decl_tensor requires `allocated_addr` (PrimExpr)"; allocated_addr_arr = ffi::Array({allocated_addr.value()}); } else if (scope == "global" || scope == "shared" || scope == "shared.dyn" || scope == "local") { TVM_FFI_ICHECK(!allocated_addr.has_value()) << "ValueError: For `" << scope - << "` scope, T.decl_buffer does not accept `allocated_addr`"; + << "` scope, T.decl_tensor does not accept `allocated_addr`"; allocated_addr_arr = ffi::Array(); } else { // Other scopes: fall back to provided value if any @@ -641,29 +641,29 @@ DeclBufferFrame DeclBuffer(ffi::Array shape, PrimType dtype, ffi::Stri } } - ffi::ObjectPtr n = ffi::make_object(); + ffi::ObjectPtr n = ffi::make_object(); n->buffer = BufferDecl(shape, dtype, buffer_name, data, strides, elem_offset, storage_scope, align, offset_factor, layout, allocated_addr_arr); if (data.has_value()) { n->data = data.value(); } else if (scope == "tmem") { // Tensor memory is an externally allocated address space. Make that - // address-to-pointer relationship explicit so every DeclBuffer has a + // address-to-pointer relationship explicit so every DeclTensor has a // physical data binding. n->data = Call(n->buffer.DataPointerType(), tvm::tirx::builtin::reinterpret(), {allocated_addr.value()}); } // For tmem, even without `data`, we should not emit an Allocate node. n->allocated = (scope == "tmem") || data.has_value(); - return DeclBufferFrame(n); + return DeclTensorFrame(n); } -BufferVar AllocBuffer(ffi::Array shape, PrimType dtype, ffi::String storage_scope, +BufferVar AllocTensor(ffi::Array shape, PrimType dtype, ffi::String storage_scope, ffi::Optional> annotations) { BufferVar buffer = BufferDecl(shape, dtype, "", std::nullopt, std::nullopt, std::nullopt, storage_scope, 0, 0, std::nullopt, {}); AddToParent(tvm::tirx::Bind( - buffer.var(), Call(buffer.type(), tvm::tirx::builtin::alloc_buffer(), + buffer.var(), Call(buffer.type(), tvm::tirx::builtin::alloc_tensor(), {tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())}, DictAttrs(annotations.value_or(ffi::Map()))))); @@ -716,7 +716,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { ffi::Optional, ffi::Optional>, ffi::Optional, ffi::String, int, int, ffi::Optional, ffi::Array)>(BufferDecl)) - .def("script.ir_builder.tirx.BufferType", BufferTypeDecl) + .def("script.ir_builder.tirx.TensorType", TensorTypeDecl) .def("script.ir_builder.tirx.PrimFunc", PrimFunc) .def("script.ir_builder.tirx.DeclFunction", DeclFunction) .def("script.ir_builder.tirx.Arg", @@ -761,7 +761,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { [](ffi::Optional> extents, ffi::String parent, ffi::String name, ffi::String cur, PrimType dtype) { return ScopeId(extents, parent, name, cur, dtype); }) - .def("script.ir_builder.tirx.AllocBuffer", AllocBuffer) + .def("script.ir_builder.tirx.AllocTensor", AllocTensor) .def("script.ir_builder.tirx.Serial", Serial) .def("script.ir_builder.tirx.Parallel", Parallel) .def("script.ir_builder.tirx.Vectorized", Vectorized) @@ -779,7 +779,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { .def("script.ir_builder.tirx.If", If) .def("script.ir_builder.tirx.Then", Then) .def("script.ir_builder.tirx.Else", Else) - .def("script.ir_builder.tirx.DeclBuffer", DeclBuffer) + .def("script.ir_builder.tirx.DeclTensor", DeclTensor) .def("script.ir_builder.tirx.LaunchThread", [](ffi::Variant thread_tag_or_var, PrimExpr extent) { if (auto var = thread_tag_or_var.as()) { diff --git a/src/tirx/script/printer/buffer.cc b/src/tirx/script/printer/buffer.cc index d5f965c8bc42..09f3f98bc9b8 100644 --- a/src/tirx/script/printer/buffer.cc +++ b/src/tirx/script/printer/buffer.cc @@ -58,9 +58,9 @@ ffi::Optional BufferOperationDocTranslate(DocTranslatorObj* d, ffi::Any << "Buffer operation destination must be a variable"; auto var = ffi::GetRef(static_cast(destination)); if (!ffi::StructuralEqual()(var->ty, call->ty)) return RawCall(d, call); - bool is_alloc = call->op.same_as(tirx::builtin::alloc_buffer()); + bool is_alloc = call->op.same_as(tirx::builtin::alloc_tensor()); size_t shape_index = is_alloc ? 0 : 1; - auto buffer = call->ty.as(); + auto buffer = call->ty.as(); if (!buffer || call->args.size() != shape_index + 3 || !call->ty_args.empty() || call->attrs.defined() != is_alloc || (call->attrs.defined() && !call->attrs.as())) { @@ -106,8 +106,8 @@ ffi::Optional BufferOperationDocTranslate(DocTranslatorObj* d, ffi::Any if (!pointer || pointer->storage_scope != scope->value) return RawCall(d, call); } CallDoc rhs = d->Translate(buffer.value()).value().as_or_throw(); - if (rhs->callee.as_or_throw()->name != "Buffer") return RawCall(d, call); - ffi::String method = is_alloc ? "alloc_buffer" : "decl_buffer"; + if (rhs->callee.as_or_throw()->name != "Tensor") return RawCall(d, call); + ffi::String method = is_alloc ? "alloc_tensor" : "decl_tensor"; if (is_alloc && (scope->value == "local" || scope->value == "shared")) { method = scope->value == "local" ? "alloc_local" : "alloc_shared"; for (size_t i = 0; i < rhs->kwargs_keys.size(); ++i) { @@ -146,16 +146,16 @@ ffi::Optional BufferOperationDocTranslate(DocTranslatorObj* d, ffi::Any } TVM_FFI_STATIC_INIT_BLOCK() { - for (const char* name : {"tirx.alloc_buffer", "tirx.decl_buffer"}) { + for (const char* name : {"tirx.alloc_tensor", "tirx.decl_tensor"}) { OpDef(name).set_attr(kOpCallDocTranslate, FDocTranslate::FromNative<&BufferOperationDocTranslate>()); } } -ffi::Optional BufferTypeDocTranslate(DocTranslatorObj* d, ffi::AnyView input, +ffi::Optional TensorTypeDocTranslate(DocTranslatorObj* d, ffi::AnyView input, const ffi::Object*) { const auto* buffer = - ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(input); + ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck(input); bool default_offset = ffi::StructuralEqual()(buffer->elem_offset, IntImm(PrimType(buffer->DefaultIndexType()), 0)); // The buffer type constructor normalizes these fields and constrains allocated addresses. @@ -164,7 +164,7 @@ ffi::Optional BufferTypeDocTranslate(DocTranslatorObj* d, ffi::AnyView (!buffer->allocated_addr.empty() && (!default_offset || buffer->offset_factor != 1))) { return NamespaceDoc("ir") ->Attr("make_node") - ->Call({LiteralDoc::Str("tirx.BufferType", std::nullopt)}, + ->Call({LiteralDoc::Str("tirx.TensorType", std::nullopt)}, {"dtype", "storage_scope", "shape", "strides", "elem_offset", "data_alignment", "offset_factor", "layout", "allocated_addr"}, {TypeValue(d, buffer->dtype, false), @@ -229,13 +229,13 @@ ffi::Optional BufferTypeDocTranslate(DocTranslatorObj* d, ffi::AnyView keys.push_back("allocated_addr"); values.push_back(addresses.size() == 1 ? addresses[0] : ExprDoc(TupleDoc(addresses))); } - return NamespaceDoc("tirx")->Attr("Buffer")->Call( + return NamespaceDoc("tirx")->Attr("Tensor")->Call( {TupleDoc(shape), LiteralDoc::DataType(buffer->dtype->dtype, std::nullopt)}, keys, values); } TVM_FFI_STATIC_INIT_BLOCK() { - ffi::reflection::TypeAttrDef().attr( - kDocTranslate, FDocTranslate::FromNative<&BufferTypeDocTranslate>()); + ffi::reflection::TypeAttrDef().attr( + kDocTranslate, FDocTranslate::FromNative<&TensorTypeDocTranslate>()); } } // namespace @@ -338,7 +338,7 @@ ffi::Optional BufferLoadDocTranslate(DocTranslatorObj* d, ffi::AnyView } TVM_FFI_STATIC_INIT_BLOCK() { - ffi::reflection::TypeAttrDef().attr( + ffi::reflection::TypeAttrDef().attr( kTensorLoadDocTranslate, FDocTranslate::FromNative<&BufferLoadDocTranslate>()); } diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc index 959fe8d1a21b..9baf881187ea 100644 --- a/src/tirx/script/printer/expr.cc +++ b/src/tirx/script/printer/expr.cc @@ -55,7 +55,7 @@ ffi::Optional VarDocTranslate(DocTranslatorObj* d, ffi::AnyView input, rhs = NamespaceDoc("ir")->Attr("dynamic")->Call( {LiteralDoc::Str(var->name, std::nullopt)}, {"dtype"}, {LiteralDoc::DataType(primitive.value()->dtype, std::nullopt)}); - } else if (var->ty.as()) { + } else if (var->ty.as()) { rhs = NamespaceDoc("tirx")->Attr("Var")->Call( {LiteralDoc::Str(var->name, std::nullopt), d->Translate(var->ty).value()}); } else { diff --git a/src/tirx/script/printer/stmt.cc b/src/tirx/script/printer/stmt.cc index 304b7d42425b..da5579c43782 100644 --- a/src/tirx/script/printer/stmt.cc +++ b/src/tirx/script/printer/stmt.cc @@ -312,7 +312,7 @@ ffi::Optional SeqStmtDocTranslate(DocTranslatorObj* d, ffi::AnyView inp const auto* allocation = alloc ? alloc->value.as() : nullptr; const auto* store = stmt->seq[i + 1].as(); auto docs = d->CurrentScopeDocs(); - if (!allocation || !allocation->op.same_as(tirx::builtin::alloc_buffer()) || !store || + if (!allocation || !allocation->op.same_as(tirx::builtin::alloc_tensor()) || !store || !alloc->var.same_as(store->buffer) || docs.empty()) continue; auto scalar = docs.back().as(); diff --git a/src/tirx/transform/flatten_buffer.cc b/src/tirx/transform/flatten_buffer.cc index 535ced8c7c72..8433be6c35f1 100644 --- a/src/tirx/transform/flatten_buffer.cc +++ b/src/tirx/transform/flatten_buffer.cc @@ -48,7 +48,7 @@ using namespace tvm::prim; * same data origin, dtype, alignment and scope; no layout, no elem_offset. * * The pass walks the AST top-down. At each buffer definition point - * (AllocBuffer/DeclBuffer; PrimFunc params are seeded up front) it derives, + * (AllocTensor/DeclTensor; PrimFunc params are seeded up front) it derives, * exactly once: * - the fold view: the original geometry with its expression fields * (runtime elem_offset, symbolic shapes/strides, layout iters) rewritten @@ -85,7 +85,7 @@ class BufferFlattener : public IRMutatorWithAnalyzer { auto new_buf = pass->Lookup(old_buf.value()).flattened; if (!old_buf.value().same_as(new_buf)) { body = SeqStmt::Flatten( - Bind(new_buf, Call(new_buf.type(), builtin::decl_buffer(), + Bind(new_buf, Call(new_buf.type(), builtin::decl_tensor(), {old_buf.value().data(), tvm::Tuple(new_buf->shape), DataTypeImm(new_buf->dtype->dtype), StringImm(new_buf.scope())}, {})), @@ -122,7 +122,7 @@ class BufferFlattener : public IRMutatorWithAnalyzer { } // Fold view: rewrite the geometry's expression leaves. - auto view_type = CopyBufferType(buf); + auto view_type = CopyTensorType(buf); auto mutate_expr = [this](const PrimExpr& expr) { return Mutate(expr).ValueOrUnchanged(expr); }; view_type->shape = view_type->shape.Map(mutate_expr); view_type->strides = view_type->strides.Map(mutate_expr); @@ -152,7 +152,7 @@ class BufferFlattener : public IRMutatorWithAnalyzer { // buf': the storage husk. The linearized indices carry layout and // elem_offset, so the husk keeps neither. auto flat = fold_view.GetFlattenedBuffer(); - auto type = CopyBufferType(flat); + auto type = CopyTensorType(flat); for (size_t i = 0; i < type->shape.size(); ++i) { type->shape.Set(i, analyzer_->canonical_simplify(type->shape[i])); } @@ -162,10 +162,10 @@ class BufferFlattener : public IRMutatorWithAnalyzer { } // Body-local buffers keep their identity when flattening changes nothing. // PrimFunc-parameter buffers always rebuild: the epilogue aliases the - // rebuilt view onto the argument buffer with an explicit DeclBuffer, and + // rebuilt view onto the argument buffer with an explicit DeclTensor, and // downstream s_tir passes pin that shape. BufferVar flattened = - (!extern_buffers_.count(buf) && ffi::StructuralEqual()(BufferType(type), buf.type())) + (!extern_buffers_.count(buf) && ffi::StructuralEqual()(TensorType(type), buf.type())) ? buf : RebuildBufferVar(buf, std::move(type)); @@ -179,27 +179,27 @@ class BufferFlattener : public IRMutatorWithAnalyzer { auto it = flat_map_.find(buf.var()); TVM_FFI_ICHECK(it != flat_map_.end()) << "Buffer " << buf.name() - << " is used before its definition (AllocBuffer/DeclBuffer/PrimFunc param)"; + << " is used before its definition (AllocTensor/DeclTensor/PrimFunc param)"; return it->second; } UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(builtin::alloc_buffer())) - return MutateAllocBuffer(op, call, inplace_mode); - if (call->op.same_as(builtin::decl_buffer())) return MutateDeclBuffer(op, call, inplace_mode); + if (call->op.same_as(builtin::alloc_tensor())) + return MutateAllocTensor(op, call, inplace_mode); + if (call->op.same_as(builtin::decl_tensor())) return MutateDeclTensor(op, call, inplace_mode); } return IRMutatorWithAnalyzer::Mutate_(op, inplace_mode); } - UnchangedOr MutateAllocBuffer(const BindNode* op, const CallNode* buffer_call, + UnchangedOr MutateAllocTensor(const BindNode* op, const CallNode* buffer_call, InplaceMode inplace_mode) { const FlatInfo& info = Define(op->var.as_or_throw()); if (info.flattened.same_as(op->var.as_or_throw())) { return ffi::Unchanged(); } return Bind(info.flattened.var(), - Call(info.flattened.type(), tirx::builtin::alloc_buffer(), + Call(info.flattened.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(info.flattened->shape, buffer_call->args[0]->span), DataTypeImm(info.flattened->dtype->dtype, buffer_call->args[1]->span), StringImm(info.flattened.scope(), buffer_call->args[2]->span)}, @@ -207,13 +207,13 @@ class BufferFlattener : public IRMutatorWithAnalyzer { op->span); } - UnchangedOr MutateDeclBuffer(const BindNode* op, const CallNode* buffer_call, + UnchangedOr MutateDeclTensor(const BindNode* op, const CallNode* buffer_call, InplaceMode inplace_mode) { Expr data = buffer_call->args[0]; bool is_extern_buffer_source = false; if (const auto* call = buffer_call->args[0].as(); call && call->op.same_as(builtin::buffer_data()) && call->args.size() == 1) { - if (const auto* var = call->args[0].as(); var && var->ty.as()) { + if (const auto* var = call->args[0].as(); var && var->ty.as()) { is_extern_buffer_source = extern_buffers_.count(ffi::GetRef(var).as_or_throw()); } @@ -227,7 +227,7 @@ class BufferFlattener : public IRMutatorWithAnalyzer { return ffi::Unchanged(); } return Bind(info.flattened, - Call(info.flattened.type(), builtin::decl_buffer(), + Call(info.flattened.type(), builtin::decl_tensor(), {std::move(data), tvm::Tuple(info.flattened->shape), DataTypeImm(info.flattened->dtype->dtype), StringImm(info.flattened.scope())}, buffer_call->attrs, buffer_call->ty_args, buffer_call->span), @@ -280,7 +280,7 @@ class BufferFlattener : public IRMutatorWithAnalyzer { } if (op->op.same_as(builtin::buffer_data()) && op->args.size() == 1) { if (auto var = op->args[0].as()) { - if (var.value()->ty.as()) { + if (var.value()->ty.as()) { BufferVar original = var.value().as_or_throw(); buffers_used_.insert(original); return Lookup(original).flattened.data(); @@ -314,7 +314,7 @@ class BufferFlattener : public IRMutatorWithAnalyzer { return BufferLoad(info.flattened, FoldIndices(info, node->indices), node->span); } - /*! \brief Set of buffers accessed during visitation (used to emit DeclBuffer for param buffers). + /*! \brief Set of buffers accessed during visitation (used to emit DeclTensor for param buffers). */ std::unordered_set buffers_used_; diff --git a/src/tirx/transform/force_narrow_index_to_i32.h b/src/tirx/transform/force_narrow_index_to_i32.h index c0a5dbe2999e..aaabfbfe0da0 100644 --- a/src/tirx/transform/force_narrow_index_to_i32.h +++ b/src/tirx/transform/force_narrow_index_to_i32.h @@ -90,12 +90,12 @@ class Int32DTypeNarrowerBase : public Normalizer { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); - call && call->op.same_as(builtin::alloc_buffer())) - return MutateAllocBuffer(op, inplace_mode); + call && call->op.same_as(builtin::alloc_tensor())) + return MutateAllocTensor(op, inplace_mode); return Normalizer::Mutate_(op, inplace_mode); } - UnchangedOr MutateAllocBuffer(const BindNode* op, InplaceMode inplace_mode) { + UnchangedOr MutateAllocTensor(const BindNode* op, InplaceMode inplace_mode) { auto result = Normalizer::Mutate_(op, inplace_mode); auto alloc = std::move(result).ValueOrUnchanged(ffi::GetRef(op)).template as_or_throw(); diff --git a/src/tirx/transform/inline_private_functions.cc b/src/tirx/transform/inline_private_functions.cc index 9644c30495d2..4e62432d0e1e 100644 --- a/src/tirx/transform/inline_private_functions.cc +++ b/src/tirx/transform/inline_private_functions.cc @@ -118,7 +118,7 @@ bool IsInlinablePrimFunc(const GlobalVar& gvar, const PrimFunc& prim_func, // We do not currently support inlining of functions that accept // buffer arguments. for (const Var& param : prim_func->params) { - if (param->ty.as()) return false; + if (param->ty.as()) return false; } // Generalize the old SBlockRealize exclusion to all non-native statement roots: @@ -243,9 +243,9 @@ class PrimFuncInliner : public StmtExprMutator { << ")"; for (const Var& param : callee->params) { - TVM_FFI_ICHECK(!param->ty.as()) + TVM_FFI_ICHECK(!param->ty.as()) << "Inlining of PrimFuncs with buffer arguments is not yet supported, " - << "but callee " << gvar << " has BufferType-annotated parameter " << param; + << "but callee " << gvar << " has TensorType-annotated parameter " << param; } ffi::Map> param_map; diff --git a/src/tirx/transform/ir_utils.cc b/src/tirx/transform/ir_utils.cc index 29e621fbbb19..3565298ab259 100644 --- a/src/tirx/transform/ir_utils.cc +++ b/src/tirx/transform/ir_utils.cc @@ -376,7 +376,7 @@ Stmt ConvertSSA(Stmt stmt) { } ffi::String GetPtrStorageScope(Var buffer_var) { - if (const auto* buffer_type = buffer_var->ty.as()) { + if (const auto* buffer_type = buffer_var->ty.as()) { return buffer_type->storage_scope; } const auto* ptr_type = buffer_var->ty.as(); diff --git a/src/tirx/transform/ir_utils.h b/src/tirx/transform/ir_utils.h index f1cf677eed98..e00055d0cc4f 100644 --- a/src/tirx/transform/ir_utils.h +++ b/src/tirx/transform/ir_utils.h @@ -119,7 +119,7 @@ inline Call AddressOffset(Var handle, PrimType dtype, int offset) { ffi::Array shape = {offset_expr + 1}; auto pointer_type = handle->ty.as_or_throw(); BufferVar dummy_buf(handle->name, - BufferType(pointer_type->storage_scope, dtype, shape, {}, 0, 0, 0)); + TensorType(pointer_type->storage_scope, dtype, shape, {}, 0, 0, 0)); TensorLoad buf_load = BufferLoad(dummy_buf, {offset_expr}); return Call(handle->ty, builtin::address_of(), {buf_load}); @@ -140,7 +140,7 @@ inline Call AddressOffset(Var handle, PrimType dtype, PrimExpr offset) { ffi::Array shape = {offset + 1}; auto pointer_type = handle->ty.as_or_throw(); - BufferVar dummy_buf(handle->name, BufferType(pointer_type->storage_scope, dtype.WithLanes(1), + BufferVar dummy_buf(handle->name, TensorType(pointer_type->storage_scope, dtype.WithLanes(1), shape, {}, 0, 0, 0)); TensorLoad buf_load = BufferLoad(dummy_buf, {offset}); diff --git a/src/tirx/transform/lower_intrin.cc b/src/tirx/transform/lower_intrin.cc index 4625c84128bc..d7f0b6d2577b 100644 --- a/src/tirx/transform/lower_intrin.cc +++ b/src/tirx/transform/lower_intrin.cc @@ -97,7 +97,7 @@ static Expr LowerAccessPtr(const CallNode* call, BufferVar access_buffer{nullptr}; ffi::String storage_scope; Expr access_data; - if (buffer_var->ty.as()) { + if (buffer_var->ty.as()) { BufferVar source_buffer = buffer_var.as_or_throw(); if (source_buffer->dtype == scalar_dtype && source_buffer->shape.size() == 1) { access_buffer = source_buffer; @@ -114,11 +114,11 @@ static Expr LowerAccessPtr(const CallNode* call, } if (!access_buffer.defined()) { - // BufferVar identity includes its immutable BufferType. Bind an explicit + // BufferVar identity includes its immutable TensorType. Bind an explicit // scalar physical view instead of retyping a vector, padded, or packed source. access_buffer = BufferVar(buffer_var->name + "_access", - BufferType(storage_scope, scalar_dtype, {scalar_extent}, {}, 0, 0, 0)); + TensorType(storage_scope, scalar_dtype, {scalar_extent}, {}, 0, 0, 0)); buffer_aliases->push_back({access_buffer, access_data}); } TensorLoad buf_load = BufferLoad(access_buffer, {offset}); @@ -175,7 +175,7 @@ class IntrinInjecter : public IRMutatorWithAnalyzer { const auto& alias = access_ptr_buffer_aliases_[i - 1]; result = SeqStmt::Flatten( Bind(alias.buffer, - Call(alias.buffer.type(), builtin::decl_buffer(), + Call(alias.buffer.type(), builtin::decl_tensor(), {alias.data, tvm::Tuple(alias.buffer->shape), DataTypeImm(alias.buffer->dtype->dtype), StringImm(alias.buffer.scope())}, {})), diff --git a/src/tirx/transform/lower_thread_allreduce.h b/src/tirx/transform/lower_thread_allreduce.h index b866df5aa2c5..dc5f5162c1f4 100644 --- a/src/tirx/transform/lower_thread_allreduce.h +++ b/src/tirx/transform/lower_thread_allreduce.h @@ -65,7 +65,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { warp_size_(target->GetAttr("thread_warp_size", 1).value()), max_num_threads_(target->GetAttr("max_num_threads", -1).value()) { for (const Var& param : params) { - if (param->ty.as()) { + if (param->ty.as()) { buffer_aliases_.Set(param, param); } } @@ -100,15 +100,15 @@ class ThreadAllreduceBuilder final : public DialectMutator { } UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(builtin::alloc_buffer())) return MutateAllocBuffer(op, inplace_mode); - if (call->op.same_as(builtin::decl_buffer())) return MutateDeclBuffer(op, call, inplace_mode); + if (call->op.same_as(builtin::alloc_tensor())) return MutateAllocTensor(op, inplace_mode); + if (call->op.same_as(builtin::decl_tensor())) return MutateDeclTensor(op, call, inplace_mode); } return DialectMutator::Mutate_(op, inplace_mode); } - UnchangedOr MutateAllocBuffer(const BindNode* op, InplaceMode inplace_mode) { + UnchangedOr MutateAllocTensor(const BindNode* op, InplaceMode inplace_mode) { buffer_aliases_.Set(op->var, op->var); - // In flat IR, alloc_remap_ may not yet be populated when this AllocBuffer is visited + // In flat IR, alloc_remap_ may not yet be populated when this AllocTensor is visited // (the remap is set up by MakeAllreduce which runs during AttrStmt/Evaluate visit // that appears later in the sequence). We record the original data pointer and // attempt the remap; if it's not ready, the post-processing pass will handle it. @@ -118,7 +118,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { .template as_or_throw(); if (auto it = alloc_remap_.find(orig_data_ptr); it != alloc_remap_.end()) { - return RemapAllocBuffer(node, it->second); + return RemapAllocTensor(node, it->second); } // Record for deferred remapping (flat IR case) pending_alloc_buffers_.emplace_back(orig_data_ptr); @@ -126,19 +126,19 @@ class ThreadAllreduceBuilder final : public DialectMutator { } /*! - * \brief Remap an AllocBuffer node to use the replacement buffer. - * \param node The original AllocBuffer node. + * \brief Remap an AllocTensor node to use the replacement buffer. + * \param node The original AllocTensor node. * \param replacement The replacement buffer. * \return The remapped statement(s). */ - Stmt RemapAllocBuffer(Bind node, const BufferVar& replacement) { + Stmt RemapAllocTensor(Bind node, const BufferVar& replacement) { const CallNode* call = node->value.template as(); DictAttrs annotations = call->attrs.as_or_throw(); if (replacement.scope() == "shared") { annotations.CopyOnWrite()->dict.Set(tirx::attr::kVolatile, true); } return Bind(replacement.var(), - Call(replacement.type(), tirx::builtin::alloc_buffer(), + Call(replacement.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(replacement->shape, call->args[0]->span), DataTypeImm(replacement->dtype->dtype, call->args[1]->span), StringImm(replacement.scope(), call->args[2]->span)}, @@ -155,7 +155,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { return std::nullopt; } - UnchangedOr MutateDeclBuffer(const BindNode* op, const CallNode* buffer_call, + UnchangedOr MutateDeclTensor(const BindNode* op, const CallNode* buffer_call, InplaceMode inplace_mode) { RegisterBufferAlias(op->var.as_or_throw(), buffer_call->args[0]); // Remap declarations only after the complete traversal has populated the @@ -388,7 +388,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { std::vector staging_shared_bufs; staging_shared_bufs.reserve(size); for (size_t i = 0; i < size; ++i) { - BufferVar staging_shared_buf = decl_buffer( + BufferVar staging_shared_buf = decl_tensor( /*shape=*/{IntImm(reduce_index.ty(), n_warps * group_extent)}, /*dtype=*/buffers[i]->dtype, /*name=*/"red_buf_staging", /*storage_scope=*/"shared"); staging_shared_bufs.push_back(staging_shared_buf); @@ -436,7 +436,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { new_alloc_bufs.push_back(reduce_results[i] .as_or_throw() ->source.as_or_throw()); - BufferVar broadcast_shared_buf = decl_buffer( + BufferVar broadcast_shared_buf = decl_tensor( /*shape=*/{IntImm(reduce_index.ty(), group_extent)}, /*dtype=*/buffers[i]->dtype, /*name=*/"red_result", /*storage_scope=*/"shared"); write_result.push_back( @@ -457,8 +457,8 @@ class ThreadAllreduceBuilder final : public DialectMutator { TVM_FFI_ICHECK_EQ(reduce_results[i].ty(), dtypes[i]); load_remap_[alloc_key] = reduce_results[i]; - // The AllocBuffer doesn't need to be emitted here since alloc_remap_ - // will cause the existing allocation to be rewritten in MutateAllocBuffer. + // The AllocTensor doesn't need to be emitted here since alloc_remap_ + // will cause the existing allocation to be rewritten in MutateAllocTensor. alloc_remap_[alloc_key] = buf; allreduce_var_remap_[alloc_key] = buf.var(); allreduce_var_remap_[buffers[i].get()] = buf.var(); @@ -477,7 +477,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { // previous iteration on the same buffer. seq.emplace_back(SyncThread("shared")); for (size_t idx = 0; idx < size; ++idx) { - shared_bufs[idx] = decl_buffer({IntImm(group_index.ty(), group_extent * reduce_extent)}, + shared_bufs[idx] = decl_tensor({IntImm(group_index.ty(), group_extent * reduce_extent)}, dtypes[idx], "red_buf" + std::to_string(idx), "shared"); seq.emplace_back(BufferStore(shared_bufs[idx], values[idx], {BufIndex(reduce_index, group_index, reduce_extent)})); @@ -505,7 +505,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { for (BufferVar buf : new_alloc_bufs) { alloc_stmts.push_back(Bind( buf.var(), - Call(buf.type(), tirx::builtin::alloc_buffer(), + Call(buf.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(buf->shape), DataTypeImm(buf->dtype->dtype), StringImm(buf.scope())}, DictAttrs()))); } @@ -542,12 +542,12 @@ class ThreadAllreduceBuilder final : public DialectMutator { load_values.reserve(n_buffers); for (int idx = 0; idx < n_buffers; ++idx) { shared_bufs.push_back( - decl_buffer(shape, dtypes[idx], "red_buf" + std::to_string(idx), "local")); + decl_tensor(shape, dtypes[idx], "red_buf" + std::to_string(idx), "local")); load_values.push_back(BufferStore(shared_bufs[idx], src_values[idx], zero_indices)); // Uses a local variable to store the shuffled data. Later // on, an allocation will be built for this local variable. - local_bufs.push_back(decl_buffer(shape, dtypes[idx], "t" + std::to_string(idx), "local")); + local_bufs.push_back(decl_tensor(shape, dtypes[idx], "t" + std::to_string(idx), "local")); } if (predicate.has_value()) { @@ -561,7 +561,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { // active channels. ffi::Optional mask_buffer; if (need_warp_shuffle_mask_) { - mask_buffer = decl_buffer(shape, mask.ty(), "mask", "local"); + mask_buffer = decl_tensor(shape, mask.ty(), "mask", "local"); seq->emplace_back(BufferStore(mask_buffer.value(), mask, zero_indices)); // Push the buffer description. Later this will have an // allocation built for it. @@ -869,10 +869,10 @@ class ThreadAllreduceBuilder final : public DialectMutator { void RegisterBufferAlias(BufferVar buffer, const Expr& data) { Var root = buffer.var(); if (auto source = GetBufferDataVar(data); - source.has_value() && source.value()->ty.as()) { + source.has_value() && source.value()->ty.as()) { auto source_root = buffer_aliases_.Get(source.value()); TVM_FFI_ICHECK(source_root.has_value()) << "Buffer alias source " << source.value()->name - << " must be registered before its DeclBuffer alias"; + << " must be registered before its DeclTensor alias"; root = source_root.value(); } buffer_aliases_.Set(buffer.var(), root); @@ -898,7 +898,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { public: const VarNode* GetAllocationKey(const VarNode* buffer) const { - if (buffer->ty.as()) { + if (buffer->ty.as()) { Var var = ffi::GetRef(buffer); return buffer_aliases_.Get(var).value_or(var).get(); } @@ -910,7 +910,7 @@ class ThreadAllreduceBuilder final : public DialectMutator { std::unordered_map alloc_remap_; // BufferVar remap std::unordered_map allreduce_var_remap_; - // Pending AllocBuffer original data pointers (for flat IR deferred remapping) + // Pending AllocTensor original data pointers (for flat IR deferred remapping) std::vector pending_alloc_buffers_; // Physical roots of buffer aliases, flattened at each declaration. ffi::Map buffer_aliases_; @@ -919,9 +919,9 @@ class ThreadAllreduceBuilder final : public DialectMutator { /*! * \brief Post-processing pass to apply deferred remappings for flat IR. * - * In flat IR, AllocBuffer nodes may be visited before the alloc_remap_ is populated + * In flat IR, AllocTensor nodes may be visited before the alloc_remap_ is populated * (since MakeAllreduce runs when Evaluate is visited, which is later in the flat sequence). - * Handles AllocBuffer, DeclBuffer, and TensorLoad nodes whose remappings + * Handles AllocTensor, DeclTensor, and TensorLoad nodes whose remappings * were not available during the main traversal. */ template @@ -951,13 +951,13 @@ class DeferredRemapper : public DialectMutator { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(builtin::alloc_buffer())) return MutateAllocBuffer(op, inplace_mode); - if (call->op.same_as(builtin::decl_buffer())) return MutateDeclBuffer(op, inplace_mode); + if (call->op.same_as(builtin::alloc_tensor())) return MutateAllocTensor(op, inplace_mode); + if (call->op.same_as(builtin::decl_tensor())) return MutateDeclTensor(op, inplace_mode); } return DialectMutator::Mutate_(op, inplace_mode); } - UnchangedOr MutateAllocBuffer(const BindNode* op, InplaceMode inplace_mode) { + UnchangedOr MutateAllocTensor(const BindNode* op, InplaceMode inplace_mode) { auto node = DialectMutator::Mutate_(op, inplace_mode) .ValueOrUnchanged(ffi::GetRef(op)) .template as_or_throw(); @@ -971,7 +971,7 @@ class DeferredRemapper : public DialectMutator { annotations.CopyOnWrite()->dict.Set(tirx::attr::kVolatile, true); } return Bind(replacement.var(), - Call(replacement.type(), tirx::builtin::alloc_buffer(), + Call(replacement.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(replacement->shape, call->args[0]->span), DataTypeImm(replacement->dtype->dtype, call->args[1]->span), StringImm(replacement.scope(), call->args[2]->span)}, @@ -982,7 +982,7 @@ class DeferredRemapper : public DialectMutator { return node; } - UnchangedOr MutateDeclBuffer(const BindNode* op, InplaceMode inplace_mode) { + UnchangedOr MutateDeclTensor(const BindNode* op, InplaceMode inplace_mode) { const VarNode* root = buffer_aliases_.Get(op->var).value_or(op->var).get(); if (pending_set_.count(root) && alloc_remap_.count(root)) { return Evaluate(0); @@ -994,7 +994,7 @@ class DeferredRemapper : public DialectMutator { const CallNode* call = node->value.template as(); return Bind( new_buf.value(), - Call(new_buf.value().type(), builtin::decl_buffer(), + Call(new_buf.value().type(), builtin::decl_tensor(), {call->args[0], tvm::Tuple(new_buf.value()->shape), DataTypeImm(new_buf.value()->dtype->dtype), StringImm(new_buf.value().scope())}, call->attrs, call->ty_args, call->span), diff --git a/src/tirx/transform/lower_tirx_cleanup.cc b/src/tirx/transform/lower_tirx_cleanup.cc index b880278939a4..09954531e962 100644 --- a/src/tirx/transform/lower_tirx_cleanup.cc +++ b/src/tirx/transform/lower_tirx_cleanup.cc @@ -64,7 +64,7 @@ class LayoutApplier : public IRMutatorWithAnalyzer { storage_lower->buffer_aliases_.Set(buffer.value().var(), buffer.value().var()); if (buffer.value()->layout.has_value()) { BufferVar flattened = storage_lower->GetFlattenedBuffer(buffer.value()); - auto type = CopyBufferType(buffer.value()); + auto type = CopyTensorType(buffer.value()); type->layout = std::nullopt; BufferVar source = RebuildBufferVar(buffer.value(), std::move(type)); param_flattened_buffers.emplace_back(flattened, source); @@ -76,7 +76,7 @@ class LayoutApplier : public IRMutatorWithAnalyzer { auto new_stmt = storage_lower->Mutate(stmt, InplaceMode::kAllow).ValueOrUnchanged(stmt); for (const auto& [buf, source] : param_flattened_buffers) { new_stmt = - SeqStmt::Flatten(Bind(buf, Call(buf.type(), builtin::decl_buffer(), + SeqStmt::Flatten(Bind(buf, Call(buf.type(), builtin::decl_tensor(), {source.data(), tvm::Tuple(buf->shape), DataTypeImm(buf->dtype->dtype), StringImm(buf.scope())}, {})), @@ -107,11 +107,11 @@ class LayoutApplier : public IRMutatorWithAnalyzer { UnchangedOr Mutate_(const CallNode* op, InplaceMode inplace_mode) final { if (op->op.same_as(builtin::buffer_data()) && op->args.size() == 1) { if (auto var = op->args[0].as(); - var.has_value() && var.value()->ty.as()) { + var.has_value() && var.value()->ty.as()) { auto root_opt = buffer_aliases_.Get(var.value()); TVM_FFI_ICHECK(root_opt.has_value()) << "buffer_data projects " << var.value()->name << ", which has no visible definition " - << "(AllocBuffer/DeclBuffer/PrimFunc parameter) at this point"; + << "(AllocTensor/DeclTensor/PrimFunc parameter) at this point"; Var root = root_opt.value(); if (auto mapped = VarRemapGet(root); mapped != nullptr) { root = mapped.as_or_throw(); @@ -124,14 +124,14 @@ class LayoutApplier : public IRMutatorWithAnalyzer { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(builtin::alloc_buffer())) - return MutateAllocBuffer(op, call, inplace_mode); - if (call->op.same_as(builtin::decl_buffer())) return MutateDeclBuffer(op, call, inplace_mode); + if (call->op.same_as(builtin::alloc_tensor())) + return MutateAllocTensor(op, call, inplace_mode); + if (call->op.same_as(builtin::decl_tensor())) return MutateDeclTensor(op, call, inplace_mode); } return IRMutatorWithAnalyzer::Mutate_(op, inplace_mode); } - UnchangedOr MutateAllocBuffer(const BindNode* op, const CallNode* buffer_call, + UnchangedOr MutateAllocTensor(const BindNode* op, const CallNode* buffer_call, InplaceMode inplace_mode) { buffer_aliases_.Set(op->var, op->var); auto mutate = [this](BufferVar buf) { @@ -145,7 +145,7 @@ class LayoutApplier : public IRMutatorWithAnalyzer { return ffi::Unchanged(); } return Bind(buffer.var(), - Call(buffer.type(), tirx::builtin::alloc_buffer(), + Call(buffer.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(buffer->shape, buffer_call->args[0]->span), DataTypeImm(buffer->dtype->dtype, buffer_call->args[1]->span), StringImm(buffer.scope(), buffer_call->args[2]->span)}, @@ -153,7 +153,7 @@ class LayoutApplier : public IRMutatorWithAnalyzer { op->span); } - UnchangedOr MutateDeclBuffer(const BindNode* op, const CallNode* buffer_call, + UnchangedOr MutateDeclTensor(const BindNode* op, const CallNode* buffer_call, InplaceMode inplace_mode) { RegisterBufferAlias(op->var.as_or_throw(), buffer_call->args[0]); auto data_result = Mutate(buffer_call->args[0], inplace_mode); @@ -164,7 +164,7 @@ class LayoutApplier : public IRMutatorWithAnalyzer { return ffi::Unchanged(); } return Bind(buffer, - Call(buffer.type(), builtin::decl_buffer(), + Call(buffer.type(), builtin::decl_tensor(), {std::move(data), tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())}, buffer_call->attrs, buffer_call->ty_args, buffer_call->span), @@ -177,7 +177,7 @@ class LayoutApplier : public IRMutatorWithAnalyzer { } auto trn_layout = buf->layout.as(); BufferVar flattened; - ffi::ObjectPtr type; + ffi::ObjectPtr type; if (trn_layout && trn_layout->IsTrainium()) { ffi::Array new_shape = buf.scope() == "trn.psum" ? ffi::Array{trn_layout->GetSpan(ffi::String("Bank")), @@ -186,13 +186,13 @@ class LayoutApplier : public IRMutatorWithAnalyzer { : ffi::Array{trn_layout->GetSize(ffi::String("P")), trn_layout->GetSpan(ffi::String("F"))}; flattened = buf; - type = CopyBufferType(flattened); + type = CopyTensorType(flattened); type->shape = new_shape; type->strides = {}; } else if (is_alloc) { if (auto tile_layout = buf->layout.as(); tile_layout && tile_layout->HasThreadAxis()) { - // Logical alloc_buffer with thread axes: physical shape = memory-axis span + // Logical alloc_tensor with thread axes: physical shape = memory-axis span sym::Analyzer ana; PrimExpr mem_span = IntImm::Int32(1); for (const auto& iter : tile_layout->shard) { @@ -211,16 +211,16 @@ class LayoutApplier : public IRMutatorWithAnalyzer { } } flattened = buf; - type = CopyBufferType(flattened); + type = CopyTensorType(flattened); type->shape = {ana->Simplify(mem_span)}; type->strides = {}; } else { flattened = buf.GetFlattenedBuffer(); - type = CopyBufferType(flattened); + type = CopyTensorType(flattened); } } else { flattened = buf.GetFlattenedBuffer(); - type = CopyBufferType(flattened); + type = CopyTensorType(flattened); } // Remap variables the pass has already rebuilt (a shape may load from // another local buffer), then canonicalize. @@ -242,7 +242,7 @@ class LayoutApplier : public IRMutatorWithAnalyzer { StmtExprMutator::Mutate(ffi::AnyView(buf->elem_offset), InplaceMode::kDisallow) .ValueOrUnchanged(buf->elem_offset) .as_or_throw(); - if (ffi::StructuralEqual()(buf.type(), BufferType(type))) { + if (ffi::StructuralEqual()(buf.type(), TensorType(type))) { return buf; } flattened = RebuildBufferVar(flattened, std::move(type)); @@ -355,11 +355,11 @@ class LayoutApplier : public IRMutatorWithAnalyzer { if (const auto* call = data.as(); call && call->op.same_as(builtin::buffer_data()) && call->args.size() == 1) { if (auto source = call->args[0].as(); - source.has_value() && source.value()->ty.as()) { + source.has_value() && source.value()->ty.as()) { auto source_root = buffer_aliases_.Get(source.value()); TVM_FFI_ICHECK(source_root.has_value()) << "Buffer alias source " << source.value()->name - << " must be registered before its DeclBuffer alias"; + << " must be registered before its DeclTensor alias"; root = source_root.value(); } } diff --git a/src/tirx/transform/lower_tirx_opaque.cc b/src/tirx/transform/lower_tirx_opaque.cc index b7df131a0c77..1546d6224520 100644 --- a/src/tirx/transform/lower_tirx_opaque.cc +++ b/src/tirx/transform/lower_tirx_opaque.cc @@ -21,7 +21,7 @@ * \file lower_tirx_opaque.cc * \brief Lower opaque constructs in TIRX programs. This is the tirx-specific * counterpart of s_tirx::LowerOpaqueBlock, handling only the non-SBlock - * parts: AllocBuffer lowering, For(thread_binding) → AttrStmt(thread_extent), + * parts: AllocTensor lowering, For(thread_binding) → AttrStmt(thread_extent), * unit loop elimination, and pragma annotation handling. */ @@ -37,7 +37,7 @@ namespace tirx { using namespace tvm::prim; /*! - * \brief Lower opaque constructs for TIRX: AllocBuffer, thread bindings, unit loops. + * \brief Lower opaque constructs for TIRX: AllocTensor, thread bindings, unit loops. * * Unlike s_tirx::LowerOpaqueBlock, this pass does NOT handle SBlock/SBlockRealize, * since TIRX programs do not contain SBlock nodes. diff --git a/src/tirx/transform/lower_tvm_builtin.cc b/src/tirx/transform/lower_tvm_builtin.cc index 2cb38708c02f..cc375fc57e67 100644 --- a/src/tirx/transform/lower_tvm_builtin.cc +++ b/src/tirx/transform/lower_tvm_builtin.cc @@ -147,7 +147,7 @@ class BuiltinLower : public StmtExprMutator { { // NOTE: this scope reference is invalid after any mutation is applied to alloca_scope_. auto& scope = precheck->alloca_scope_.back(); - scope.stack_shape = decl_buffer({IntImm::Int64(0)}, PrimType::Int(64), "stack_shape"); + scope.stack_shape = decl_tensor({IntImm::Int64(0)}, PrimType::Int(64), "stack_shape"); } precheck->Mutate(stmt, InplaceMode::kDisallow).ValueOrUnchanged(stmt); @@ -188,10 +188,10 @@ class BuiltinLower : public StmtExprMutator { } if (scope.max_sizes.shape_stack != -1) { - scope.stack_shape = decl_buffer({IntImm::Int64(scope.max_sizes.shape_stack)}, + scope.stack_shape = decl_tensor({IntImm::Int64(scope.max_sizes.shape_stack)}, PrimType::Int(64), "stack_shape"); stmt = SeqStmt::Flatten( - Bind(scope.stack_shape, Call(scope.stack_shape.type(), builtin::decl_buffer(), + Bind(scope.stack_shape, Call(scope.stack_shape.type(), builtin::decl_tensor(), {StackAlloca(scope.stack_shape.DataPointerType(), "shape", scope.max_sizes.shape_stack), tvm::Tuple(scope.stack_shape->shape), @@ -254,8 +254,8 @@ class BuiltinLower : public StmtExprMutator { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); - call && call->op.same_as(builtin::alloc_buffer())) - return MutateAllocBuffer(op, inplace_mode); + call && call->op.same_as(builtin::alloc_tensor())) + return MutateAllocTensor(op, inplace_mode); if (const CallNode* call = op->value.as()) { if (call->op.same_as(builtin::nd_mem_alloc_with_scope())) { return MakeNdMemAllocWithScope(op, call); @@ -264,9 +264,9 @@ class BuiltinLower : public StmtExprMutator { return StmtExprMutator::Mutate_(op, inplace_mode); } - UnchangedOr MutateAllocBuffer(const BindNode* op, InplaceMode inplace_mode) { - // Lower AllocBuffer to device allocate when needed. - // AllocBuffer is flat (no body). Visit buffer fields via base class. + UnchangedOr MutateAllocTensor(const BindNode* op, InplaceMode inplace_mode) { + // Lower AllocTensor to device allocate when needed. + // AllocTensor is flat (no body). Visit buffer fields via base class. Stmt stmt = StmtExprMutator::Mutate_(op, inplace_mode).ValueOrUnchanged(ffi::GetRef(op)); op = stmt.as(); const auto* buffer_call = op->value.as(); @@ -287,7 +287,7 @@ class BuiltinLower : public StmtExprMutator { if (const auto* dev_type = device_type_.as(); dev_type && dev_type->value == kDLCPU) { if (scope == "global") { - auto constant_size = op->var->ty.as_or_throw()->ConstantAllocationSize(); + auto constant_size = op->var->ty.as_or_throw()->ConstantAllocationSize(); if (constant_size.has_value() && constant_size.value() > 0 && static_cast(constant_size.value()) * nbytes < runtime::kMaxStackAlloca) { return stmt; @@ -322,7 +322,7 @@ class BuiltinLower : public StmtExprMutator { Stmt alloc_bind = Bind(op->var.as_or_throw(), - Call(op->var.as_or_throw().type(), builtin::decl_buffer(), + Call(op->var.as_or_throw().type(), builtin::decl_tensor(), {Call(op->var.as_or_throw().DataPointerType(), alloc_workspace_op, {prim::cast(PrimType::Int(32), device_type_.value()), prim::cast(PrimType::Int(32), device_id_.value()), total_bytes, @@ -725,7 +725,7 @@ class BuiltinLower : public StmtExprMutator { * free_nd stmt is pushed to the current scope's pending_frees. Body-carrying * stmts (For, IfThenElse, AttrStmt) create new scopes via * WithNewScope. On scope exit, pending_frees are appended after the body. - * AllocBuffer (flat, no body) pushes its free to the enclosing scope. + * AllocTensor (flat, no body) pushes its free to the enclosing scope. */ struct ScopeLevel { std::vector pending_frees; diff --git a/src/tirx/transform/lower_warp_memory.cc b/src/tirx/transform/lower_warp_memory.cc index a7749cb64c57..5fe301c0ff6c 100644 --- a/src/tirx/transform/lower_warp_memory.cc +++ b/src/tirx/transform/lower_warp_memory.cc @@ -268,7 +268,7 @@ class WarpAccessRewriter : public StmtExprMutator { using StmtExprMutator::Mutate_; explicit WarpAccessRewriter(int warp_size, sym::AnalyzerObj* analyzer) : warp_size_(warp_size), analyzer_(analyzer) {} - // Rewrite the AllocBuffer statement which transforms + // Rewrite the AllocTensor statement which transforms // warp memory to local memory. // \param op The allocation binding for warp memory. // \param buffer_call The matched allocation Call. @@ -299,7 +299,7 @@ class WarpAccessRewriter : public StmtExprMutator { warp_group_ = (alloc_size + (factor - 1)) / factor; alloc_size = warp_group_ * factor; - auto type = CopyBufferType(op->var.as_or_throw()); + auto type = CopyTensorType(op->var.as_or_throw()); type->storage_scope = "local"; type->shape = {IntImm::Int32(alloc_size / width_)}; type->strides = {}; @@ -309,7 +309,7 @@ class WarpAccessRewriter : public StmtExprMutator { Stmt rewritten_body = this->Mutate(body, InplaceMode::kDisallow).ValueOrUnchanged(body); return SeqStmt::Flatten( Bind(new_buf.var(), - Call(new_buf.type(), tirx::builtin::alloc_buffer(), + Call(new_buf.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(new_buf->shape, buffer_call->args[0]->span), DataTypeImm(new_buf->dtype->dtype, buffer_call->args[1]->span), StringImm(new_buf.scope(), buffer_call->args[2]->span)}, @@ -553,13 +553,13 @@ class WarpMemoryRewriter : public StmtExprMutator { private: UnchangedOr Mutate_(const SeqStmtNode* op, InplaceMode inplace_mode) { - // Process SeqStmt to find warp AllocBuffer and gather remaining siblings as body. + // Process SeqStmt to find warp AllocTensor and gather remaining siblings as body. ffi::Array new_seq; bool changed = false; for (size_t i = 0; i < op->seq.size(); ++i) { const auto* alloc = op->seq[i].as(); if (const auto* call = alloc ? alloc->value.as() : nullptr; - call && call->op.same_as(builtin::alloc_buffer()) && + call && call->op.same_as(builtin::alloc_tensor()) && call->args[2].as_or_throw()->value == "warp") { new_storage_scopes_[alloc->var] = "local"; // Gather remaining siblings as the "body" for rewriting. diff --git a/src/tirx/transform/split_host_device.cc b/src/tirx/transform/split_host_device.cc index d7ab2ea998f0..fc2101f6d6ed 100644 --- a/src/tirx/transform/split_host_device.cc +++ b/src/tirx/transform/split_host_device.cc @@ -220,7 +220,7 @@ class HostDeviceSplitter : public StmtExprMutator { std::sort(params.begin(), params.end(), [](const Var& a, const Var& b) { auto sort_key = [](const Var& var) { bool is_handle = - var->ty.as() != nullptr || var->ty.as() != nullptr; + var->ty.as() != nullptr || var->ty.as() != nullptr; return std::tuple{ !is_handle, var->name, @@ -242,13 +242,13 @@ class HostDeviceSplitter : public StmtExprMutator { // Buffer Vars are compiler-side values, not ABI values. Thread their // physical pointer projection through the kernel call and recover the - // typed buffer at the kernel entry with an explicit DeclBuffer source. + // typed buffer at the kernel entry with an explicit DeclTensor source. ffi::Array kernel_params; ffi::Array call_args; ffi::Map buffer_data_params; auto kernel_rewriter = ffi::make_object(); for (const Var& param : params) { - if (param->ty.as()) { + if (param->ty.as()) { BufferVar buffer = param.as_or_throw(); BufferVar kernel_buffer(buffer.name(), buffer.type(), buffer.span()); Var data_param(buffer.name() + "_ptr", buffer.DataPointerType()); @@ -288,7 +288,7 @@ class HostDeviceSplitter : public StmtExprMutator { TVM_FFI_ICHECK(kernel_buffer != nullptr); body = SeqStmt::Flatten( Bind(kernel_buffer.as_or_throw(), - Call(kernel_buffer.as_or_throw().type(), builtin::decl_buffer(), + Call(kernel_buffer.as_or_throw().type(), builtin::decl_tensor(), {data_param.value(), tvm::Tuple(kernel_buffer.as_or_throw()->shape), DataTypeImm(kernel_buffer.as_or_throw()->dtype->dtype), StringImm(kernel_buffer.as_or_throw().scope())}, @@ -487,8 +487,8 @@ class DeviceInfoCollector : public StmtExprVisitor { ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(builtin::alloc_buffer())) - return DispatchAllocBuffer(op, call); + call && call->op.same_as(builtin::alloc_tensor())) + return DispatchAllocTensor(op, call); // Track Bind definitions so that thread_extent values and // dyn_shmem_size expressions that reference locally-bound // variables (e.g. CSE variables) can be inlined back to @@ -554,7 +554,7 @@ class DeviceInfoCollector : public StmtExprVisitor { return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchAllocBuffer(const BindNode* op, const CallNode* call) { + ffi::Optional DispatchAllocTensor(const BindNode* op, const CallNode* call) { ffi::String scope = call->args[2].as_or_throw()->value; auto storage_scope = runtime::StorageScope::Create(scope); if (storage_scope.rank == runtime::StorageRank::kShared && storage_scope.tag == ".dyn") { diff --git a/src/tirx/transform/stmt_simplify.cc b/src/tirx/transform/stmt_simplify.cc index 9e72dc4671a2..9bd15c10069b 100644 --- a/src/tirx/transform/stmt_simplify.cc +++ b/src/tirx/transform/stmt_simplify.cc @@ -119,7 +119,7 @@ PrimFunc StmtSimplifier::Run(PrimFunc func) { } UnchangedOr StmtSimplifier::Mutate(ffi::AnyView input, InplaceMode inplace_mode) { - if (input.as()) { + if (input.as()) { return ffi::Unchanged(); } if (auto expr = input.as()) { @@ -141,8 +141,8 @@ UnchangedOr StmtSimplifier::Mutate_(const ForNode* op, InplaceMode inplace UnchangedOr StmtSimplifier::Mutate_(const BindNode* op, InplaceMode inplace_mode) { if (const auto* call = op->value.as()) { // Preserve buffer metadata and shape operands; only declaration data is simplified. - if (call->op.same_as(builtin::alloc_buffer())) return ffi::Unchanged(); - if (call->op.same_as(builtin::decl_buffer())) { + if (call->op.same_as(builtin::alloc_tensor())) return ffi::Unchanged(); + if (call->op.same_as(builtin::decl_tensor())) { // The Call and its arguments may be shared even when the Bind is writable. auto data = this->Mutate(call->args[0], InplaceMode::kDisallow); if (data.UnchangedOrSameAs(call->args[0])) return ffi::Unchanged(); diff --git a/src/tirx/transform/storage_rewrite.cc b/src/tirx/transform/storage_rewrite.cc index 3442a16a6516..597f3ffb64b1 100644 --- a/src/tirx/transform/storage_rewrite.cc +++ b/src/tirx/transform/storage_rewrite.cc @@ -124,11 +124,11 @@ class LinearAccessPatternFinder final : public StmtExprVisitor { size_t num_physical_dimensions{0}; // scope level size_t level{0}; - // The AllocBuffer node that created this allocation. + // The AllocTensor node that created this allocation. const BindNode* alloc{nullptr}; }; - ffi::Optional DispatchAllocBuffer(const BindNode* op, const CallNode* call) { + ffi::Optional DispatchAllocTensor(const BindNode* op, const CallNode* call) { tvm::Tuple shape = call->args[0].as_or_throw(); size_t level = scope_.size(); const VarNode* buf = op->var.get(); @@ -143,7 +143,7 @@ class LinearAccessPatternFinder final : public StmtExprVisitor { return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchDeclBuffer(const BindNode* op, + ffi::Optional DispatchDeclTensor(const BindNode* op, const CallNode* buffer_call) { RegisterBufferAlias(op->var.as_or_throw(), buffer_call->args[0]); return std::nullopt; @@ -202,7 +202,7 @@ class LinearAccessPatternFinder final : public StmtExprVisitor { ffi::Optional Visit_(const VarNode* buf) final { if (def_region_kind() != kTVMFFIDefRegionKindNone) return StmtExprVisitor::Visit_(buf); // Directly reference to the variable count as a read. - if (buf->ty.as()) { + if (buf->ty.as()) { Var var = ffi::GetRef(buf); buf = buffer_aliases_.Get(var).value_or(var).get(); } @@ -260,8 +260,8 @@ class LinearAccessPatternFinder final : public StmtExprVisitor { ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(builtin::alloc_buffer())) return DispatchAllocBuffer(op, call); - if (call->op.same_as(builtin::decl_buffer())) return DispatchDeclBuffer(op, call); + if (call->op.same_as(builtin::alloc_tensor())) return DispatchAllocTensor(op, call); + if (call->op.same_as(builtin::decl_tensor())) return DispatchDeclTensor(op, call); } scope_.push_back(StmtEntry()); // visit subexpr (the value may contain BufferLoad) @@ -292,17 +292,17 @@ class LinearAccessPatternFinder final : public StmtExprVisitor { std::vector linear_seq_; // The storage scope of each buffer std::unordered_map alloc_info_; - // Physical roots of buffer aliases, flattened when each DeclBuffer is visited. + // Physical roots of buffer aliases, flattened when each DeclTensor is visited. ffi::Map buffer_aliases_; private: void RegisterBufferAlias(BufferVar buffer, const Expr& data) { Var root = buffer.var(); if (auto source = GetBufferDataVar(data); - source.has_value() && source.value()->ty.as()) { + source.has_value() && source.value()->ty.as()) { auto source_root = buffer_aliases_.Get(source.value()); TVM_FFI_ICHECK(source_root.has_value()) << "Buffer alias source " << source.value()->name - << " must be registered before its DeclBuffer alias"; + << " must be registered before its DeclTensor alias"; root = source_root.value(); } buffer_aliases_.Set(buffer.var(), root); @@ -397,12 +397,12 @@ class InplaceOpVerifier : public StmtExprVisitor { ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(builtin::alloc_buffer())) - return DispatchAllocBuffer(op, call); + call && call->op.same_as(builtin::alloc_tensor())) + return DispatchAllocTensor(op, call); return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchAllocBuffer(const BindNode* op, + ffi::Optional DispatchAllocTensor(const BindNode* op, const CallNode* buffer_call) { // reject inplace for volatile buffers if (buffer_call->attrs.as()->dict.count(attr::kVolatile)) { @@ -550,7 +550,7 @@ class StoragePlanRewriter : public StmtExprMutator { BufferVar backing = new_backing_array.as_or_throw(); BufferVar remapped = buf.same_as(backing) ? buf - : RebuildBufferVar(buf, CopyBufferType(buf), new_backing_array->name); + : RebuildBufferVar(buf, CopyTensorType(buf), new_backing_array->name); VarRemapSet(buf, remapped); remapped_backing_[remapped.get()] = backing; return remapped; @@ -582,7 +582,7 @@ class StoragePlanRewriter : public StmtExprMutator { UnchangedOr Mutate_(const VarNode* op, InplaceMode inplace_mode) final { if (def_region_kind() == kTVMFFIDefRegionKindNone) { const VarNode* root = op; - if (op->ty.as()) { + if (op->ty.as()) { Var var = ffi::GetRef(op); root = buffer_aliases_.Get(var).value_or(var).get(); } @@ -643,7 +643,7 @@ class StoragePlanRewriter : public StmtExprMutator { return StmtExprMutator::Mutate_(op, inplace_mode); } const VarNode* buffer = buffer_var.value().get(); - if (buffer->ty.as()) { + if (buffer->ty.as()) { Var var = buffer_var.value(); buffer = buffer_aliases_.Get(var).value_or(var).get(); } @@ -705,29 +705,29 @@ class StoragePlanRewriter : public StmtExprMutator { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(builtin::alloc_buffer())) - return MutateAllocBuffer(op, call, inplace_mode); - if (call->op.same_as(builtin::decl_buffer())) return MutateDeclBuffer(op, call, inplace_mode); + if (call->op.same_as(builtin::alloc_tensor())) + return MutateAllocTensor(op, call, inplace_mode); + if (call->op.same_as(builtin::decl_tensor())) return MutateDeclTensor(op, call, inplace_mode); } return StmtExprMutator::Mutate_(op, inplace_mode); } - UnchangedOr MutateAllocBuffer(const BindNode* op, const CallNode* buffer_call, + UnchangedOr MutateAllocTensor(const BindNode* op, const CallNode* buffer_call, InplaceMode inplace_mode) { - // AllocBuffer combines allocation and buffer declaration. + // AllocTensor combines allocation and buffer declaration. // Storage rewrite may merge this allocation with others. if (auto it = alloc_map_.find(op->var.get()); it != alloc_map_.end()) { if (it->second->alloc_var.get() == op->var.get()) { - // This is the "winner" allocation -- its AllocBuffer was already + // This is the "winner" allocation -- its AllocTensor was already // hoisted by PrepareNewAlloc. Strip this one. return Evaluate(0); } - // This allocation was merged into another. Emit a DeclBuffer + // This allocation was merged into another. Emit a DeclTensor // aliasing the winner's data variable. BufferVar buf = RemapBuffer(op->var.as_or_throw(), it->second->alloc_var); return Bind( buf, - Call(buf.type(), builtin::decl_buffer(), + Call(buf.type(), builtin::decl_tensor(), {it->second->alloc_var.as_or_throw().data(), tvm::Tuple(buf->shape), DataTypeImm(buf->dtype->dtype), StringImm(buf.scope())}, {}, buffer_call->ty_args, buffer_call->span), @@ -737,7 +737,7 @@ class StoragePlanRewriter : public StmtExprMutator { return Evaluate(0); } - UnchangedOr MutateDeclBuffer(const BindNode* op, const CallNode* buffer_call, + UnchangedOr MutateDeclTensor(const BindNode* op, const CallNode* buffer_call, InplaceMode inplace_mode) { const VarNode* root = buffer_aliases_.Get(op->var).value_or(op->var).get(); auto it = alloc_map_.find(root); @@ -745,7 +745,7 @@ class StoragePlanRewriter : public StmtExprMutator { BufferVar buffer = RemapBuffer(op->var.as_or_throw(), it->second->alloc_var); return Bind( buffer, - Call(buffer.type(), builtin::decl_buffer(), + Call(buffer.type(), builtin::decl_tensor(), {it->second->alloc_var.as_or_throw().data(), tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())}, buffer_call->attrs, buffer_call->ty_args, buffer_call->span), @@ -772,7 +772,7 @@ class StoragePlanRewriter : public StmtExprMutator { std::vector allocs; // The children of this entry, not including itself. std::vector merged_children; - // The replacement AllocBuffer, if any. + // The replacement AllocTensor, if any. std::vector alloc_nest; // The var expr of new allocation. Var alloc_var; @@ -782,7 +782,7 @@ class StoragePlanRewriter : public StmtExprMutator { bool is_volatile{false}; // Fragment metadata belongs to this allocation, not interchangeable storage. bool has_fragment_metadata{false}; - // This is non-zero if this alloc_buffer is folded into another one + // This is non-zero if this alloc_tensor is folded into another one // the address(in bits) becomes alloc_var + bits_offset; // can be effectively converted to the element type. // We need to convert bit_offset to offset of specific element type later. @@ -884,7 +884,7 @@ class StoragePlanRewriter : public StmtExprMutator { } e->alloc_nest.push_back(Bind( buf.var(), - Call(buf.type(), tirx::builtin::alloc_buffer(), + Call(buf.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(buf->shape), DataTypeImm(buf->dtype->dtype), StringImm(buf.scope())}, DictAttrs(annotations)))); continue; @@ -924,7 +924,7 @@ class StoragePlanRewriter : public StmtExprMutator { }); if (all_allocs_identical) { - // Emit AllocBuffer for the hoisted allocation. + // Emit AllocTensor for the hoisted allocation. BufferVar buf = RemapBuffer(e->allocs[0]->var.as_or_throw(), e->alloc_var); ffi::Map annotations; if (e->is_volatile) { @@ -932,7 +932,7 @@ class StoragePlanRewriter : public StmtExprMutator { } e->alloc_nest.push_back(Bind( buf.var(), - Call(buf.type(), tirx::builtin::alloc_buffer(), + Call(buf.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(buf->shape), DataTypeImm(buf->dtype->dtype), StringImm(buf.scope())}, DictAttrs(annotations)))); } else { @@ -977,7 +977,7 @@ class StoragePlanRewriter : public StmtExprMutator { combo_size = combo_size + IntImm::Int32(1); } combo_size = analyzer_->Simplify(combo_size); - BufferVar buf(e->alloc_var->name, BufferType(e->scope.to_string(), alloc_type, + BufferVar buf(e->alloc_var->name, TensorType(e->scope.to_string(), alloc_type, {combo_size}, {}, PrimExpr(), 0, 0)); e->alloc_var = buf.var(); ffi::Map annotations; @@ -986,7 +986,7 @@ class StoragePlanRewriter : public StmtExprMutator { } e->alloc_nest.push_back(Bind( buf.var(), - Call(buf.type(), tirx::builtin::alloc_buffer(), + Call(buf.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(buf->shape), DataTypeImm(buf->dtype->dtype), StringImm(buf.scope())}, DictAttrs(annotations)))); } @@ -1022,7 +1022,7 @@ class StoragePlanRewriter : public StmtExprMutator { tvm::Tuple shape = call->args[0].as_or_throw(); PrimExpr alloc_size = MakeConst(shape->fields[0].as_or_throw().ty(), (total_bits + type_bits - 1) / type_bits); - BufferVar buf(e->alloc_var->name, BufferType(e->scope.to_string(), e->elem_type, {alloc_size}, + BufferVar buf(e->alloc_var->name, TensorType(e->scope.to_string(), e->elem_type, {alloc_size}, {}, PrimExpr(), 0, 0)); e->alloc_var = buf.var(); for (StorageEntry* child : e->merged_children) { @@ -1038,7 +1038,7 @@ class StoragePlanRewriter : public StmtExprMutator { } e->alloc_nest.push_back( Bind(buf.var(), - Call(buf.type(), tirx::builtin::alloc_buffer(), + Call(buf.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(buf->shape), DataTypeImm(buf->dtype->dtype), StringImm(buf.scope())}, DictAttrs(annotations)))); } @@ -1135,7 +1135,7 @@ class StoragePlanRewriter : public StmtExprMutator { src_entry->elem_type == element_type.WithLanes(1) && visitor->Check(s.stmt, var, src)) { int64_t const_size = - alloc->var->ty.as_or_throw()->ConstantAllocationSize().value_or( + alloc->var->ty.as_or_throw()->ConstantAllocationSize().value_or( 0); uint64_t const_nbits = static_cast(const_size) * element_type.bits() * element_type.lanes(); @@ -1223,7 +1223,7 @@ class StoragePlanRewriter : public StmtExprMutator { bool is_scalable_vector = element_type.IsScalableVector(); uint64_t op_elem_bits = is_scalable_vector ? 0 : element_type.bits() * element_type.lanes(); int64_t const_size = - op->var->ty.as_or_throw()->ConstantAllocationSize().value_or(0); + op->var->ty.as_or_throw()->ConstantAllocationSize().value_or(0); uint64_t const_nbits = is_scalable_vector ? 0 : static_cast(const_size * op_elem_bits); @@ -1351,9 +1351,9 @@ struct BufferVarInfo { enum DeclarationLocation { kPrimFuncBufferParam = (1 << 0), kPrimFuncPointerParam = (1 << 1), - kAllocBufferCall = (1 << 2), + kAllocTensorCall = (1 << 2), kLetNode = (1 << 3), - kDeclBufferCall = (1 << 4), + kDeclTensorCall = (1 << 4), }; // The tirx::Var that represents this buffer. @@ -1365,7 +1365,7 @@ struct BufferVarInfo { /* The extent of the buffer. * * If multidimensional, the extent of the last dimension of the buffer. If the - * size is unknown (e.g. pointer arguments to PrimFunc without a BufferType + * size is unknown (e.g. pointer arguments to PrimFunc without a TensorType * annotation), then extent is zero. */ PrimExpr extent; @@ -1510,19 +1510,19 @@ class VectorTypeAccessChecker : public StmtExprVisitor { return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchAllocBuffer(const BindNode* op, const CallNode* call) { + ffi::Optional DispatchAllocTensor(const BindNode* op, const CallNode* call) { tvm::Tuple shape = call->args[0].as_or_throw(); DLDataType dtype = call->args[1].as_or_throw()->value; PrimType element_type(dtype); buffer_aliases_.Set(op->var, op->var); PrimExpr extent = shape->fields.size() ? shape->fields.back().as_or_throw() : PrimExpr(0); - OnArrayDeclaration(op->var, element_type, extent, BufferVarInfo::kAllocBufferCall); + OnArrayDeclaration(op->var, element_type, extent, BufferVarInfo::kAllocTensorCall); return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchDeclBuffer(const BindNode* op, + ffi::Optional DispatchDeclTensor(const BindNode* op, const CallNode* buffer_call) { RegisterBufferAlias(op->var.as_or_throw(), buffer_call->args[0]); tvm::Tuple shape = buffer_call->args[1].as_or_throw(); @@ -1530,7 +1530,7 @@ class VectorTypeAccessChecker : public StmtExprVisitor { PrimType element_type(dtype); PrimExpr extent = shape->fields.size() ? shape->fields.back().as_or_throw() : PrimExpr(0); - OnArrayDeclaration(op->var, element_type, extent, BufferVarInfo::kDeclBufferCall); + OnArrayDeclaration(op->var, element_type, extent, BufferVarInfo::kDeclTensorCall); return StmtExprVisitor::Visit_(op); } @@ -1541,8 +1541,8 @@ class VectorTypeAccessChecker : public StmtExprVisitor { ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(builtin::alloc_buffer())) return DispatchAllocBuffer(op, call); - if (call->op.same_as(builtin::decl_buffer())) return DispatchDeclBuffer(op, call); + if (call->op.same_as(builtin::alloc_tensor())) return DispatchAllocTensor(op, call); + if (call->op.same_as(builtin::decl_tensor())) return DispatchDeclTensor(op, call); } HandleLetNode(op->var); return StmtExprVisitor::Visit_(op); @@ -1694,10 +1694,10 @@ class VectorTypeAccessChecker : public StmtExprVisitor { void RegisterBufferAlias(BufferVar buffer, const Expr& data) { Var root = buffer.var(); if (auto source = GetBufferDataVar(data); - source.has_value() && source.value()->ty.as()) { + source.has_value() && source.value()->ty.as()) { auto source_root = buffer_aliases_.Get(source.value()); TVM_FFI_ICHECK(source_root.has_value()) << "Buffer alias source " << source.value()->name - << " must be registered before its DeclBuffer alias"; + << " must be registered before its DeclTensor alias"; root = source_root.value(); } buffer_aliases_.Set(buffer.var(), root); @@ -1752,7 +1752,7 @@ class VectorTypeRewriter : public StmtExprMutator { * @param checker The VectorTypeAccessChecker that has previously read out * information from the PrimFunc * - * @param rewrite_buffer_params Whether BufferType-annotated parameters should + * @param rewrite_buffer_params Whether TensorType-annotated parameters should * be rewritten from scalar element types to vectorized element types. * * @param rewrite_pointer_params Whether pointer-typed parameters should be @@ -1781,7 +1781,7 @@ class VectorTypeRewriter : public StmtExprMutator { rewrite_mask |= BufferVarInfo::kPrimFuncPointerParam; } if (rewrite_alloc_buffer_node) { - rewrite_mask |= BufferVarInfo::kAllocBufferCall; + rewrite_mask |= BufferVarInfo::kAllocTensorCall; } if (rewrite_let_node) { rewrite_mask |= BufferVarInfo::kLetNode; @@ -1793,9 +1793,9 @@ class VectorTypeRewriter : public StmtExprMutator { if (preferred != var_info.element_dtype && (rewrite_mask & var_info.declaration_location)) { Var old_buffer_var = var_info.var; Var new_buffer_var = [&]() -> Var { - if (old_buffer_var->ty.as()) { + if (old_buffer_var->ty.as()) { BufferVar old_buffer = old_buffer_var.as_or_throw(); - auto type = CopyBufferType(old_buffer); + auto type = CopyTensorType(old_buffer); type->dtype = preferred; if (!type->shape.empty()) { PrimExpr last_dim = type->shape.back(); @@ -1988,9 +1988,9 @@ class VectorTypeRewriter : public StmtExprMutator { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(builtin::alloc_buffer())) - return MutateAllocBuffer(op, call, inplace_mode); - if (call->op.same_as(builtin::decl_buffer())) return MutateDeclBuffer(op, call, inplace_mode); + if (call->op.same_as(builtin::alloc_tensor())) + return MutateAllocTensor(op, call, inplace_mode); + if (call->op.same_as(builtin::decl_tensor())) return MutateDeclTensor(op, call, inplace_mode); } auto it = rewrite_map_.find(op->var.get()); auto value_result = this->Mutate(op->value, inplace_mode); @@ -2008,14 +2008,14 @@ class VectorTypeRewriter : public StmtExprMutator { return Bind(var, value); } - UnchangedOr MutateAllocBuffer(const BindNode* op, const CallNode* buffer_call, + UnchangedOr MutateAllocTensor(const BindNode* op, const CallNode* buffer_call, InplaceMode inplace_mode) { BufferVar new_buf = RemapBuffer(op->var.as_or_throw()); if (new_buf.same_as(op->var.as_or_throw())) { return ffi::Unchanged(); } return Bind(new_buf.var(), - Call(new_buf.type(), tirx::builtin::alloc_buffer(), + Call(new_buf.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(new_buf->shape, buffer_call->args[0]->span), DataTypeImm(new_buf->dtype->dtype, buffer_call->args[1]->span), StringImm(new_buf.scope(), buffer_call->args[2]->span)}, @@ -2023,14 +2023,14 @@ class VectorTypeRewriter : public StmtExprMutator { op->span); } - UnchangedOr MutateDeclBuffer(const BindNode* op, const CallNode* buffer_call, + UnchangedOr MutateDeclTensor(const BindNode* op, const CallNode* buffer_call, InplaceMode inplace_mode) { Expr data = Mutate(buffer_call->args[0], inplace_mode).ValueOrUnchanged(buffer_call->args[0]); BufferVar buffer = RemapBuffer(op->var.as_or_throw()); if (buffer.same_as(op->var.as_or_throw()) && data.same_as(buffer_call->args[0])) return ffi::Unchanged(); return Bind(buffer, - Call(buffer.type(), builtin::decl_buffer(), + Call(buffer.type(), builtin::decl_tensor(), {data, tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())}, buffer_call->attrs, buffer_call->ty_args, buffer_call->span), @@ -2049,7 +2049,7 @@ class VectorTypeRewriter : public StmtExprMutator { if (root.same_as(buf.var())) { buf = info.new_buffer_var.as_or_throw(); } else { - auto type = CopyBufferType(buf); + auto type = CopyTensorType(buf); type->dtype = info.new_element_dtype; if (!type->shape.empty()) { PrimExpr last_dim = type->shape.back(); @@ -2071,7 +2071,7 @@ class VectorTypeRewriter : public StmtExprMutator { } if (op->op.same_as(builtin::buffer_data()) && op->args.size() == 1) { if (auto var = op->args[0].as(); - var.has_value() && var.value()->ty.as()) { + var.has_value() && var.value()->ty.as()) { return RemapBuffer(var.value().as_or_throw()).data(); } } @@ -2104,7 +2104,7 @@ class VectorTypeRewriter : public StmtExprMutator { int factor = info.factor(); extent = extent / MakeConst(extent.ty(), factor); index = index / MakeConst(index.ty(), factor); - Expr data = info.new_buffer_var->ty.as() + Expr data = info.new_buffer_var->ty.as() ? info.new_buffer_var.as_or_throw().data() : Expr(info.new_buffer_var); ffi::Array acc_args{e_dtype, data, index, extent, flag}; diff --git a/src/tirx/transform/tile_primitive_dispatch.cc b/src/tirx/transform/tile_primitive_dispatch.cc index df6350bd7ea9..4b6df950d09c 100644 --- a/src/tirx/transform/tile_primitive_dispatch.cc +++ b/src/tirx/transform/tile_primitive_dispatch.cc @@ -310,7 +310,7 @@ class TilePrimitiveDispatcher : public StmtExprMutator { seq.reserve(alloc_buffers_.size() + 1); for (const auto& buffer : alloc_buffers_) { seq.push_back( - Bind(buffer.var(), Call(buffer.type(), tirx::builtin::alloc_buffer(), + Bind(buffer.var(), Call(buffer.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())}, DictAttrs()))); @@ -389,8 +389,8 @@ class TilePrimitiveDispatcher : public StmtExprMutator { rebuilt.push_back(s); if (const auto* bind = s.as()) { if (const auto* call = bind->value.as(); - call && (call->op.same_as(builtin::alloc_buffer()) || - call->op.same_as(builtin::decl_buffer()))) { + call && (call->op.same_as(builtin::alloc_tensor()) || + call->op.same_as(builtin::decl_tensor()))) { changed |= AppendPostBufferDefStmts(&rebuilt, bind->var.as_or_throw(), bind->var.as_or_throw()); } @@ -404,8 +404,8 @@ class TilePrimitiveDispatcher : public StmtExprMutator { UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { if (const auto* call = op->value.as(); call) { - if (call->op.same_as(builtin::alloc_buffer())) return MutateAllocBuffer(op, inplace_mode); - if (call->op.same_as(builtin::decl_buffer())) return MutateDeclBuffer(op, inplace_mode); + if (call->op.same_as(builtin::alloc_tensor())) return MutateAllocTensor(op, inplace_mode); + if (call->op.same_as(builtin::decl_tensor())) return MutateDeclTensor(op, inplace_mode); } Stmt stmt = StmtExprMutator::Mutate_(op, inplace_mode).ValueOrUnchanged(ffi::GetRef(op)); const auto* bind = stmt.as(); @@ -431,7 +431,7 @@ class TilePrimitiveDispatcher : public StmtExprMutator { /*! * \brief Track the storage root of a buffer variable. * - * A ``DeclBuffer`` whose data is ``buffer_data(src)`` is a view over + * A ``DeclTensor`` whose data is ``buffer_data(src)`` is a view over * ``src``'s storage, so it inherits ``src``'s root; anything else owns its * storage. Buffers with no definition in the body (PrimFunc parameters) * are absent from the map and are their own root. @@ -443,7 +443,7 @@ class TilePrimitiveDispatcher : public StmtExprMutator { if (const auto* call = data.value().as(); call && call->op.same_as(builtin::buffer_data()) && call->args.size() == 1) { if (auto src = call->args[0].as(); - src.has_value() && src.value()->ty.as()) { + src.has_value() && src.value()->ty.as()) { root = StorageRootOf(src.value()); } } @@ -477,7 +477,7 @@ class TilePrimitiveDispatcher : public StmtExprMutator { UnchangedOr Mutate_(const CallNode* op, InplaceMode inplace_mode) final { if (op->op.same_as(builtin::buffer_data()) && op->args.size() == 1) { if (auto var = op->args[0].as(); - var.has_value() && var.value()->ty.as()) { + var.has_value() && var.value()->ty.as()) { auto it = buffer_root_.find(var.value()); if (it != buffer_root_.end() && !it->second.same_as(var.value())) { return it->second.as_or_throw().data(); @@ -490,7 +490,7 @@ class TilePrimitiveDispatcher : public StmtExprMutator { const std::unordered_map& buffer_root_; }; - UnchangedOr MutateAllocBuffer(const BindNode* op, InplaceMode inplace_mode) { + UnchangedOr MutateAllocTensor(const BindNode* op, InplaceMode inplace_mode) { BufferVar old_buffer = op->var.as_or_throw(); Stmt stmt = StmtExprMutator::Mutate_(op, inplace_mode).ValueOrUnchanged(ffi::GetRef(op)); op = stmt.as(); @@ -502,7 +502,7 @@ class TilePrimitiveDispatcher : public StmtExprMutator { return SeqStmt::Flatten(seq); } - UnchangedOr MutateDeclBuffer(const BindNode* op, InplaceMode inplace_mode) { + UnchangedOr MutateDeclTensor(const BindNode* op, InplaceMode inplace_mode) { BufferVar old_buffer = op->var.as_or_throw(); Stmt stmt = StmtExprMutator::Mutate_(op, inplace_mode).ValueOrUnchanged(ffi::GetRef(op)); op = stmt.as(); diff --git a/src/tirx/transform/tvm_ffi_binder.cc b/src/tirx/transform/tvm_ffi_binder.cc index b11666ecbe46..97705413a024 100644 --- a/src/tirx/transform/tvm_ffi_binder.cc +++ b/src/tirx/transform/tvm_ffi_binder.cc @@ -494,7 +494,7 @@ Expr TVMFFIABIBuilder::LoadTVMFFIAnyUnionValue(const Var& v_packed_args, int par Expr TVMFFIABIBuilder::DecodeParamOpaqueHandle(int param_index, const PrimExpr& type_index) { // ── Type check: accept handle-like types ─────────────────── - std::string expected_type = params_[param_index]->ty.as() ? "Tensor" : "pointer"; + std::string expected_type = params_[param_index]->ty.as() ? "Tensor" : "pointer"; EmitTypeIndexCheck(param_index, type_index == ffi::TypeIndex::kTVMFFINone || type_index == ffi::TypeIndex::kTVMFFIOpaquePtr || @@ -565,7 +565,7 @@ void TVMFFIABIBuilder::DecodeParam(int param_index) { ffi::reflection::AccessPath param_path = ffi::reflection::AccessPath::Root()->Extend(AccessStep::ArrayItem(param_index)); - if (param->ty.as()) { + if (param->ty.as()) { Var handle(param->name + ".handle", PointerType::VoidPointerTy()); Expr handle_value = DecodeParamOpaqueHandle(param_index, type_index.as_or_throw()); BindPointer(handle, handle_value, param_path, true); @@ -622,7 +622,7 @@ void TVMFFIABIBuilder::DecodeAllParams() { func_name_ + "." + param->name, param_path); decl_buffers_.push_back( Bind(buffer.value(), - Call(buffer.value().type(), builtin::decl_buffer(), + Call(buffer.value().type(), builtin::decl_tensor(), {data, tvm::Tuple(buffer.value()->shape), DataTypeImm(buffer.value()->dtype->dtype), StringImm(buffer.value().scope())}, {}))); diff --git a/src/tirx/transform/tvm_ffi_binder.h b/src/tirx/transform/tvm_ffi_binder.h index 4f51955a1234..0d6e26d792b7 100644 --- a/src/tirx/transform/tvm_ffi_binder.h +++ b/src/tirx/transform/tvm_ffi_binder.h @@ -59,10 +59,10 @@ namespace tirx { * by a later buffer's shape (batch_size). Separating definitions from * checks guarantees all variables are in scope when assertions reference them. * - * - init_nest: Binds, DeclBuffers for shape/strides arrays, AttrStmts — + * - init_nest: Binds, DeclTensors for shape/strides arrays, AttrStmts — * all value-loading code that defines variables. * - asserts: AssertStmts — all validation checks. - * - decl_buffers: DeclBuffer for buffer-typed parameters — buffer declarations. + * - decl_buffers: DeclTensor for buffer-typed parameters — buffer declarations. * * ## Calling Protocol * @@ -96,7 +96,7 @@ class TVMFFIABIBuilder { struct Result { /*! \brief Var -> VarDefInfo map for defined variables. */ std::unordered_map var_defs; - /*! \brief Variable definitions (Binds, shape/strides DeclBuffers, AttrStmts). */ + /*! \brief Variable definitions (Binds, shape/strides DeclTensors, AttrStmts). */ std::vector init_nest; /*! \brief Validation checks (all AssertStmts). */ std::vector asserts; @@ -385,7 +385,7 @@ class TVMFFIABIBuilder { /*! \brief The definition map: VarNode* -> VarDefInfo (value + first_def_path). */ std::unordered_map var_defs_; - /*! \brief Variable definitions: Binds, shape/strides DeclBuffers, AttrStmts. */ + /*! \brief Variable definitions: Binds, shape/strides DeclTensors, AttrStmts. */ std::vector init_nest_; /*! \brief Validation checks: all AssertStmts. */ std::vector asserts_; diff --git a/src/tirx/transform/unsupported_dtype_legalize.cc b/src/tirx/transform/unsupported_dtype_legalize.cc index 175ba6407af7..d738e25f2f72 100644 --- a/src/tirx/transform/unsupported_dtype_legalize.cc +++ b/src/tirx/transform/unsupported_dtype_legalize.cc @@ -85,12 +85,12 @@ class ComputeLegalizePlanner : public StmtExprVisitor { ffi::Optional Visit_(const BindNode* op) final { if (const auto* call = op->value.as(); - call && call->op.same_as(builtin::alloc_buffer())) - return DispatchAllocBuffer(op, call); + call && call->op.same_as(builtin::alloc_tensor())) + return DispatchAllocTensor(op, call); return StmtExprVisitor::Visit_(op); } - ffi::Optional DispatchAllocBuffer(const BindNode* op, const CallNode* call) { + ffi::Optional DispatchAllocTensor(const BindNode* op, const CallNode* call) { DLDataType dtype = call->args[1].as_or_throw()->value; PrimType alloc_dtype(dtype); // Select intermediate buffers with an unsupported element type. @@ -230,7 +230,7 @@ class ComputeLegalizer : public StmtExprMutator { } UnchangedOr Mutate_(const CallNode* op, InplaceMode inplace_mode) final { - if (op->op.same_as(builtin::alloc_buffer()) || op->op.same_as(builtin::decl_buffer())) { + if (op->op.same_as(builtin::alloc_tensor()) || op->op.same_as(builtin::decl_tensor())) { Call call = StmtExprMutator::Mutate_(op, inplace_mode) .ValueOrUnchanged(ffi::GetRef(op)) .as_or_throw(); @@ -532,7 +532,7 @@ class StorageLegalizer : public StmtExprMutator { using StmtExprMutator::Mutate_; PrimFunc Legalize(PrimFunc func) { for (const Var& param : func->params) { - TVM_FFI_ICHECK(!param->ty.as()) + TVM_FFI_ICHECK(!param->ty.as()) << "This pass must be called after MakePackedAPI"; } auto* n = func.CopyOnWrite(); @@ -559,7 +559,7 @@ class StorageLegalizer : public StmtExprMutator { } UnchangedOr Mutate_(const BindNode* op, InplaceMode inplace_mode) final { - if (op->value->ty.as()) { + if (op->value->ty.as()) { return StmtExprMutator::Mutate_(op, inplace_mode); } auto value_result = Mutate(op->value, inplace_mode); @@ -616,11 +616,11 @@ class StorageLegalizer : public StmtExprMutator { } UnchangedOr Mutate_(const CallNode* op, InplaceMode inplace_mode) final { - if (op->op.same_as(builtin::alloc_buffer()) || op->op.same_as(builtin::decl_buffer())) { + if (op->op.same_as(builtin::alloc_tensor()) || op->op.same_as(builtin::decl_tensor())) { Call call = StmtExprMutator::Mutate_(op, inplace_mode) .ValueOrUnchanged(ffi::GetRef(op)) .as_or_throw(); - int dtype_index = op->op.same_as(builtin::alloc_buffer()) ? 1 : 2; + int dtype_index = op->op.same_as(builtin::alloc_tensor()) ? 1 : 2; auto dtype = call->args[dtype_index].as_or_throw(); if (MatchType(PrimType(dtype->value))) { call.CopyOnWrite()->args.Set( diff --git a/src/tirx/transform/update_pointer_storage_scope.cc b/src/tirx/transform/update_pointer_storage_scope.cc index 3051554e3975..1bb70796b4ed 100644 --- a/src/tirx/transform/update_pointer_storage_scope.cc +++ b/src/tirx/transform/update_pointer_storage_scope.cc @@ -51,9 +51,9 @@ UpdatePointerStorageScope::UpdatePointerStorageScope( const std::unordered_map& new_storage_scopes) { for (auto& kv : new_storage_scopes) { - if (kv.first->ty.as()) { + if (kv.first->ty.as()) { BufferVar buffer = GetBufferVar(kv.first.get()); - auto type = CopyBufferType(buffer); + auto type = CopyTensorType(buffer); type->storage_scope = kv.second; BufferVar replacement = RebuildBufferVar(buffer, std::move(type)); VarRemapSet(kv.first, replacement); @@ -66,7 +66,7 @@ UpdatePointerStorageScope::UpdatePointerStorageScope( UnchangedOr UpdatePointerStorageScope::Mutate_(const BindNode* op, InplaceMode inplace_mode) { const auto* call = op->value.as(); if (call && - (call->op.same_as(builtin::alloc_buffer()) || call->op.same_as(builtin::decl_buffer()))) { + (call->op.same_as(builtin::alloc_tensor()) || call->op.same_as(builtin::decl_tensor()))) { if (auto mapped = VarRemapGet(op->var); mapped != nullptr) { buffer_scopes_.emplace(call, mapped.as_or_throw().scope()); auto result = StmtExprMutator::Mutate_(op, inplace_mode); @@ -82,7 +82,7 @@ UnchangedOr UpdatePointerStorageScope::Mutate_(const CallNode* op, Inplace if (auto it = buffer_scopes_.find(op); it != buffer_scopes_.end()) { Expr value = std::move(result).ValueOrUnchanged(ffi::GetRef(op)); auto call = value.as_or_throw(); - size_t scope_index = call->op.same_as(builtin::alloc_buffer()) ? 2 : 3; + size_t scope_index = call->op.same_as(builtin::alloc_tensor()) ? 2 : 3; if (call->args[scope_index].as_or_throw()->value != it->second) { auto copy = ffi::make_object(*call.get()); copy->args.Set(scope_index, StringImm(it->second, call->args[scope_index]->span)); diff --git a/src/tirx/transform/vectorize_loop.cc b/src/tirx/transform/vectorize_loop.cc index 6c631240bc84..4973cbb145d4 100644 --- a/src/tirx/transform/vectorize_loop.cc +++ b/src/tirx/transform/vectorize_loop.cc @@ -399,7 +399,7 @@ class VecAllocAccess : public StmtExprMutator { } // Copy everything into the new buffer. - auto type = CopyBufferType(node->buffer); + auto type = CopyTensorType(node->buffer); type->shape = shape; type->strides = strides; buf = RebuildBufferVar(node->buffer, std::move(type)); @@ -434,7 +434,7 @@ class VecAllocAccess : public StmtExprMutator { if (i + 1 != strides.size()) stride *= var_lanes_; strides.Set(i, analyzer_->Simplify(stride)); } - auto type = CopyBufferType(buffer); + auto type = CopyTensorType(buffer); type->shape = shape; type->strides = strides; buf = RebuildBufferVar(buffer, std::move(type)); diff --git a/tests/cpp/ir_functor_test.cc b/tests/cpp/ir_functor_test.cc index 468df824050b..f60aa11f8950 100644 --- a/tests/cpp/ir_functor_test.cc +++ b/tests/cpp/ir_functor_test.cc @@ -209,7 +209,7 @@ TEST(IRF, StmtVisitor) { // implementation ffi::Optional Visit_(const VarNode* op) final { // Buffer variables now share this hook; this fixture counts other Var operands. - if (!op->ty.as()) ++count; + if (!op->ty.as()) ++count; return std::nullopt; } }; @@ -218,16 +218,16 @@ TEST(IRF, StmtVisitor) { auto z = x + 1; Stmt eval_body = Evaluate(z); PrimType dtype = PrimType::Float(32); - BufferVar buf("b", BufferType("global", dtype, {z, z}, {}, PrimExpr(), 0, 0)); - // AllocBuffer is flat (no body). Return as SeqStmt with eval. - return SeqStmt({Bind(buf.var(), Call(buf.type(), tirx::builtin::alloc_buffer(), + BufferVar buf("b", TensorType("global", dtype, {z, z}, {}, PrimExpr(), 0, 0)); + // AllocTensor is flat (no body). Return as SeqStmt with eval. + return SeqStmt({Bind(buf.var(), Call(buf.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(buf->shape), DataTypeImm(buf->dtype->dtype), StringImm(buf.scope())}, DictAttrs())), eval_body}); }; v->Visit(fmaketest()); - // AllocBuffer visits buffer shape at its definition site. + // AllocTensor visits buffer shape at its definition site. // shape = {z, z} where z = x + 1, so x is visited twice from shape + once from eval = 3 TVM_FFI_ICHECK_EQ(v->count, 3); @@ -236,14 +236,14 @@ TEST(IRF, StmtVisitor) { Stmt body = fmaketest(); PrimType dtype = PrimType::Float(32); tirx::Var buf_var("b", PointerType(dtype)); - BufferVar buffer = decl_buffer({16}); + BufferVar buffer = decl_tensor({16}); body = - SeqStmt({Bind(buffer, Call(buffer.type(), tvm::tirx::builtin::decl_buffer(), + SeqStmt({Bind(buffer, Call(buffer.type(), tvm::tirx::builtin::decl_tensor(), {buf_var, tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())})), std::move(body)}); TensorRegion buffer_region = BufferRegion(buffer, {Range::FromMinExtent(x + 1, 1)}); - s_tir::MatchBufferRegion match_buffer_region(decl_buffer({1}), buffer_region); + s_tir::MatchBufferRegion match_buffer_region(decl_tensor({1}), buffer_region); // construct block and block_realize s_tir::SBlock block = s_tir::SBlock({}, {buffer_region}, {buffer_region}, "block", body, body, @@ -253,8 +253,8 @@ TEST(IRF, StmtVisitor) { v->count = 0; v->Visit(block_realize); // x visited in: reads range (1), writes range (1), match_buffers range (1). - // init: DeclBuffer data b(1) + AllocBuffer shape x,x(2) + Evaluate x(1) = 4. - // body: DeclBuffer data b(1) + AllocBuffer shape x,x(2) + Evaluate x(1) = 4. + // init: DeclTensor data b(1) + AllocTensor shape x,x(2) + Evaluate x(1) = 4. + // body: DeclTensor data b(1) + AllocTensor shape x,x(2) + Evaluate x(1) = 4. // Total: 1 + 1 + 1 + 4 + 4 = 11. TVM_FFI_ICHECK_EQ(v->count, 11); } @@ -273,8 +273,8 @@ TEST(IRF, StmtExprMutator) { auto fmakealloc = [&]() { auto z = x + 1; PrimType dtype = PrimType::Float(32); - BufferVar buf("b", BufferType("global", dtype, {1, z}, {}, PrimExpr(), 0, 0)); - return Bind(buf.var(), Call(buf.type(), tirx::builtin::alloc_buffer(), + BufferVar buf("b", TensorType("global", dtype, {1, z}, {}, PrimExpr(), 0, 0)); + return Bind(buf.var(), Call(buf.type(), tirx::builtin::alloc_tensor(), {tvm::Tuple(buf->shape), DataTypeImm(buf->dtype->dtype), StringImm(buf.scope())}, DictAttrs())); @@ -295,7 +295,7 @@ TEST(IRF, StmtExprMutator) { auto* arrptr = arr.get(); arr.MutateByApply([&](Stmt s) { return v->Mutate(s).ValueOrUnchanged(std::move(s)); }); TVM_FFI_ICHECK(arr.get() == arrptr); - // buffer IS mutated now (AllocBuffer mutator visits buffer shape at the buffer definition) + // buffer IS mutated now (AllocTensor mutator visits buffer shape at the buffer definition) // shape was {1, x+1}, mutator transforms x+1 -> x, so buffer changes TVM_FFI_ICHECK(arr[0].as()->var.get() != bufptr); } @@ -360,25 +360,25 @@ TEST(IRF, StmtExprMutator) { auto* alloc_node = body.as()->seq[0].as(); TVM_FFI_ICHECK(alloc_node != nullptr); auto* alloc_call = alloc_node->value.as(); - TVM_FFI_ICHECK(alloc_call && alloc_call->op.same_as(tirx::builtin::alloc_buffer())); + TVM_FFI_ICHECK(alloc_call && alloc_call->op.same_as(tirx::builtin::alloc_tensor())); // bref still holds the old SeqStmt (not shared with new one due to copy) TVM_FFI_ICHECK(!bref.same_as(body)); } { // tests for block and block_realize - // AllocBuffer and DeclBuffer are flat (no body), placed as siblings in SeqStmt + // AllocTensor and DeclTensor are flat (no body), placed as siblings in SeqStmt Stmt eval_body = Evaluate(x + 1); - BufferVar buffer = decl_buffer({16}); + BufferVar buffer = decl_tensor({16}); tirx::Var buffer_data("buffer_data", buffer.DataPointerType()); - Stmt decl = Bind(buffer, Call(buffer.type(), tvm::tirx::builtin::decl_buffer(), + Stmt decl = Bind(buffer, Call(buffer.type(), tvm::tirx::builtin::decl_tensor(), {buffer_data, tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())})); Stmt alloc = fmakealloc(); - // body is: DeclBuffer, AllocBuffer, Evaluate + // body is: DeclTensor, AllocTensor, Evaluate Stmt body = SeqStmt({decl, alloc, eval_body}); TensorRegion buffer_region = BufferRegion(buffer, {Range::FromMinExtent(x + 1, 1)}); - s_tir::MatchBufferRegion match_buffer_region(decl_buffer({1}), buffer_region); + s_tir::MatchBufferRegion match_buffer_region(decl_tensor({1}), buffer_region); // construct block and block_realize s_tir::SBlock block = s_tir::SBlock({}, {buffer_region}, {buffer_region}, "block", body, body, {}, {match_buffer_region}); @@ -685,7 +685,7 @@ TEST(IRF, StructuralMapBufferDefinition) { PrimVar n("n", PrimType::Int(32)); auto fmakebuffer = [&]() { - return BufferVar("buf", BufferType(/*storage_scope=*/"global", + return BufferVar("buf", TensorType(/*storage_scope=*/"global", /*dtype=*/PrimType::Float(32), /*shape=*/{n}, /*strides=*/{}, @@ -695,15 +695,15 @@ TEST(IRF, StructuralMapBufferDefinition) { }; { - // Test substitution of an explicit DeclBuffer source and a dependent - // BufferType shape. Changing the type creates one fresh Var identity + // Test substitution of an explicit DeclTensor source and a dependent + // TensorType shape. Changing the type creates one fresh Var identity // that is shared by the declaration and every use. tirx::Var y = x.CopyWithSuffix("subst"); PrimVar m("m", PrimType::Int(32)); BufferVar buffer = fmakebuffer(); Stmt store = BufferStore(buffer, FloatImm(dtype, 0), {IntImm::Int32(0)}); Stmt decl = - SeqStmt({Bind(buffer, Call(buffer.type(), tvm::tirx::builtin::decl_buffer(), + SeqStmt({Bind(buffer, Call(buffer.type(), tvm::tirx::builtin::decl_tensor(), {x, tvm::Tuple(buffer->shape), DataTypeImm(buffer->dtype->dtype), StringImm(buffer.scope())})), store}); @@ -719,7 +719,7 @@ TEST(IRF, StructuralMapBufferDefinition) { auto* decl_node = seq_node->seq[0].as(); TVM_FFI_ICHECK(decl_node != nullptr); auto* decl_call = decl_node->value.as(); - TVM_FFI_ICHECK(decl_call && decl_call->op.same_as(tirx::builtin::decl_buffer())); + TVM_FFI_ICHECK(decl_call && decl_call->op.same_as(tirx::builtin::decl_tensor())); TVM_FFI_ICHECK(decl_call->args[0].same_as(y)); TVM_FFI_ICHECK(decl_node->var.as_or_throw()->shape[0].same_as(m)); TVM_FFI_ICHECK(!decl_node->var.same_as(buffer)); diff --git a/tests/cpp/sym_simplify_test.cc b/tests/cpp/sym_simplify_test.cc index 6f47e1279f51..fc30621df636 100644 --- a/tests/cpp/sym_simplify_test.cc +++ b/tests/cpp/sym_simplify_test.cc @@ -101,7 +101,7 @@ TEST(Simplify, AssumeConstraintKeepsBufferLoadStable) { using namespace tvm; sym::Analyzer analyzer; - tirx::BufferVar buffer = tirx::decl_buffer({1}, PrimType::Int(32)); + tirx::BufferVar buffer = tirx::decl_tensor({1}, PrimType::Int(32)); PrimExpr load = tirx::BufferLoad(buffer, {IntImm::Int32(0)}); PrimExpr constraint = load > 0; diff --git a/tests/cpp/tir_analysis_side_effect.cc b/tests/cpp/tir_analysis_side_effect.cc index b2d095289564..f91429a182a5 100644 --- a/tests/cpp/tir_analysis_side_effect.cc +++ b/tests/cpp/tir_analysis_side_effect.cc @@ -29,7 +29,7 @@ TEST(SimplePasses, SideEffect) { using namespace tvm::prim; using namespace tvm; - auto buf = tirx::decl_buffer({16}, PrimType::Float(32)); + auto buf = tirx::decl_tensor({16}, PrimType::Float(32)); auto i = PrimVar("i", PrimType::Int(32)); TVM_FFI_ICHECK(SideEffect(tirx::BufferLoad(buf, {i})) == CallEffectKind::kReadState); TVM_FFI_ICHECK(SideEffect(exp(prim::Cast(PrimType::Float(32), i + 1))) == CallEffectKind::kPure); diff --git a/tests/python/codegen/test_codegen_error_handling.py b/tests/python/codegen/test_codegen_error_handling.py index 8e76d7614e99..93a7d2857131 100644 --- a/tests/python/codegen/test_codegen_error_handling.py +++ b/tests/python/codegen/test_codegen_error_handling.py @@ -43,7 +43,7 @@ def test_wrong_argument_count_error(codegen_target): n0 = T.dynamic("n0") @T.prim_func - def func(A: T.Buffer((n0,), "float32"), B: T.Buffer((n0,), "float32")): + def func(A: T.Tensor((n0,), "float32"), B: T.Tensor((n0,), "float32")): for i in range(n0): B[i] = A[i] + T.float32(1) @@ -71,7 +71,7 @@ def test_type_mismatch_non_tensor(codegen_target): n0 = T.dynamic("n0") @T.prim_func - def func(A: T.Buffer((n0,), "float32"), B: T.Buffer((n0,), "float32")): + def func(A: T.Tensor((n0,), "float32"), B: T.Tensor((n0,), "float32")): for i in range(n0): B[i] = A[i] + T.float32(1) @@ -100,7 +100,7 @@ def test_shape_mismatch_shared_variable(codegen_target): n0 = T.dynamic("n0") @T.prim_func - def func(A: T.Buffer((n0,), "float32"), B: T.Buffer((n0,), "float32")): + def func(A: T.Tensor((n0,), "float32"), B: T.Tensor((n0,), "float32")): for i in range(n0): B[i] = A[i] + T.float32(1) @@ -125,7 +125,7 @@ def test_invalid_shape_fixed(codegen_target): """Passing wrong shape for a fixed buffer dimension raises ValueError.""" @T.prim_func - def func(a: T.Buffer((128,), "float32"), b: T.Buffer((128,), "float32")): + def func(a: T.Tensor((128,), "float32"), b: T.Tensor((128,), "float32")): for i in range(128): b[i] = a[i] + T.float32(1) @@ -154,7 +154,7 @@ def test_ndim_mismatch_error(codegen_target): """ndim mismatch produces ValueError with function signature.""" @T.prim_func - def func(a: T.Buffer((4, 8), "float32"), b: T.Buffer((4, 8), "float32")): + def func(a: T.Tensor((4, 8), "float32"), b: T.Tensor((4, 8), "float32")): for i, j in T.grid(4, 8): b[i, j] = a[i, j] @@ -183,7 +183,7 @@ def test_dtype_mismatch_error(codegen_target): """dtype mismatch produces TypeError with function signature.""" @T.prim_func - def func(a: T.Buffer((8,), "float32"), b: T.Buffer((8,), "float32")): + def func(a: T.Tensor((8,), "float32"), b: T.Tensor((8,), "float32")): for i in range(8): b[i] = a[i] @@ -213,7 +213,7 @@ def test_data_alignment_error(codegen_target): """Misaligned buffer data pointer raises ValueError.""" @T.prim_func - def func(a: T.Buffer((128,), "float32"), b: T.Buffer((128,), "float32")): + def func(a: T.Tensor((128,), "float32"), b: T.Tensor((128,), "float32")): for i in range(128): b[i] = a[i] + T.float32(1) @@ -245,7 +245,7 @@ def test_strides_mismatch_transposed(codegen_target): """Transposed (non-compact) strides raise ValueError.""" @T.prim_func - def func(a: T.Buffer((128, 128), "float32"), b: T.Buffer((128, 128), "float32")): + def func(a: T.Tensor((128, 128), "float32"), b: T.Tensor((128, 128), "float32")): for i, j in T.grid(128, 128): b[i, j] = a[i, j] + T.float32(1) @@ -279,7 +279,7 @@ def test_device_mismatch_error(): """Passing GPU tensor to CPU function raises ValueError.""" @T.prim_func - def func(a: T.Buffer((128,), "float32"), b: T.Buffer((128,), "float32")): + def func(a: T.Tensor((128,), "float32"), b: T.Tensor((128,), "float32")): for i in range(128): b[i] = a[i] + T.float32(1) @@ -392,7 +392,7 @@ def test_forward_reference_symbolic_shape(codegen_target): batch_size = T.dynamic("batch_size") @T.prim_func - def func(A: T.Buffer((batch_size + 1,), "int32"), B: T.Buffer((batch_size,), "int32")): + def func(A: T.Tensor((batch_size + 1,), "int32"), B: T.Tensor((batch_size,), "int32")): for i in range(batch_size): B[i] = A[i] + A[i + 1] @@ -425,7 +425,7 @@ def test_invalid_arguments_mixed_params(codegen_target): """Mixed bool + tensor function: type, dtype, and shape errors.""" @T.prim_func - def func(a0: T.bool, a1: T.Buffer([10], "float32")) -> T.int32: + def func(a0: T.bool, a1: T.Tensor([10], "float32")) -> T.int32: return 0 lib = tvm.compile(func, target=codegen_target) diff --git a/tests/python/codegen/test_gpu_codegen_allreduce.py b/tests/python/codegen/test_gpu_codegen_allreduce.py index a94d760b0fd3..bbbbf62e5727 100644 --- a/tests/python/codegen/test_gpu_codegen_allreduce.py +++ b/tests/python/codegen/test_gpu_codegen_allreduce.py @@ -35,12 +35,12 @@ def _reduce_module(d1, d2, d3, is_max=False): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1, d1, d2, d3), "float32"), B: T.Buffer((1, d1, d2), "float32")): + def main(A: T.Tensor((1, d1, d2, d3), "float32"), B: T.Tensor((1, d1, d2), "float32")): for i in T.thread_binding(1, thread="blockIdx.x"): for j in T.thread_binding(d1, thread="threadIdx.z"): for k in T.thread_binding(d2, thread="threadIdx.y"): for l in T.thread_binding(d3, thread="threadIdx.x"): - reduced = T.alloc_buffer((1,), "float32", scope="local") + reduced = T.alloc_tensor((1,), "float32", scope="local") with T.attr(reducer, "reduce_scope", 0): T.tvm_thread_allreduce( T.uint32(1), A[i, j, k, l], True, reduced[0], l diff --git a/tests/python/codegen/test_target_codegen.py b/tests/python/codegen/test_target_codegen.py index 205340963a20..6c354e84141e 100644 --- a/tests/python/codegen/test_target_codegen.py +++ b/tests/python/codegen/test_target_codegen.py @@ -27,7 +27,7 @@ def test_buffer_store_predicate_not_supported(): target = "c" @T.prim_func - def func(B: T.Buffer((8,), "float32")): + def func(B: T.Tensor((8,), "float32")): T.evaluate( T.call_intrin( "void", @@ -60,7 +60,7 @@ def test_buffer_store_predicate_not_supported_gpu(target): pytest.skip(f"{target} not enabled") @T.prim_func - def func(A: T.Buffer((2, 3), "float32"), B: T.Buffer((6,), "float32")): + def func(A: T.Tensor((2, 3), "float32"), B: T.Tensor((6,), "float32")): T.func_attr({"global_symbol": "main"}) for i_0 in T.thread_binding(3, thread="threadIdx.x"): T.evaluate( @@ -84,7 +84,7 @@ def test_buffer_load_predicate_not_supported(): target = "c" @T.prim_func - def func(A: T.Buffer((8,), "float32"), B: T.Buffer((8,), "float32")): + def func(A: T.Tensor((8,), "float32"), B: T.Tensor((8,), "float32")): for i_0 in range(4): B.vstore( [T.Ramp(0, 2, 4)], @@ -118,7 +118,7 @@ def test_buffer_load_predicate_not_supported_gpu(target): pytest.skip(f"{target} not enabled") @T.prim_func - def func(A: T.Buffer((8,), "float32"), B: T.Buffer((8,), "float32")): + def func(A: T.Tensor((8,), "float32"), B: T.Tensor((8,), "float32")): for i_0 in T.thread_binding(3, thread="threadIdx.x"): B.vstore( [T.Ramp(0, 2, 4)], @@ -152,8 +152,8 @@ def kernel(A_ptr: T.handle("float32", "global")): "tirx.noalias": True, } ) - A = T.decl_buffer((8,), "float32", data=A_ptr) - B = T.decl_buffer((4,), "float32", data=T.address_of(A[4])) + A = T.decl_tensor((8,), "float32", data=A_ptr) + B = T.decl_tensor((4,), "float32", data=T.address_of(A[4])) B[0] = T.float32(1) mod = tvm.IRModule({"kernel": kernel}) @@ -170,9 +170,9 @@ def test_codegen_loop_step(target): @T.prim_func def test_loop_step( - A: T.Buffer((1024,), "float32"), - B: T.Buffer((1024,), "float32"), - C: T.Buffer((1024,), "float32"), + A: T.Tensor((1024,), "float32"), + B: T.Tensor((1024,), "float32"), + C: T.Tensor((1024,), "float32"), ): for i in T.serial(3, 1024, step=96): C[i] = A[i] + B[i] diff --git a/tests/python/codegen/test_target_codegen_aarch64.py b/tests/python/codegen/test_target_codegen_aarch64.py index 63bb8ad28365..6c3a9dcc7070 100644 --- a/tests/python/codegen/test_target_codegen_aarch64.py +++ b/tests/python/codegen/test_target_codegen_aarch64.py @@ -45,9 +45,9 @@ def test_mul(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), dtype=dtype), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), dtype=dtype), ): T.func_attr({"tirx.noalias": True}) @@ -83,9 +83,9 @@ def test_add(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), dtype=dtype), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), dtype=dtype), ): T.func_attr({"tirx.noalias": True}) @@ -121,9 +121,9 @@ def test_sub(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), dtype=dtype), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), dtype=dtype), ): T.func_attr({"tirx.noalias": True}) @@ -159,10 +159,10 @@ def test_muladd(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), dtype=dtype), - D: T.Buffer((m,), dtype=dtype), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), dtype=dtype), + D: T.Tensor((m,), dtype=dtype), ): T.func_attr({"tirx.noalias": True}) @@ -208,9 +208,9 @@ def test_max(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), dtype=dtype), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), dtype=dtype), ): T.func_attr({"tirx.noalias": True}) @@ -250,9 +250,9 @@ def test_min(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), dtype=dtype), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), dtype=dtype), ): T.func_attr({"tirx.noalias": True}) @@ -292,9 +292,9 @@ def test_div(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), dtype=dtype), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), dtype=dtype), ): T.func_attr({"tirx.noalias": True}) @@ -329,9 +329,9 @@ def test_mod(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), dtype=dtype), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), dtype=dtype), ): T.func_attr({"tirx.noalias": True}) @@ -367,9 +367,9 @@ def test_eq(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), "bool"), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), "bool"), ): T.func_attr({"tirx.noalias": True}) @@ -408,9 +408,9 @@ def test_neq(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), "bool"), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), "bool"), ): T.func_attr({"tirx.noalias": True}) @@ -448,9 +448,9 @@ def test_or(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), dtype=dtype), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), dtype=dtype), ): T.func_attr({"tirx.noalias": True}) @@ -485,9 +485,9 @@ def test_and(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), dtype=dtype), - C: T.Buffer((m,), dtype=dtype), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), dtype=dtype), + C: T.Tensor((m,), dtype=dtype), ): T.func_attr({"tirx.noalias": True}) @@ -521,7 +521,7 @@ def test_not(dtype): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((m,), dtype=dtype), C: T.Buffer((m,), dtype=dtype)): + def main(A: T.Tensor((m,), dtype=dtype), C: T.Tensor((m,), dtype=dtype)): T.func_attr({"tirx.noalias": True}) for i in range(m): @@ -559,9 +559,9 @@ def test_memcpy(dtype): class Module: @T.prim_func def main( - A: T.Buffer((m,), dtype=dtype), - B: T.Buffer((m,), "int32"), - C: T.Buffer((m,), dtype=dtype), + A: T.Tensor((m,), dtype=dtype), + B: T.Tensor((m,), "int32"), + C: T.Tensor((m,), dtype=dtype), ): T.func_attr({"tirx.noalias": True}) @@ -601,7 +601,7 @@ def test_vscale_range_function_attribute(mattr, expect_attr): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((m,)), C: T.Buffer((m,))): + def main(A: T.Tensor((m,)), C: T.Tensor((m,))): T.func_attr({"tirx.noalias": True}) 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 6e1fbcd501da..dee8dc022126 100644 --- a/tests/python/codegen/test_target_codegen_arm.py +++ b/tests/python/codegen/test_target_codegen_arm.py @@ -33,7 +33,7 @@ def check_correct_assembly(type, elements, counts): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((elements,), type), B: T.Buffer((elements,), type)): + def main(A: T.Tensor((elements,), type), B: T.Tensor((elements,), type)): T.func_attr({"tirx.noalias": True}) for i in T.vectorized(elements): B[i] = T.popcount(A[i]) @@ -68,7 +68,7 @@ def check_correct_assembly(N): class Module: @T.prim_func def main( - A: T.Buffer((K, N), "int8"), B: T.Buffer((K, N), "int8"), C: T.Buffer((N,), "int32") + A: T.Tensor((K, N), "int8"), B: T.Tensor((K, N), "int8"), C: T.Tensor((N,), "int32") ): T.func_attr({"tirx.noalias": True}) @@ -96,7 +96,7 @@ def check_broadcast_correct_assembly(N): class Module: @T.prim_func def main( - A: T.Buffer((K, N), "int8"), B: T.Buffer((K,), "int8"), C: T.Buffer((N,), "int32") + A: T.Tensor((K, N), "int8"), B: T.Tensor((K,), "int8"), C: T.Tensor((N,), "int32") ): T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/codegen/test_target_codegen_blob.py b/tests/python/codegen/test_target_codegen_blob.py index 237d4829883e..b0ce15967f27 100644 --- a/tests/python/codegen/test_target_codegen_blob.py +++ b/tests/python/codegen/test_target_codegen_blob.py @@ -44,7 +44,7 @@ class ModA: I.module_attrs({"system_lib_prefix": "modA_"}) @T.prim_func - def my_inplace_update(x: T.Buffer((12), "float32")) -> None: + def my_inplace_update(x: T.Tensor((12), "float32")) -> None: T.func_attr({"global_symbol": "modA_my_inplace_update"}) for bx in T.thread_binding(T.int64(1), thread="blockIdx.x"): for tx in T.thread_binding(T.int64(12), thread="threadIdx.x"): @@ -55,7 +55,7 @@ class ModB: I.module_attrs({"system_lib_prefix": "modB_"}) @T.prim_func - def my_inplace_update(x: T.Buffer((12), "float32")) -> None: + def my_inplace_update(x: T.Tensor((12), "float32")) -> None: T.func_attr({"global_symbol": "modB_my_inplace_update"}) for bx in T.thread_binding(T.int64(1), thread="blockIdx.x"): for tx in T.thread_binding(T.int64(12), thread="threadIdx.x"): diff --git a/tests/python/codegen/test_target_codegen_bool.py b/tests/python/codegen/test_target_codegen_bool.py index b4b85f6f669f..0792169def99 100644 --- a/tests/python/codegen/test_target_codegen_bool.py +++ b/tests/python/codegen/test_target_codegen_bool.py @@ -35,14 +35,14 @@ def test_cmp_load_store(target): class GPUModule: @T.prim_func def main( - A: T.Buffer((32,), "float32"), - B: T.Buffer((32,), "float32"), - D: T.Buffer((32,), "float32"), + A: T.Tensor((32,), "float32"), + B: T.Tensor((32,), "float32"), + D: T.Tensor((32,), "float32"), ): T.func_attr({"tirx.noalias": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): for tx in T.thread_binding(4, thread="threadIdx.x"): - C = T.alloc_buffer((1,), "bool", scope="local") + C = T.alloc_tensor((1,), "bool", scope="local") C[0] = B[bx * 4 + tx] < A[bx * 4 + tx] D[bx * 4 + tx] = T.Cast("float32", C[0] and T.float32(1.0) < A[bx * 4 + tx]) @@ -50,12 +50,12 @@ def main( class CPUModule: @T.prim_func def main( - A: T.Buffer((32,), "float32"), - B: T.Buffer((32,), "float32"), - D: T.Buffer((32,), "float32"), + A: T.Tensor((32,), "float32"), + B: T.Tensor((32,), "float32"), + D: T.Tensor((32,), "float32"), ): T.func_attr({"tirx.noalias": True}) - C = T.alloc_buffer((32,), "bool") + C = T.alloc_tensor((32,), "bool") for i0 in range(32): C[i0] = B[i0] < A[i0] for i0 in range(32): @@ -89,7 +89,7 @@ def run_and_check(): def test_bitwise_not_c(tmp_path): @T.prim_func - def complement(values: T.Buffer((5,), "int32"), output: T.Buffer((5,), "bool")): + def complement(values: T.Tensor((5,), "int32"), output: T.Tensor((5,), "bool")): for i in range(5): output[i] = T.bitwise_not(values[i] != 0) diff --git a/tests/python/codegen/test_target_codegen_c_host.py b/tests/python/codegen/test_target_codegen_c_host.py index bf90341066c8..68448d66de31 100644 --- a/tests/python/codegen/test_target_codegen_c_host.py +++ b/tests/python/codegen/test_target_codegen_c_host.py @@ -31,9 +31,9 @@ def test_add(): class Module: @T.prim_func def test_fadd( - A: T.Buffer((1024,), "float32"), - B: T.Buffer((1024,), "float32"), - C: T.Buffer((1024,), "float32"), + A: T.Tensor((1024,), "float32"), + B: T.Tensor((1024,), "float32"), + C: T.Tensor((1024,), "float32"), ): T.func_attr({"tirx.noalias": True}) for i0 in range(1024): @@ -64,8 +64,8 @@ def test_reinterpret(): class Module: @T.prim_func def test_reinterpret( - A: T.Buffer((1024,), "int32"), - B: T.Buffer((1024,), "float32"), + A: T.Tensor((1024,), "int32"), + B: T.Tensor((1024,), "float32"), ): T.func_attr({"tirx.noalias": True}) for i0 in range(1024): @@ -95,8 +95,8 @@ def test_ceil(): class Module: @T.prim_func def test_ceil( - A: T.Buffer((1024,), "float32"), - B: T.Buffer((1024,), "float32"), + A: T.Tensor((1024,), "float32"), + B: T.Tensor((1024,), "float32"), ): T.func_attr({"tirx.noalias": True}) for i0 in range(1024): @@ -126,8 +126,8 @@ def test_floor(): class Module: @T.prim_func def test_floor( - A: T.Buffer((1024,), "float32"), - B: T.Buffer((1024,), "float32"), + A: T.Tensor((1024,), "float32"), + B: T.Tensor((1024,), "float32"), ): T.func_attr({"tirx.noalias": True}) for i0 in range(1024): @@ -157,8 +157,8 @@ def test_round(): class Module: @T.prim_func def test_round( - A: T.Buffer((1024,), "float32"), - B: T.Buffer((1024,), "float32"), + A: T.Tensor((1024,), "float32"), + B: T.Tensor((1024,), "float32"), ): T.func_attr({"tirx.noalias": True}) for i0 in range(1024): @@ -193,12 +193,12 @@ def test_subroutine_call(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer(1, dtype="float32")): + def main(A: T.Tensor(1, dtype="float32")): Module.subroutine(A.data) @T.prim_func(private=True) def subroutine(A_data: T.handle("float32")): - A = T.decl_buffer(1, dtype="float32", data=A_data) + A = T.decl_tensor(1, dtype="float32", data=A_data) A[0] = 42.0 built = tvm.tirx.build(Module, target="c") @@ -219,8 +219,8 @@ def test_workspace_allocation_cast(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((256,), "float32")): - workspace = T.alloc_buffer((256,), "float32", scope="global") + def main(A: T.Tensor((256,), "float32")): + workspace = T.alloc_tensor((256,), "float32", scope="global") for i in range(256): workspace[i] = A[i] for i in range(256): @@ -237,8 +237,8 @@ def test_local_alloc_buffer_uses_plain_c_pointer(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1,), "float32")): - B = T.alloc_buffer((1,), "float32", scope="local") + def main(A: T.Tensor((1,), "float32")): + B = T.alloc_tensor((1,), "float32", scope="local") for i in range(1): B[i] = A[i] + T.float32(1) A[i] = B[i] @@ -257,7 +257,7 @@ def main(A: T.Buffer((1,), "float32")): def test_vector_access_ptr_address_uses_ramp_base(): - buffer = tvm.tirx.decl_buffer((8,), "float32x2", name="A") + buffer = tvm.tirx.decl_tensor((8,), "float32x2", name="A") access_ptr = buffer.access_ptr(access_mask=3, offset=2, extent=4) body = tvm.tirx.Evaluate(tvm.tirx.call_extern("void", "consume", access_ptr)) func = tvm.tirx.PrimFunc([buffer], body).with_attr("global_symbol", "main") @@ -273,7 +273,7 @@ def test_if_then_else_avoids_extraneous_parentheses(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((8,), "int32"), B: T.Buffer((8,), "int32")): + def main(A: T.Tensor((8,), "int32"), B: T.Tensor((8,), "int32")): for i in range(8): B[i] = T.if_then_else(i == 0, 1, A[i]) diff --git a/tests/python/codegen/test_target_codegen_cross_llvm.py b/tests/python/codegen/test_target_codegen_cross_llvm.py index 31e2279fc740..10755b5f1c6c 100644 --- a/tests/python/codegen/test_target_codegen_cross_llvm.py +++ b/tests/python/codegen/test_target_codegen_cross_llvm.py @@ -36,9 +36,9 @@ class AddModule: @T.prim_func def main( - A: T.Buffer((1024,), "float32"), - B: T.Buffer((1024,), "float32"), - C: T.Buffer((1024,), "float32"), + A: T.Tensor((1024,), "float32"), + B: T.Tensor((1024,), "float32"), + C: T.Tensor((1024,), "float32"), ): T.func_attr({"tirx.noalias": True}) for i0_0 in T.parallel(256): diff --git a/tests/python/codegen/test_target_codegen_cuda.py b/tests/python/codegen/test_target_codegen_cuda.py index 37f4b41a25af..997798ae393f 100644 --- a/tests/python/codegen/test_target_codegen_cuda.py +++ b/tests/python/codegen/test_target_codegen_cuda.py @@ -67,7 +67,7 @@ def test_cuda_host_bundle(tmp_path): pytest.skip("CUDA-host compilation requires NVCC") @T.prim_func - def add_one(A: T.Buffer((32,), "float32"), B: T.Buffer((32,), "float32")): + def add_one(A: T.Tensor((32,), "float32"), B: T.Tensor((32,), "float32")): for tx in T.thread_binding(32, "threadIdx.x"): B[tx] = A[tx] + T.float32(1) @@ -113,9 +113,9 @@ def test_cuda_host_bundle_bf16_cluster(tmp_path): @T.prim_func def main( - A: T.Buffer((128,), "bfloat16"), - B: T.Buffer((128,), "bfloat16"), - R: T.Buffer((4,), "int32"), + A: T.Tensor((128,), "bfloat16"), + B: T.Tensor((128,), "bfloat16"), + R: T.Tensor((4,), "int32"), ): T.device_entry() for cx in T.thread_binding(2, thread="clusterCtaIdx.x"): @@ -170,7 +170,7 @@ def test_cuda_host_bundle_programmatic_dependent_launch(tmp_path): pytest.skip("programmatic dependent launch requires SM90 or newer") @T.prim_func - def add_one(A: T.Buffer((32,), "float32"), B: T.Buffer((32,), "float32")): + def add_one(A: T.Tensor((32,), "float32"), B: T.Tensor((32,), "float32")): T.func_attr({"tirx.kernel_launch_params": ["tirx.use_programtic_dependent_launch"]}) for tx in T.thread_binding(32, "threadIdx.x"): B[tx] = A[tx] + T.float32(1) @@ -251,7 +251,7 @@ def check_cuda(dtype, n, lanes): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), vec_dtype), B: T.Buffer((n,), vec_dtype)): + def main(A: T.Tensor((n,), vec_dtype), B: T.Tensor((n,), vec_dtype)): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(num_blocks, thread="blockIdx.x"): for i_1 in T.thread_binding(num_thread, thread="threadIdx.x"): @@ -313,7 +313,7 @@ def check_cuda(n, lanes): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), vec_dtype), B: T.Buffer((n,), vec_dtype)): + def main(A: T.Tensor((n,), vec_dtype), B: T.Tensor((n,), vec_dtype)): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(num_blocks, thread="blockIdx.x"): for i_1 in T.thread_binding(num_thread, thread="threadIdx.x"): @@ -358,10 +358,10 @@ def check_cuda(dtype, n, lanes): class Module: @T.prim_func def main( - A: T.Buffer((n,), vec_dtype), - B: T.Buffer((n,), vec_dtype), - C: T.Buffer((n,), "int32"), - D: T.Buffer((n,), "int32"), + A: T.Tensor((n,), vec_dtype), + B: T.Tensor((n,), vec_dtype), + C: T.Tensor((n,), "int32"), + D: T.Tensor((n,), "int32"), ): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(num_blocks, thread="blockIdx.x"): @@ -407,7 +407,7 @@ def check_cuda(dtype, n, lanes): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), vec_dtype), B: T.Buffer((n,), vec_dtype)): + def main(A: T.Tensor((n,), vec_dtype), B: T.Tensor((n,), vec_dtype)): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(num_blocks, thread="blockIdx.x"): for i_1 in T.thread_binding(num_thread, thread="threadIdx.x"): @@ -442,7 +442,7 @@ def check_cuda(n, value, lanes): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n, lanes), dtype)): + def main(A: T.Tensor((n, lanes), dtype)): T.func_attr({"tirx.noalias": True}) for i in T.thread_binding(n, thread="blockIdx.x"): for j in T.vectorized(lanes): @@ -479,7 +479,7 @@ def check_inf_nan(n, value, dtype): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), dtype), C: T.Buffer((n,), dtype)): + def main(A: T.Tensor((n,), dtype), C: T.Tensor((n,), dtype)): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(1, thread="blockIdx.x"): for i_1 in T.thread_binding(8, thread="threadIdx.x"): @@ -522,13 +522,13 @@ def sched(nthd): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n, m)), B: T.Buffer((n,))): + def main(A: T.Tensor((n, m)), B: T.Tensor((n,))): T.func_attr({"tirx.noalias": True}) for i in T.thread_binding(n, thread="blockIdx.x"): for m_0 in T.thread_binding(nthd, thread="threadIdx.x"): - partial = T.alloc_buffer((1,), "float32", scope="local") - reduced = T.alloc_buffer((1,), "float32", scope="local") + partial = T.alloc_tensor((1,), "float32", scope="local") + reduced = T.alloc_tensor((1,), "float32", scope="local") partial[0] = T.float32(0) for m_1 in range((m + nthd - 1) // nthd): if m_0 * ((m + nthd - 1) // nthd) + m_1 < m: @@ -584,14 +584,14 @@ def sched(nthdx, nthdy): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n, k0, k1)), B: T.Buffer((n,))): + def main(A: T.Tensor((n, k0, k1)), B: T.Tensor((n,))): T.func_attr({"tirx.noalias": True}) for i in T.thread_binding(n, thread="blockIdx.x"): for k0_0 in T.thread_binding(nthdx, thread="threadIdx.x"): for k1_0 in T.thread_binding(nthdy, thread="threadIdx.y"): - partial = T.alloc_buffer((1,), "float32", scope="local") - reduced = T.alloc_buffer((1,), "float32", scope="local") + partial = T.alloc_tensor((1,), "float32", scope="local") + reduced = T.alloc_tensor((1,), "float32", scope="local") partial[0] = T.float32(0) for k0_1, k1_1 in T.grid( (k0 + nthdx - 1) // nthdx, (k1 + nthdy - 1) // nthdy @@ -652,7 +652,7 @@ def test_cuda_reduction_binding(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((96, 32), "float32"), B: T.Buffer((96,), "float32")): + def main(A: T.Tensor((96, 32), "float32"), B: T.Tensor((96,), "float32")): T.func_attr({"tirx.noalias": True}) for k in range(32): for m_0 in T.thread_binding(3, thread="blockIdx.x"): @@ -675,7 +675,7 @@ def test_cuda_const_float_to_half(): @I.ir_module class Module: @T.prim_func - def main(a: T.Buffer((2, 3, 4), "float16"), C: T.Buffer((2, 3, 4), "bool")): + def main(a: T.Tensor((2, 3, 4), "float16"), C: T.Tensor((2, 3, 4), "bool")): T.func_attr({"tirx.noalias": True}) for i_j_k_fused_0 in T.thread_binding(1, thread="blockIdx.x"): for i_j_k_fused_1 in T.thread_binding(64, thread="threadIdx.x"): @@ -720,7 +720,7 @@ def test_cuda_floordiv_with_vectorization(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((256,), "float32"), B: T.Buffer((256,), "float32")): + def main(A: T.Tensor((256,), "float32"), B: T.Tensor((256,), "float32")): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(1, thread="blockIdx.x"): for i_1_0 in T.thread_binding(64, thread="threadIdx.x"): @@ -754,7 +754,7 @@ def test_cuda_floormod_with_vectorization(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((256,), "float32"), B: T.Buffer((256,), "float32")): + def main(A: T.Tensor((256,), "float32"), B: T.Tensor((256,), "float32")): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(1, thread="blockIdx.x"): for i_1_0 in T.thread_binding(64, thread="threadIdx.x"): @@ -838,7 +838,7 @@ def test_vectorized_casts(t0, t1, factor): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), t0), B: T.Buffer((n,), t1), C: T.Buffer((n,), t0)): + def main(A: T.Tensor((n,), t0), B: T.Tensor((n,), t1), C: T.Tensor((n,), t0)): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(num_thread, thread="threadIdx.x"): for i_1 in T.vectorized(factor): @@ -875,7 +875,7 @@ def sched(compute_fn, dtype, n=128): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), dtype), B: T.Buffer((n,), dtype)): + def main(A: T.Tensor((n,), dtype), B: T.Tensor((n,), dtype)): T.func_attr({"tirx.noalias": True}) for i0_0 in T.thread_binding(1, thread="blockIdx.x"): for i0_1_0 in T.thread_binding(32, thread="threadIdx.x"): @@ -1045,7 +1045,7 @@ def check_cuda(dtype, n, l, padding, lanes): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n, l), dtype), B: T.Buffer((dim0, dim1, lanes), dtype)): + def main(A: T.Tensor((n, l), dtype), B: T.Tensor((dim0, dim1, lanes), dtype)): T.func_attr({"tirx.noalias": True}) for i in T.thread_binding(dim0, thread="blockIdx.x"): for j in T.thread_binding(dim1, thread="threadIdx.x"): @@ -1087,7 +1087,7 @@ def build(N, C_N, offset): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((N,), "float16"), C: T.Buffer((C_N,), "float16")): + def main(A: T.Tensor((N,), "float16"), C: T.Tensor((C_N,), "float16")): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(C_N // 2, thread="threadIdx.x"): for i_1 in T.vectorized(2): @@ -1129,8 +1129,8 @@ def run_and_check(): @pytest.mark.skipif(not env.has_cuda(), reason="need cuda") def test_cuda_thread_sync_inside_condition(): @T.prim_func - def func2(A: T.Buffer((4, 4), "float32")) -> None: - A_shared = T.alloc_buffer((4, 4), "float32", scope="shared") + def func2(A: T.Tensor((4, 4), "float32")) -> None: + A_shared = T.alloc_tensor((4, 4), "float32", scope="shared") for bx in T.thread_binding(1, "blockIdx.x"): for tx in T.thread_binding(32, "threadIdx.x"): if T.tvm_thread_invariant(A[0, 0] > 1.0): @@ -1141,8 +1141,8 @@ def func2(A: T.Buffer((4, 4), "float32")) -> None: A[i, j] = A_shared[i, j] + 1.0 @T.prim_func - def func3(A: T.Buffer((4, 4), "float32")) -> None: - A_shared = T.alloc_buffer((4, 4), "float32", scope="shared") + def func3(A: T.Tensor((4, 4), "float32")) -> None: + A_shared = T.alloc_tensor((4, 4), "float32", scope="shared") for bx in T.thread_binding(1, "blockIdx.x"): for tx in T.thread_binding(32, "threadIdx.x"): while T.tvm_thread_invariant(A[0, 0] > 1.0): @@ -1162,7 +1162,7 @@ def func3(A: T.Buffer((4, 4), "float32")) -> None: @pytest.mark.skipif(not env.has_cuda(), reason="need cuda") def test_invalid_reinterpret(): @T.prim_func - def func(A: T.Buffer((4,), "uint32"), B: T.Buffer((4,), "uint8")) -> None: + def func(A: T.Tensor((4,), "uint32"), B: T.Tensor((4,), "uint8")) -> None: for tx in T.thread_binding(4, "threadIdx.x"): B[tx] = T.call_intrin("uint8", "tirx.reinterpret", A[tx]) @@ -1175,7 +1175,7 @@ def func(A: T.Buffer((4,), "uint32"), B: T.Buffer((4,), "uint8")) -> None: def test_cuda_tensormap(): # fmt: off @T.prim_func - def main(A: T.Buffer((16, 16), dtype='float32', align=16)): + def main(A: T.Tensor((16, 16), dtype='float32', align=16)): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) T.call_packed("runtime.cuTensorMapInit", A_map, "float32", 2, A.data, @@ -1211,9 +1211,9 @@ def add(a: T.float32, b: T.float32) -> T.float32: @T.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ): for bx in T.thread_binding(1024, "blockIdx.x"): for tx in T.thread_binding(1024, "threadIdx.x"): @@ -1233,7 +1233,7 @@ def test_cuda_float_const_hex_format(): class Module: @T.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), ): for bx in T.thread_binding(1024, "blockIdx.x"): for tx in T.thread_binding(1024, "threadIdx.x"): @@ -1255,9 +1255,9 @@ def add(a: T.int32, b: T.int32) -> T.int32: @T.prim_func def main( - A: T.Buffer((128, 128), "int32"), - B: T.Buffer((128, 128), "int32"), - C: T.Buffer((128, 128), "int32"), + A: T.Tensor((128, 128), "int32"), + B: T.Tensor((128, 128), "int32"), + C: T.Tensor((128, 128), "int32"), ): length: T.let[T.int32] = Module.add(64, 64) # Call from host for bx in T.thread_binding(length, "blockIdx.x"): @@ -1293,7 +1293,7 @@ def test_thread_return(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")): + def main(A: T.Tensor((16, 16), "float32"), B: T.Tensor((16, 16), "float32")): for bx in T.thread_binding(32, "blockIdx.x"): for tx in T.thread_binding(32, "threadIdx.x"): if bx >= 16 or tx >= 16: @@ -1311,9 +1311,9 @@ def main(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")): def test_cuda_loop_step(): @T.prim_func def cuda_loop_step( - A: T.Buffer((1024,), "float32"), - B: T.Buffer((1024,), "float32"), - C: T.Buffer((1024,), "float32"), + A: T.Tensor((1024,), "float32"), + B: T.Tensor((1024,), "float32"), + C: T.Tensor((1024,), "float32"), ): # Each thread computes a strided subset of the i loop: start = tx*3, step = 96 (3 * 32 threads) for bx in T.thread_binding(1, "blockIdx.x"): @@ -1350,7 +1350,7 @@ def test_export_load_with_fallback(monkeypatch, tmp_path): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), "float32"), B: T.Buffer((n,), "float32")): + def main(A: T.Tensor((n,), "float32"), B: T.Tensor((n,), "float32")): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(n // 32, thread="blockIdx.x"): for i_1 in T.thread_binding(32, thread="threadIdx.x"): diff --git a/tests/python/codegen/test_target_codegen_cuda_fastmath.py b/tests/python/codegen/test_target_codegen_cuda_fastmath.py index c90fc5e92fd8..fbddc809ef93 100644 --- a/tests/python/codegen/test_target_codegen_cuda_fastmath.py +++ b/tests/python/codegen/test_target_codegen_cuda_fastmath.py @@ -46,8 +46,8 @@ def make_prim_func( @T.prim_func def kernel( - A: T.Buffer((VECTOR_N_INPUTS,), dtype), - B: T.Buffer((VECTOR_N_INPUTS,), dtype), + A: T.Tensor((VECTOR_N_INPUTS,), dtype), + B: T.Tensor((VECTOR_N_INPUTS,), dtype), ): T.func_attr({"global_symbol": name + "_kernel", "tirx.noalias": True}) for i in T.thread_binding(VECTOR_N_INPUTS, thread="threadIdx.x"): @@ -58,9 +58,9 @@ def kernel( @T.prim_func def kernel( - A: T.Buffer((VECTOR_N_INPUTS,), dtype), - E: T.Buffer((VECTOR_N_INPUTS,), dtype), - B: T.Buffer((VECTOR_N_INPUTS,), dtype), + A: T.Tensor((VECTOR_N_INPUTS,), dtype), + E: T.Tensor((VECTOR_N_INPUTS,), dtype), + B: T.Tensor((VECTOR_N_INPUTS,), dtype), ): T.func_attr({"global_symbol": name + "_kernel", "tirx.noalias": True}) for i in T.thread_binding(VECTOR_N_INPUTS, thread="threadIdx.x"): diff --git a/tests/python/codegen/test_target_codegen_cuda_fp4.py b/tests/python/codegen/test_target_codegen_cuda_fp4.py index eb998f244ccc..ecc5fdb3615b 100644 --- a/tests/python/codegen/test_target_codegen_cuda_fp4.py +++ b/tests/python/codegen/test_target_codegen_cuda_fp4.py @@ -45,9 +45,9 @@ def test_e2m1_vector_conversions(promoted_dtype): class Module: @T.prim_func def main( - A: T.Buffer((vector_length,), native_dtype), - B: T.Buffer((vector_length,), native_dtype), - C: T.Buffer((vector_length,), native_dtype), + A: T.Tensor((vector_length,), native_dtype), + B: T.Tensor((vector_length,), native_dtype), + C: T.Tensor((vector_length,), native_dtype), ): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(vector_length // 32, thread="blockIdx.x"): @@ -114,8 +114,8 @@ def _shuffle_reinterpret_module(n, num_blocks, vector_length, num_elem_per_stora class Module: @T.prim_func def main( - A: T.Buffer((n // num_elem_per_storage,), "uint32"), - B: T.Buffer((n,), "float16"), + A: T.Tensor((n // num_elem_per_storage,), "uint32"), + B: T.Tensor((n,), "float16"), ): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(num_blocks, thread="blockIdx.x"): @@ -162,8 +162,8 @@ def _scalar_reinterpret_module(n, num_blocks, vector_length, num_elem_per_storag class Module: @T.prim_func def main( - A: T.Buffer((n // num_elem_per_storage,), "uint32"), - B: T.Buffer((n,), "float16"), + A: T.Tensor((n // num_elem_per_storage,), "uint32"), + B: T.Tensor((n,), "float16"), ): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(num_blocks, thread="blockIdx.x"): @@ -232,9 +232,9 @@ def test_e2m1_scalar_buffer_offset(): n = 128 @T.prim_func - def func(A_raw: T.Buffer((n // 2,), "uint8"), B: T.Buffer((n,), "float16")): + def func(A_raw: T.Tensor((n // 2,), "uint8"), B: T.Tensor((n,), "float16")): T.func_attr({"tir.noalias": True}) - A = T.decl_buffer((n,), "float4_e2m1fn", data=A_raw.data) + A = T.decl_tensor((n,), "float4_e2m1fn", data=A_raw.data) for bx in T.thread_binding(n // 32, thread="blockIdx.x"): for tx in T.thread_binding(32, thread="threadIdx.x"): B[bx * 32 + tx] = T.Cast("float16", A[bx * 32 + tx]) diff --git a/tests/python/codegen/test_target_codegen_cuda_fp8.py b/tests/python/codegen/test_target_codegen_cuda_fp8.py index bc001e38dba2..ec330668f8cd 100644 --- a/tests/python/codegen/test_target_codegen_cuda_fp8.py +++ b/tests/python/codegen/test_target_codegen_cuda_fp8.py @@ -50,9 +50,9 @@ def _create_mod(dtype): class Module: @T.prim_func def main( - A: T.Buffer((64,), dtype), - B: T.Buffer((64,), dtype), - C: T.Buffer((64,), dtype), + A: T.Tensor((64,), dtype), + B: T.Tensor((64,), dtype), + C: T.Tensor((64,), dtype), ): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(2, thread="blockIdx.x"): @@ -100,9 +100,9 @@ def _create_mod(native_dtype, packed_dtype, length): class Module: @T.prim_func def main( - A: T.Buffer((length,), native_dtype), - R: T.Buffer((length,), packed_dtype), - B: T.Buffer((length,), native_dtype), + A: T.Tensor((length,), native_dtype), + R: T.Tensor((length,), packed_dtype), + B: T.Tensor((length,), native_dtype), ): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(2, thread="blockIdx.x"): @@ -157,9 +157,9 @@ def _create_mod(native_dtype, promoted_dtype): class Module: @T.prim_func def main( - A: T.Buffer((64,), native_dtype), - B: T.Buffer((64,), native_dtype), - C: T.Buffer((64,), native_dtype), + A: T.Tensor((64,), native_dtype), + B: T.Tensor((64,), native_dtype), + C: T.Tensor((64,), native_dtype), ): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(2, thread="blockIdx.x"): @@ -215,7 +215,7 @@ def _create_mod(bcast_length, dtype): @I.ir_module class Module: @T.prim_func - def main(a: T.Buffer((), dtype), vec: T.Buffer((bcast_length,), dtype)): + def main(a: T.Tensor((), dtype), vec: T.Tensor((bcast_length,), dtype)): for i_0 in T.thread_binding(1, thread="blockIdx.x"): for i_1 in T.thread_binding(1, thread="threadIdx.x"): vec[0:bcast_length] = T.broadcast(a[()], bcast_length) @@ -250,7 +250,7 @@ def test_half_misaligned_vector_load(vector_length): @T.prim_func def vector_load( - A: T.Buffer((length,), dtype), B: T.Buffer((length // vector_length,), vec_dtype) + A: T.Tensor((length,), dtype), B: T.Tensor((length // vector_length,), vec_dtype) ): for b in T.thread_binding(1, thread="blockIdx.x"): for i in T.thread_binding(length // vector_length, thread="threadIdx.x"): @@ -289,9 +289,9 @@ def test_half4_vector_add(): class Module: @T.prim_func def main( - A: T.Buffer((64,), "float16x4"), - B: T.Buffer((64,), "float16x4"), - C: T.Buffer((64,), "float16x4"), + A: T.Tensor((64,), "float16x4"), + B: T.Tensor((64,), "float16x4"), + C: T.Tensor((64,), "float16x4"), ): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(2, thread="blockIdx.x"): @@ -341,13 +341,13 @@ def compile_quant_and_dequant_by_scale( @T.prim_func def quantize( - A: T.Buffer(weight_shape, model_dtype), - packed: T.Buffer(quant_weight_shape, storage_dtype), - scale: T.Buffer(scales_shape, model_dtype), + A: T.Tensor(weight_shape, model_dtype), + packed: T.Tensor(quant_weight_shape, storage_dtype), + scale: T.Tensor(scales_shape, model_dtype), ): for row in T.thread_binding(rows, thread="blockIdx.x"): for group in T.thread_binding(groups, thread="threadIdx.x"): - maximum = T.alloc_buffer((1,), model_dtype, scope="local") + maximum = T.alloc_tensor((1,), model_dtype, scope="local") maximum[0] = T.Cast(model_dtype, 0) for k in range(group_size): if group * group_size + k < columns: @@ -368,9 +368,9 @@ def quantize( @T.prim_func def dequantize( - packed: T.Buffer(quant_weight_shape, storage_dtype), - scale: T.Buffer(scales_shape, model_dtype), - output: T.Buffer(weight_shape, model_dtype), + packed: T.Tensor(quant_weight_shape, storage_dtype), + scale: T.Tensor(scales_shape, model_dtype), + output: T.Tensor(weight_shape, model_dtype), ): for row in T.thread_binding(rows, thread="blockIdx.x"): for k in T.thread_binding(packed_columns, thread="threadIdx.x"): @@ -497,8 +497,8 @@ def test_main(self, weight_shape, model_dtype, target_str, compiled_functions): @pytest.mark.parametrize("dtype", ["float8_e5m2", "float8_e4m3fn", "float8_e8m0fnu"]) def test_const(dtype): @T.prim_func - def func(A: T.Buffer((4,), dtype)) -> None: - A_local = T.alloc_buffer((4,), dtype=dtype, scope="local") + def func(A: T.Tensor((4,), dtype)) -> None: + A_local = T.alloc_tensor((4,), dtype=dtype, scope="local") for tx in T.thread_binding(0, 4, "threadIdx.x"): for i in T.vectorized(4): A_local[i] = T.float32(1.0).astype(dtype) @@ -515,14 +515,14 @@ def func(A: T.Buffer((4,), dtype)) -> None: def test_copy(dtype, vec_len): @T.prim_func def func( - A: T.Buffer( + A: T.Tensor( ( 4, vec_len, ), dtype, ), - B: T.Buffer( + B: T.Tensor( ( 4, vec_len, @@ -553,18 +553,18 @@ def test_moe_gemv_shfl_down_illegal_instr(): @T.prim_func def moe_dequantize_gemv( - x: T.Buffer((1, reduce_size), "float16"), - indptr: T.Buffer((1, 2), "int32"), - w: T.Buffer((num_experts, spatial_size, reduce_size), "float8_e4m3fn"), - scale: T.Buffer((1,), "float32"), - output: T.Buffer((2, spatial_size), "float16"), + x: T.Tensor((1, reduce_size), "float16"), + indptr: T.Tensor((1, 2), "int32"), + w: T.Tensor((num_experts, spatial_size, reduce_size), "float8_e4m3fn"), + scale: T.Tensor((1,), "float32"), + output: T.Tensor((2, spatial_size), "float16"), ): for expert in T.thread_binding(2, thread="blockIdx.y"): for block in T.thread_binding(spatial_size // 4, thread="blockIdx.x"): for spatial in T.thread_binding(4, thread="threadIdx.y"): for reduction in T.thread_binding(64, thread="threadIdx.x"): - partial = T.alloc_buffer((1,), "float16", scope="local") - reduced = T.alloc_buffer((1,), "float16", scope="local") + partial = T.alloc_tensor((1,), "float16", scope="local") + reduced = T.alloc_tensor((1,), "float16", scope="local") partial[0] = T.float16(0) for k in range(reduce_size // 64): partial[0] = partial[0] + x[0, k * 64 + reduction] * ( @@ -617,9 +617,9 @@ def _create_mod(vec_length, dtype): class Module: @T.prim_func def main( - A: T.Buffer((128,), "float8_e4m3fn"), - B: T.Buffer((128,), dtype), - C: T.Buffer((128,), dtype), + A: T.Tensor((128,), "float8_e4m3fn"), + B: T.Tensor((128,), dtype), + C: T.Tensor((128,), dtype), ) -> None: for i_0 in T.thread_binding(num_threads, thread="threadIdx.x"): for i_1 in T.vectorized(vec_length): diff --git a/tests/python/codegen/test_target_codegen_device.py b/tests/python/codegen/test_target_codegen_device.py index d78a22e5c2d9..01e0b0ae9c98 100644 --- a/tests/python/codegen/test_target_codegen_device.py +++ b/tests/python/codegen/test_target_codegen_device.py @@ -33,7 +33,7 @@ def test_large_uint_imm(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((12,), "uint64")): + def main(A: T.Tensor((12,), "uint64")): T.func_attr({"tirx.noalias": True}) for i0_0 in T.thread_binding(6, thread="blockIdx.x"): for i0_1 in T.thread_binding(2, thread="threadIdx.x"): @@ -65,10 +65,10 @@ def test_add_pipeline(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,)), B: T.Buffer((), "float32"), D: T.Buffer((n,))): + def main(A: T.Tensor((n,)), B: T.Tensor((), "float32"), D: T.Tensor((n,))): T.func_attr({"tirx.noalias": True}) - C = T.alloc_buffer((n,)) + C = T.alloc_tensor((n,)) for i0_0 in T.thread_binding((n + 255) // 256, thread="blockIdx.x"): for i0_1 in T.thread_binding(256, thread="threadIdx.x"): if i0_0 * 256 + i0_1 < n: diff --git a/tests/python/codegen/test_target_codegen_extern.py b/tests/python/codegen/test_target_codegen_extern.py index 433ae2ea64fc..98b38b813a3b 100644 --- a/tests/python/codegen/test_target_codegen_extern.py +++ b/tests/python/codegen/test_target_codegen_extern.py @@ -34,7 +34,7 @@ def test_add_pipeline(): @I.ir_module class ModuleCPU: @T.prim_func - def main(A: T.Buffer((64,), "float32"), C: T.Buffer((64,), "float32")): + def main(A: T.Tensor((64,), "float32"), C: T.Tensor((64,), "float32")): for i in T.serial((64 + 1) // 2): C[T.Ramp(i * 2, 1, 2)] = A[T.Ramp(i * 2, 1, 2)] + T.Broadcast(T.float32(1), 2) @@ -42,7 +42,7 @@ def main(A: T.Buffer((64,), "float32"), C: T.Buffer((64,), "float32")): @I.ir_module class ModuleGPU: @T.prim_func - def main(A: T.Buffer((64,), "float32"), C: T.Buffer((64,), "float32")): + def main(A: T.Tensor((64,), "float32"), C: T.Tensor((64,), "float32")): bx = T.launch_thread("blockIdx.x", (64 + 4 - 1) // 4) tx = T.launch_thread("threadIdx.x", 4) idx = bx * 4 + tx @@ -81,7 +81,7 @@ def test_pack_buffer_simple(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1024,), "float32"), C: T.Buffer((1024,), "float32")): + def main(A: T.Tensor((1024,), "float32"), C: T.Tensor((1024,), "float32")): T.evaluate(T.call_packed("my_extern_array_func1", A, C)) @tvm.register_global_func diff --git a/tests/python/codegen/test_target_codegen_gpu_common.py b/tests/python/codegen/test_target_codegen_gpu_common.py index 567555c3725f..8d234c4d1cd6 100644 --- a/tests/python/codegen/test_target_codegen_gpu_common.py +++ b/tests/python/codegen/test_target_codegen_gpu_common.py @@ -52,8 +52,8 @@ def test_int_intrin(target, dtype): class Module: @T.prim_func def main( - A: T.Buffer((n,), dtype), - B: T.Buffer((n,), dtype), + A: T.Tensor((n,), dtype), + B: T.Tensor((n,), dtype), ): T.func_attr({"tirx.noalias": True}) for i0 in T.thread_binding(n, thread="threadIdx.x"): diff --git a/tests/python/codegen/test_target_codegen_hexagon.py b/tests/python/codegen/test_target_codegen_hexagon.py index f7f0c19c015e..130fe9c12d6e 100644 --- a/tests/python/codegen/test_target_codegen_hexagon.py +++ b/tests/python/codegen/test_target_codegen_hexagon.py @@ -45,9 +45,9 @@ def test_basic(): class Module: @T.prim_func def main( - C: T.Buffer((128,), "uint8"), - A: T.Buffer((128,), "uint8"), - A_1: T.Buffer((128,), "uint8"), + C: T.Tensor((128,), "uint8"), + A: T.Tensor((128,), "uint8"), + A_1: T.Tensor((128,), "uint8"), ): T.func_attr({"tirx.noalias": True}) for i in range(128): @@ -66,7 +66,7 @@ def test_llvm_target_features(): @I.ir_module class Module: @T.prim_func - def add_one(C: T.Buffer((128,), "int32"), A: T.Buffer((128,), "uint8")): + def add_one(C: T.Tensor((128,), "int32"), A: T.Tensor((128,), "uint8")): T.func_attr({"tirx.noalias": True}) for i in range(128): C[i] = T.Cast("int32", A[i]) + 1 @@ -95,7 +95,7 @@ def test_llvm_options(): @I.ir_module class Module: @T.prim_func - def main(compute: T.Buffer((10,), "int32")): + def main(compute: T.Tensor((10,), "int32")): T.func_attr({"tirx.noalias": True}) for _ in range(10): compute[_] = 0 diff --git a/tests/python/codegen/test_target_codegen_llvm.py b/tests/python/codegen/test_target_codegen_llvm.py index 6a5c0dad4a47..738edcf9e220 100644 --- a/tests/python/codegen/test_target_codegen_llvm.py +++ b/tests/python/codegen/test_target_codegen_llvm.py @@ -37,12 +37,12 @@ def test_duplicate_primfunc_global_symbol_diagnostic(): @I.ir_module class Module: @T.prim_func - def first_unique_key(A: T.Buffer((1,), "float32")): + def first_unique_key(A: T.Tensor((1,), "float32")): T.func_attr({"global_symbol": "dup_symbol", "tirx.noalias": True}) A[0] = T.float32(1) @T.prim_func - def second_unique_key(A: T.Buffer((1,), "float32")): + def second_unique_key(A: T.Tensor((1,), "float32")): T.func_attr({"global_symbol": "dup_symbol", "tirx.noalias": True}) A[0] = T.float32(2) @@ -59,12 +59,12 @@ def test_unique_primfunc_global_symbols_compile(): @I.ir_module class Module: @T.prim_func - def first_unique_key(A: T.Buffer((1,), "float32")): + def first_unique_key(A: T.Tensor((1,), "float32")): T.func_attr({"global_symbol": "dup_symbol_a", "tirx.noalias": True}) A[0] = T.float32(1) @T.prim_func - def second_unique_key(A: T.Buffer((1,), "float32")): + def second_unique_key(A: T.Tensor((1,), "float32")): T.func_attr({"global_symbol": "dup_symbol_b", "tirx.noalias": True}) A[0] = T.float32(2) @@ -77,7 +77,7 @@ def test_llvm_intrin(): class Module: @T.prim_func def main(A: T.handle("float32")): - A_buf = T.decl_buffer((4,), "float32", data=A) + A_buf = T.decl_tensor((4,), "float32", data=A) T.evaluate(T.Call("tirx.prefetch", [T.address_of(A_buf[0]), 0, 3, 1], ty="void")) fcode = tvm.compile(Module) @@ -116,7 +116,7 @@ def test_llvm_overloaded_intrin(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1, 1), "int32"), C: T.Buffer((1, 1), "int32")): + def main(A: T.Tensor((1, 1), "int32"), C: T.Tensor((1, 1), "int32")): C[0, 0] = T.call_llvm_pure_intrin("int32", "llvm.ctlz", A[0, 0], int1_zero) f = tvm.compile(Module, target="llvm") @@ -128,7 +128,7 @@ def test_llvm_lookup_intrin(): class Module: @T.prim_func def main(A: T.handle("uint8x8")): - A_buf = T.decl_buffer((1,), "uint8x8", data=A) + A_buf = T.decl_tensor((1,), "uint8x8", data=A) T.evaluate(T.call_llvm_pure_intrin("uint8x8", "llvm.ctpop.v8i8", A_buf[0])) fcode = tvm.compile(Module, None) @@ -142,7 +142,7 @@ def test_llvm_large_uintimm(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((), "uint64")): + def main(A: T.Tensor((), "uint64")): T.func_attr({"tirx.noalias": True}) A[()] = large_val + T.uint64(3) @@ -158,9 +158,9 @@ def test_llvm_multi_parallel(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((128,), "float32"), C: T.Buffer((128,), "float32")): + def main(A: T.Tensor((128,), "float32"), C: T.Tensor((128,), "float32")): T.func_attr({"tirx.noalias": True}) - B = T.alloc_buffer((128,)) + B = T.alloc_tensor((128,)) for i0_0_0 in T.parallel(1): for ax0 in range(128): B[ax0] = A[ax0] + T.float32(1.0) @@ -185,7 +185,7 @@ def check_llvm(nn, base): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((nn + base,), "float32"), C: T.Buffer((nn,), "float32")): + def main(A: T.Tensor((nn + base,), "float32"), C: T.Tensor((nn,), "float32")): T.func_attr({"tirx.noalias": True}) for i_0 in T.parallel((nn + 3) // 4): for i_1 in T.vectorized(4): @@ -212,7 +212,7 @@ def test_llvm_vadd_pipeline(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,)), B: T.Buffer((n,)), C: T.Buffer((n,))): + def main(A: T.Tensor((n,)), B: T.Tensor((n,)), C: T.Tensor((n,))): T.func_attr({"tirx.noalias": True}) for i_0 in range((n + 3) // 4): @@ -237,8 +237,8 @@ def check_llvm(nn, base, stride): class Module: @T.prim_func def main( - A: T.Buffer((nn + base, stride), "float32"), - C: T.Buffer((nn, stride), "float32"), + A: T.Tensor((nn + base, stride), "float32"), + C: T.Tensor((nn, stride), "float32"), ): T.func_attr({"tirx.noalias": True}) for i_0 in T.parallel((nn + 3) // 4): @@ -266,9 +266,9 @@ def test_llvm_temp_space(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1024,), "float32"), C: T.Buffer((1024,), "float32")): + def main(A: T.Tensor((1024,), "float32"), C: T.Tensor((1024,), "float32")): T.func_attr({"tirx.noalias": True}) - B = T.alloc_buffer((1024,)) + B = T.alloc_tensor((1024,)) for i in range(1024): B[i] = A[i] + T.float32(1.0) for i in range(1024): @@ -291,14 +291,14 @@ def test_multiple_func(): @I.ir_module class Module: @T.prim_func - def fadd1(A: T.Buffer((fadd1_n,)), B: T.Buffer((fadd1_n,)), C: T.Buffer((fadd1_n,))): + def fadd1(A: T.Tensor((fadd1_n,)), B: T.Tensor((fadd1_n,)), C: T.Tensor((fadd1_n,))): T.func_attr({"tirx.noalias": True}) for i in range(fadd1_n): C[i] = A[i] + B[i] @T.prim_func - def fadd2(A: T.Buffer((fadd2_n,)), B: T.Buffer((fadd2_n,)), C: T.Buffer((fadd2_n,))): + def fadd2(A: T.Tensor((fadd2_n,)), B: T.Tensor((fadd2_n,)), C: T.Tensor((fadd2_n,))): T.func_attr({"tirx.noalias": True}) for i in range(fadd2_n): @@ -322,7 +322,7 @@ def test_llvm_condition(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((64,), "float32"), C: T.Buffer((64,), "float32")): + def main(A: T.Tensor((64,), "float32"), C: T.Tensor((64,), "float32")): T.func_attr({"tirx.noalias": True}) for i in range(64): C[i] = T.if_then_else(8 <= i, A[i], T.float32(0.0)) @@ -344,7 +344,7 @@ def test_llvm_bool(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((64,), "int32"), C: T.Buffer((64,), "float32")): + def main(A: T.Tensor((64,), "int32"), C: T.Tensor((64,), "float32")): T.func_attr({"tirx.noalias": True}) for i in range(64): C[i] = T.Cast("float32", A[i] == 1) @@ -364,7 +364,7 @@ def test_llvm_cast_float_to_bool(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((4,), "float32"), C: T.Buffer((4,), "bool")): + def main(A: T.Tensor((4,), "float32"), C: T.Tensor((4,), "bool")): T.func_attr({"tirx.noalias": True}) for i in range(4): C[i] = T.Cast("bool", A[i]) @@ -385,12 +385,12 @@ def test_rank_zero(): class Module: @T.prim_func def main( - A: T.Buffer((64,), "float32"), - scale: T.Buffer((), "float32"), - compute: T.Buffer((), "float32"), + A: T.Tensor((64,), "float32"), + scale: T.Tensor((), "float32"), + compute: T.Tensor((), "float32"), ): T.func_attr({"tirx.noalias": True}) - C = T.alloc_buffer(()) + C = T.alloc_tensor(()) C[()] = T.float32(0.0) for k in range(64): C[()] = C[()] + A[k] * scale[()] @@ -413,12 +413,12 @@ def test_rank_zero_bound_checkers(): class Module: @T.prim_func def main( - A: T.Buffer((64,), "float32"), - scale: T.Buffer((), "float32"), - compute: T.Buffer((), "float32"), + A: T.Tensor((64,), "float32"), + scale: T.Tensor((), "float32"), + compute: T.Tensor((), "float32"), ): T.func_attr({"tirx.noalias": True}) - C = T.alloc_buffer(()) + C = T.alloc_tensor(()) C[()] = T.float32(0.0) for k in range(64): C[()] = C[()] + A[k] * scale[()] @@ -441,7 +441,7 @@ def test_alignment(): @I.ir_module class Module: @T.prim_func - def test_alignment(A: T.Buffer((1024,), "float32"), B: T.Buffer((1024,), "float32")): + def test_alignment(A: T.Tensor((1024,), "float32"), B: T.Tensor((1024,), "float32")): T.func_attr({"tirx.noalias": True}) for i_0 in range(128): for i_1 in T.vectorized(8): @@ -577,10 +577,10 @@ def check(start, end, dstart, dend, dtype, floor_div=False): class Module: @T.prim_func def main( - A: T.Buffer((a_size,), dtype), - B: T.Buffer((b_size,), dtype), - D: T.Buffer((a_size, b_size), dtype), - M: T.Buffer((a_size, b_size), dtype), + A: T.Tensor((a_size,), dtype), + B: T.Tensor((b_size,), dtype), + D: T.Tensor((a_size, b_size), dtype), + M: T.Tensor((a_size, b_size), dtype), ): T.func_attr({"tirx.noalias": True}) for i, j in T.grid(a_size, b_size): @@ -648,7 +648,7 @@ def test_llvm_fp_math(): @I.ir_module class RecipModule: @T.prim_func - def main(A: T.Buffer((n,)), B: T.Buffer((n,))): + def main(A: T.Tensor((n,)), B: T.Tensor((n,))): T.func_attr({"tirx.noalias": True}) for i in range(n): @@ -667,7 +667,7 @@ def main(A: T.Buffer((n,)), B: T.Buffer((n,))): @I.ir_module class SigmoidModule: @T.prim_func - def main(A: T.Buffer((n,)), B: T.Buffer((n,))): + def main(A: T.Tensor((n,)), B: T.Tensor((n,))): T.func_attr({"tirx.noalias": True}) for i in range(n): @@ -688,9 +688,9 @@ def test_dwarf_debug_information(): class Module: @T.prim_func def main( - A: T.Buffer((1024,), "float32"), - B: T.Buffer((1024,), "float32"), - C: T.Buffer((1024,), "float32"), + A: T.Tensor((1024,), "float32"), + B: T.Tensor((1024,), "float32"), + C: T.Tensor((1024,), "float32"), ): T.func_attr({"tirx.noalias": True}) for i0_0 in T.parallel(256): @@ -775,9 +775,9 @@ def dotest(do_vectorize): class Module: @T.prim_func def main( - A: T.Buffer((32,), "bfloat16"), - B: T.Buffer((32,), "bfloat16"), - D: T.Buffer((32,), "bfloat16"), + A: T.Tensor((32,), "bfloat16"), + B: T.Tensor((32,), "bfloat16"), + D: T.Tensor((32,), "bfloat16"), ): T.func_attr({"tirx.noalias": True}) for x in loop_kind(32): @@ -806,9 +806,9 @@ def test_llvm_crt_static_lib(): class Module: @T.prim_func def main( - A: T.Buffer((32,), "bfloat16"), - B: T.Buffer((32,), "bfloat16"), - C: T.Buffer((32,), "bfloat16"), + A: T.Tensor((32,), "bfloat16"), + B: T.Tensor((32,), "bfloat16"), + C: T.Tensor((32,), "bfloat16"), ): T.func_attr({"tirx.noalias": True}) for x in range(32): @@ -853,8 +853,8 @@ def Kirby(v: T.float32) -> T.float32: @pytest.mark.parametrize("extent", [2**32 + 1, 2**32 + 4]) def test_llvm_large_stack_allocation_uses_64bit_extent(extent): @T.prim_func - def main(A: T.Buffer((1,), "float32")): - B = T.alloc_buffer( + def main(A: T.Tensor((1,), "float32")): + B = T.alloc_tensor( (extent,), "float32", scope="global", @@ -895,7 +895,7 @@ def check_llvm(use_file): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((10,), "float32"), B: T.Buffer((10,), "float32")): + def main(A: T.Tensor((10,), "float32"), B: T.Tensor((10,), "float32")): T.func_attr({"tirx.noalias": True}) for i in T.serial(10, annotations={"pragma_import_llvm": import_val}): B[i] = T.call_pure_extern("float32", "my_add", A[i], T.float32(1.0)) @@ -916,7 +916,7 @@ def test_llvm_scalar_concat(): @I.ir_module class Module: @T.prim_func - def main(x: T.int32, y: T.int32, buffer: T.Buffer((1,), "int32x2")): + def main(x: T.int32, y: T.int32, buffer: T.Tensor((1,), "int32x2")): buffer[0] = T.Shuffle([x, y], [0, 1]) # This will crash in LLVM codegen if CodeGenLLVM::CreateVecConcat doesn't convert @@ -930,7 +930,7 @@ def test_raise_exception_during_codegen(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4, 4), "float32")) -> None: + def main(A: T.Tensor((4, 4), "float32"), B: T.Tensor((4, 4), "float32")) -> None: T.func_attr({"tirx.noalias": True}) for i in T.parallel(4): for j in T.parallel(4): @@ -952,9 +952,9 @@ def test_llvm_target_attributes(): class Module: @T.prim_func def test_func( - A: T.Buffer((tindex,)), # noqa: F821 - B: T.Buffer((tindex,)), # noqa: F821 - C: T.Buffer((tindex,)), # noqa: F821 + A: T.Tensor((tindex,)), # noqa: F821 + B: T.Tensor((tindex,)), # noqa: F821 + C: T.Tensor((tindex,)), # noqa: F821 tindex: T.int32, ): T.func_attr({"tirx.noalias": True}) @@ -1016,13 +1016,13 @@ def test_llvm_assume(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((4, 4), "int32"), B: T.Buffer((14,), "int32")): + def main(A: T.Tensor((4, 4), "int32"), B: T.Tensor((14,), "int32")): T.func_attr({"tirx.noalias": True}) - A_1 = T.decl_buffer((16,), "int32", data=A.data) + A_1 = T.decl_tensor((16,), "int32", data=A.data) for axis0, axis1 in T.grid(4, 4): T.assume(axis0 < 3 or axis1 < 2 or A_1[axis0 * 4 + axis1] == 0) for i in range(14): - B_1 = T.decl_buffer((14,), "int32", data=B.data) + B_1 = T.decl_tensor((14,), "int32", data=B.data) B_1[i] = A_1[i] * 2 m = tvm.compile(Module, target="llvm") @@ -1042,8 +1042,8 @@ class Module: @T.prim_func def main(a: T.handle("float64"), b: T.handle("float64"), n: T.int64): T.func_attr({"calling_conv": 2}) - A = T.decl_buffer(16, "float64", data=a) - B = T.decl_buffer(16, "float64", data=b) + A = T.decl_tensor(16, "float64", data=a) + B = T.decl_tensor(16, "float64", data=b) for i in range(n): B[i] = A[i] @@ -1057,8 +1057,8 @@ def test_debug_symbol_for_buffer_var(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): - C = T.alloc_buffer((16,), "float32") + def main(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): + C = T.alloc_tensor((16,), "float32") for i in T.parallel(16): C[i] = A[i] B[i] = C[i] @@ -1071,7 +1071,7 @@ def test_subroutine_call(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer(1, dtype="float32")): + def main(A: T.Tensor(1, dtype="float32")): Module.subroutine(A.data) @T.prim_func @@ -1079,7 +1079,7 @@ def subroutine(A_data: T.handle("float32")): # The calling_conv parameter is to prevent MakePackedAPI # from changing the call signature of the subroutine. T.func_attr({"calling_conv": -1}) - A = T.decl_buffer(1, dtype="float32", data=A_data) + A = T.decl_tensor(1, dtype="float32", data=A_data) A[0] = 42.0 target = "llvm" @@ -1136,7 +1136,7 @@ def test_call_packed_without_string_arg(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.Call(tvm.ir.Op.get("tirx.tvm_call_packed"), [A.data], ty="int32") with pytest.raises(RuntimeError): @@ -1160,8 +1160,8 @@ def test_invalid_volatile_masked_buffer_load(): @I.ir_module class Module: @T.prim_func - def main(B: T.Buffer([4])): - A = T.alloc_buffer((4,), annotations={"tirx.volatile": True}) + def main(B: T.Tensor([4])): + A = T.alloc_tensor((4,), annotations={"tirx.volatile": True}) B[0:4] = T.call_intrin( "float32x4", "tirx.masked_load", @@ -1180,9 +1180,9 @@ def test_invalid_volatile_masked_decl_buffer_load(): @I.ir_module class Module: @T.prim_func - def main(B: T.Buffer([4])): - A = T.alloc_buffer((4,), annotations={"tirx.volatile": True}) - A_alias = T.decl_buffer((4,), data=A.data) + def main(B: T.Tensor([4])): + A = T.alloc_tensor((4,), annotations={"tirx.volatile": True}) + A_alias = T.decl_tensor((4,), data=A.data) B[0:4] = T.call_intrin( "float32x4", "tirx.masked_load", @@ -1202,7 +1202,7 @@ def test_invalid_volatile_masked_buffer_store(): class Module: @T.prim_func def main(): - A = T.alloc_buffer((4,), annotations={"tirx.volatile": True}) + A = T.alloc_tensor((4,), annotations={"tirx.volatile": True}) T.evaluate( T.call_intrin( "void", diff --git a/tests/python/codegen/test_target_codegen_llvm_vla.py b/tests/python/codegen/test_target_codegen_llvm_vla.py index 6c9ebfc3544c..fb7a36a0ab29 100644 --- a/tests/python/codegen/test_target_codegen_llvm_vla.py +++ b/tests/python/codegen/test_target_codegen_llvm_vla.py @@ -51,7 +51,7 @@ def test_codegen_vscale(target): vscale = tvm.tirx.vscale() @T.prim_func - def main(A: T.Buffer((5,), "int32")): + def main(A: T.Tensor((5,), "int32")): for i in range(5): A[i] = 2 * vscale @@ -83,7 +83,7 @@ def test_scalable_buffer_load_store(target): pytest.skip(f"{target} not enabled") @T.prim_func - def my_func(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def my_func(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): T.func_attr({"global_symbol": "my_module", "tirx.noalias": True}) B[T.ramp(0, 1, 4 * T.vscale())] = A[T.ramp(0, 1, 4 * T.vscale())] @@ -116,7 +116,7 @@ def test_scalable_broadcast(target): pytest.skip(f"{target} not enabled") @T.prim_func - def my_func(A: T.Buffer((128,), "float32")): + def my_func(A: T.Tensor((128,), "float32")): T.func_attr({"global_symbol": "my_module", "tirx.noalias": True}) A[T.ramp(0, 1, 4 * T.vscale())] = T.broadcast(1, 4 * T.vscale()) @@ -154,7 +154,7 @@ def test_get_active_lane_mask(target): pytest.skip(f"{target} not enabled") @T.prim_func - def before(A: T.Buffer((30,), "int1")): + def before(A: T.Tensor((30,), "int1")): for i in range(T.ceildiv(30, T.vscale() * 4)): A[i : i + T.vscale() * 4] = T.get_active_lane_mask("uint1xvscalex4", i, 30) @@ -186,7 +186,7 @@ def test_predicated_scalable_buffer(target): pytest.skip(f"{target} not enabled") @T.prim_func - def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def before(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i_0 in T.serial(T.ceildiv(16, 4 * T.vscale())): for i_1 in T.vectorized(4 * T.vscale()): diff --git a/tests/python/codegen/test_target_codegen_metal.py b/tests/python/codegen/test_target_codegen_metal.py index eac1396e8ecd..36c053e3a99b 100644 --- a/tests/python/codegen/test_target_codegen_metal.py +++ b/tests/python/codegen/test_target_codegen_metal.py @@ -35,8 +35,8 @@ def check_inf_nan(n, value, dtype): class Module: @T.prim_func def main( - A: T.Buffer((1,), dtype), - C: T.Buffer((1,), dtype), + A: T.Tensor((1,), dtype), + C: T.Tensor((1,), dtype), ): T.func_attr({"tirx.noalias": True}) for i in T.thread_binding(1, thread="threadIdx.x"): @@ -66,7 +66,7 @@ def test_unaligned_vectorize(): @tvm.script.ir_module class IRModule: @T.prim_func - def main(A: T.Buffer((2, 3), "float32"), B: T.Buffer((6,), "float32")): + def main(A: T.Tensor((2, 3), "float32"), B: T.Tensor((6,), "float32")): T.func_attr({"global_symbol": "main"}) for i0_1 in T.thread_binding(3, thread="threadIdx.x"): for i0_0 in T.vectorized(2): @@ -96,8 +96,8 @@ def check_erf(n, dtype): class Module: @T.prim_func def main( - A: T.Buffer((1,), dtype), - C: T.Buffer((1,), dtype), + A: T.Tensor((1,), dtype), + C: T.Tensor((1,), dtype), ): T.func_attr({"tirx.noalias": True}) for i0 in T.thread_binding(1, thread="threadIdx.x"): @@ -125,7 +125,7 @@ def test_ramp(): @tvm.script.ir_module class IRModule: @T.prim_func - def main(A: T.Buffer((1, 2), "int32")): + def main(A: T.Tensor((1, 2), "int32")): T.func_attr({"global_symbol": "main"}) for i in T.thread_binding(1, thread="threadIdx.x"): r: T.let = T.ramp(i, 3, 2) @@ -148,7 +148,7 @@ def test_select_vectorize(): @tvm.script.ir_module class IRModule: @T.prim_func - def main(A: T.Buffer((6), "float32"), B: T.Buffer((6,), "float32")): + def main(A: T.Tensor((6), "float32"), B: T.Tensor((6,), "float32")): T.func_attr({"global_symbol": "main"}) for i0_1 in T.thread_binding(3, thread="threadIdx.x"): for i0_0 in T.vectorized(2): @@ -175,7 +175,7 @@ def run_and_check(): @pytest.mark.skipif(not env.has_metal(), reason="need metal") def test_vectorized_uint8(): @T.prim_func - def func(A: T.Buffer((16), "uint8"), B: T.Buffer((16), "float32")): + def func(A: T.Tensor((16), "uint8"), B: T.Tensor((16), "float32")): for i in T.thread_binding(4, thread="threadIdx.x"): for j in T.vectorized(4): B[i * 4 + j] = T.Cast("float32", A[i * 4 + j]) @@ -199,7 +199,7 @@ def test_func_with_trailing_pod_params(): from tvm.support import xcode # pylint: disable=import-outside-toplevel @T.prim_func - def func(A: T.Buffer((16), "float32"), B: T.Buffer((16), "float32"), x: T.float32): + def func(A: T.Tensor((16), "float32"), B: T.Tensor((16), "float32"), x: T.float32): for i in T.thread_binding(16, thread="threadIdx.x"): B[i] = A[i] + x @@ -223,7 +223,7 @@ def test_metal_compile_callback_source_passthrough(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), "float32"), B: T.Buffer((n,), "float32")): + def main(A: T.Tensor((n,), "float32"), B: T.Tensor((n,), "float32")): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(n // 32, thread="blockIdx.x"): for i_1 in T.thread_binding(32, thread="threadIdx.x"): @@ -262,9 +262,9 @@ def test_metal_compile_callback_mixed_formats_rejected(): class Module: @T.prim_func def main( - A: T.Buffer((n,), "float32"), - B: T.Buffer((n,), "float32"), - C: T.Buffer((n,), "float32"), + A: T.Tensor((n,), "float32"), + B: T.Tensor((n,), "float32"), + C: T.Tensor((n,), "float32"), ): T.func_attr({"tirx.noalias": True}) # Two independent thread-bound regions -> two device kernels, so the @@ -303,7 +303,7 @@ def test_export_load_with_fallback(monkeypatch, tmp_path): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), "float32"), B: T.Buffer((n,), "float32")): + def main(A: T.Tensor((n,), "float32"), B: T.Tensor((n,), "float32")): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(n // 32, thread="blockIdx.x"): for i_1 in T.thread_binding(32, thread="threadIdx.x"): @@ -331,10 +331,10 @@ def kernel(): "tirx.kernel_launch_params": [], } ) - A = T.alloc_buffer((64,), "float16", scope="shared") - A_frag = T.alloc_buffer((64,), "float16", scope="metal.simdgroup") - B_frag = T.alloc_buffer((64,), "float16", scope="metal.simdgroup") - C_frag = T.alloc_buffer((64,), "float16", scope="metal.simdgroup") + A = T.alloc_tensor((64,), "float16", scope="shared") + A_frag = T.alloc_tensor((64,), "float16", scope="metal.simdgroup") + B_frag = T.alloc_tensor((64,), "float16", scope="metal.simdgroup") + C_frag = T.alloc_tensor((64,), "float16", scope="metal.simdgroup") T.metal.make_filled_simdgroup_matrix(C_frag.data, 0, T.float32(0), 8, 8) T.metal.simdgroup_load(A_frag.data, 0, A.data, 8, 8, 8, T.bool(False)) T.metal.simdgroup_store(C_frag.data, 0, A.data, 8, 8, 8, T.bool(False)) @@ -371,7 +371,7 @@ def main(n: T.int32): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer((T.min(n, 64), 2), "float32", scope="local") + scratch = T.alloc_tensor((T.min(n, 64), 2), "float32", scope="local") T.evaluate(scratch.data) source = _build_metal(Module).inspect_source() @@ -398,7 +398,7 @@ def main(n: T.int32): # Common subexpression elimination can hoist the bounded extent. extent: T.let[T.int32] = T.min(n, limit) elements: T.let[T.int32] = extent * 2 - scratch = T.alloc_buffer((elements,), "float32", scope="local") + scratch = T.alloc_tensor((elements,), "float32", scope="local") T.evaluate(scratch.data) if bounded: @@ -428,14 +428,14 @@ def main(): "tirx.is_global_func": True, } ) - state = T.alloc_buffer((1,), "int32", scope="local") + state = T.alloc_tensor((1,), "int32", scope="local") state[0] = 0 snapshot: T.let[T.int32] = state[0] state[0] = 32 difference: T.let[T.int32] = state[0] - snapshot # The snapshot is immutable, but the buffer it read has changed. # Substituting the load would incorrectly reduce this extent to 1. - scratch = T.alloc_buffer( + scratch = T.alloc_tensor( (T.min(T.max(difference, 1), 32 if bounded else 2147483647),), "float32", scope=scope, @@ -469,7 +469,7 @@ def main(n: T.uint64): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer((T.min(n, T.uint64(64)),), "float32", scope="local") + scratch = T.alloc_tensor((T.min(n, T.uint64(64)),), "float32", scope="local") T.evaluate(scratch.data) source = _build_metal(Module).inspect_source() @@ -490,7 +490,7 @@ def main(n: T.int32): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer((n,), "float32", scope="local") + scratch = T.alloc_tensor((n,), "float32", scope="local") scratch[0] = 1.0 T.evaluate(scratch[0]) @@ -515,7 +515,7 @@ def main(n: T.uint64): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer((n,), "float32", scope="local") + scratch = T.alloc_tensor((n,), "float32", scope="local") scratch[0] = 1.0 T.evaluate(scratch[0]) @@ -541,7 +541,7 @@ def main(): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer((extent,), "float32", scope="local") + scratch = T.alloc_tensor((extent,), "float32", scope="local") T.evaluate(scratch.data) with pytest.raises( @@ -565,7 +565,7 @@ def main(n: T.int32, m: T.int32, k: T.int32): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer( + scratch = T.alloc_tensor( (T.min(n, 1 << 30), T.min(m, 1 << 30), T.min(k, 1 << 30)), "uint8", scope="local", @@ -592,11 +592,11 @@ def kernel(): "tirx.kernel_launch_params": [], } ) - shared = T.alloc_buffer((16,), "float16", scope="shared") + shared = T.alloc_tensor((16,), "float16", scope="shared") typed_alias = T.ptr_byte_offset(shared.data, 4, "float16") - typed_buffer = T.decl_buffer((14,), "float16", data=typed_alias, scope="shared") + typed_buffer = T.decl_tensor((14,), "float16", data=typed_alias, scope="shared") void_alias = T.handle_add_byte_offset(shared.data, 8) - void_buffer = T.decl_buffer((12,), "float16", data=void_alias, scope="shared") + void_buffer = T.decl_tensor((12,), "float16", data=void_alias, scope="shared") typed_buffer[0] = T.float16(1) void_buffer[0] = T.float16(2) @@ -617,14 +617,14 @@ def test_pointer_byte_offsets_execute_in_threadgroup_memory(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def main(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): for bx in T.thread_binding(1, thread="blockIdx.x"): for tx in T.thread_binding(1, thread="threadIdx.x"): - shared = T.alloc_buffer((16,), "float32", scope="shared") + shared = T.alloc_tensor((16,), "float32", scope="shared") typed_alias = T.ptr_byte_offset(shared.data, 4, "float32") - typed_buffer = T.decl_buffer((15,), "float32", data=typed_alias, scope="shared") + typed_buffer = T.decl_tensor((15,), "float32", data=typed_alias, scope="shared") void_alias = T.handle_add_byte_offset(shared.data, 8) - void_buffer = T.decl_buffer((14,), "float32", data=void_alias, scope="shared") + void_buffer = T.decl_tensor((14,), "float32", data=void_alias, scope="shared") shared[0] = A[0] typed_buffer[0] = A[1] void_buffer[0] = A[2] diff --git a/tests/python/codegen/test_target_codegen_opencl.py b/tests/python/codegen/test_target_codegen_opencl.py index 88e9d0fa25f9..c6280415aeb0 100644 --- a/tests/python/codegen/test_target_codegen_opencl.py +++ b/tests/python/codegen/test_target_codegen_opencl.py @@ -35,7 +35,7 @@ def check_if_then_else(n, dtype): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1,), dtype), C: T.Buffer((1,), dtype)): + def main(A: T.Tensor((1,), dtype), C: T.Tensor((1,), dtype)): T.func_attr({"tirx.noalias": True}) for i in T.thread_binding(1, thread="threadIdx.x"): C[i] = T.max( @@ -59,7 +59,7 @@ def check_select(n, dtype): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1,), dtype), C: T.Buffer((1,), dtype)): + def main(A: T.Tensor((1,), dtype), C: T.Tensor((1,), dtype)): T.func_attr({"tirx.noalias": True}) for i in T.thread_binding(1, thread="threadIdx.x"): C[i] = T.max( @@ -94,7 +94,7 @@ def check_inf_nan(n, value, dtype): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1,), dtype), C: T.Buffer((1,), dtype)): + def main(A: T.Tensor((1,), dtype), C: T.Tensor((1,), dtype)): T.func_attr({"tirx.noalias": True}) for i in T.thread_binding(1, thread="threadIdx.x"): C[i] = T.Cast(dtype, value) @@ -124,7 +124,7 @@ def check_max(n, dtype): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1,), dtype), C: T.Buffer((1,), dtype)): + def main(A: T.Tensor((1,), dtype), C: T.Tensor((1,), dtype)): T.func_attr({"tirx.noalias": True}) for i in T.thread_binding(1, thread="threadIdx.x"): C[i] = T.max(A[0] + T.Cast(dtype, 1), T.Cast(dtype, 0)) @@ -152,7 +152,7 @@ def check_erf(n, dtype): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1,), dtype), C: T.Buffer((1,), dtype)): + def main(A: T.Tensor((1,), dtype), C: T.Tensor((1,), dtype)): T.func_attr({"tirx.noalias": True}) for i0 in T.thread_binding(1, thread="threadIdx.x"): C[i0] = T.erf(A[i0]) @@ -174,7 +174,7 @@ def test_opencl_type_casting(): @I.ir_module class Module: @T.prim_func - def main(C: T.Buffer((32,), "float32")): + def main(C: T.Tensor((32,), "float32")): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(8, thread="threadIdx.x"): for i_1 in T.vectorized(4): @@ -225,7 +225,7 @@ def _check(target, n, dtype): @I.ir_module class Module: @T.prim_func - def main(C: T.Buffer((n,), "int32")): + def main(C: T.Tensor((n,), "int32")): T.func_attr({"tirx.noalias": True}) for i in T.thread_binding(n, thread="threadIdx.x"): C[i] = T.Cast("int32", T.ceil(T.log2(T.Cast(inter_dtype, i)))) @@ -267,7 +267,7 @@ def test_export_load_with_fallback(monkeypatch, tmp_path): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), "float32"), B: T.Buffer((n,), "float32")): + def main(A: T.Tensor((n,), "float32"), B: T.Tensor((n,), "float32")): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(n // 32, thread="blockIdx.x"): for i_1 in T.thread_binding(32, thread="threadIdx.x"): diff --git a/tests/python/codegen/test_target_codegen_riscv.py b/tests/python/codegen/test_target_codegen_riscv.py index 5ab3f221339d..61f01210aa44 100644 --- a/tests/python/codegen/test_target_codegen_riscv.py +++ b/tests/python/codegen/test_target_codegen_riscv.py @@ -67,7 +67,7 @@ def test_rvv(target): def check_rvv_presence(N, extent): @T.prim_func - def load_vec(A: T.Buffer((N,), "int8")): + def load_vec(A: T.Tensor((N,), "int8")): for j in T.vectorized(0, extent): A[j] = 1 @@ -109,7 +109,7 @@ def test_rvv_vscale_llvm_dbginfo(target): # fmt: off @T.prim_func - def rvv_with_vscale(A: T.Buffer((8,), dtype='float32', align=4, offset_factor=1), B: T.Buffer((4, 8), dtype='float32', align=4, offset_factor=1, strides=[8, 1]), C: T.Buffer((4,), dtype='float32', align=4, offset_factor=1)): + def rvv_with_vscale(A: T.Tensor((8,), dtype='float32', align=4, offset_factor=1), B: T.Tensor((4, 8), dtype='float32', align=4, offset_factor=1, strides=[8, 1]), C: T.Tensor((4,), dtype='float32', align=4, offset_factor=1)): zero = T.call_llvm_intrin('float32xvscalex2', 'llvm.riscv.vfmv.v.f', T.Broadcast(T.float32(0.0), T.vscale() * 2), C[0], T.uint64(1)) vec_A = T.call_llvm_intrin('float32xvscalex4', 'llvm.riscv.vle', T.Broadcast(T.float32(0.0), T.vscale() * 4), T.tvm_access_ptr(T.type_annotation('float32'), A.data, 0, 8, 1), T.int64(8)) @@ -127,8 +127,8 @@ def rvv_with_vscale(A: T.Buffer((8,), dtype='float32', align=4, offset_factor=1) def test_rvv_fixed_width_vectorized_loop_uses_scalable_chunks(): @T.prim_func def fixed16_negative( - A: T.Buffer((14, 23, 67, 99), "float32"), - B: T.Buffer((14, 23, 67, 99), "float32"), + A: T.Tensor((14, 23, 67, 99), "float32"), + B: T.Tensor((14, 23, 67, 99), "float32"), ): for n, c, h, wo in T.grid(14, 23, 67, 7): for wi in T.vectorized(0, 16): @@ -136,7 +136,7 @@ def fixed16_negative( B[n, c, h, wo * 16 + wi] = T.float32(0) - A[n, c, h, wo * 16 + wi] @T.prim_func - def fixed16_negative_int64(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def fixed16_negative_int64(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): for wi in T.vectorized(T.int64(0), T.int64(16)): B[wi] = T.float32(0) - A[wi] @@ -168,7 +168,7 @@ def check_codegen(func): @pytest.mark.skipif(not env.has_llvm_min_version(14), reason="need llvm >= 14") def test_rvv_scalable_ramp_expression(): @T.prim_func - def ramp_compare(B: T.Buffer((16,), "int32")): + def ramp_compare(B: T.Tensor((16,), "int32")): for i in T.vectorized(16): B[i] = T.Select(i * 3 + 5 < 29, i * 3 + 5, -1) diff --git a/tests/python/codegen/test_target_codegen_rocm.py b/tests/python/codegen/test_target_codegen_rocm.py index 214a8a89b9f1..19c94c969a1b 100644 --- a/tests/python/codegen/test_target_codegen_rocm.py +++ b/tests/python/codegen/test_target_codegen_rocm.py @@ -32,7 +32,7 @@ def check_inf_nan(n, value, dtype): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1,), dtype), C: T.Buffer((1,), dtype)): + def main(A: T.Tensor((1,), dtype), C: T.Tensor((1,), dtype)): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(1, thread="blockIdx.x"): for i_1 in T.thread_binding(128, thread="threadIdx.x"): @@ -88,7 +88,7 @@ def check_rocm(dtype, n, lanes): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), vec_dtype), B: T.Buffer((n,), vec_dtype)): + def main(A: T.Tensor((n,), vec_dtype), B: T.Tensor((n,), vec_dtype)): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(num_blocks, thread="blockIdx.x"): for i_1 in T.thread_binding(4, thread="threadIdx.x"): @@ -114,13 +114,13 @@ def run_and_check(): def test_rocm_warp_shuffle(): @T.prim_func def func( - A: T.Buffer((32,), dtype="float32"), + A: T.Tensor((32,), dtype="float32"), ): for bx in T.thread_binding(1, thread="blockIdx.x"): for tx in T.thread_binding(32, thread="threadIdx.x"): - A_local = T.alloc_buffer((1,), "float32", scope="local") - mask = T.alloc_buffer((1,), "uint32", scope="local") - t0 = T.alloc_buffer((1,), "float32", scope="local") + A_local = T.alloc_tensor((1,), "float32", scope="local") + mask = T.alloc_tensor((1,), "uint32", scope="local") + t0 = T.alloc_tensor((1,), "float32", scope="local") A_local[0] = A[tx] A_local[0] = T.tvm_warp_shuffle(mask[0], A_local[0], 0, 32, 32) A[tx] = A_local[0] @@ -141,8 +141,8 @@ def run_and_check(): def test_rocm_vectorized_exp(): @T.prim_func def func( - A: T.Buffer((4,), dtype="float32"), - B: T.Buffer((4,), dtype="float32"), + A: T.Tensor((4,), dtype="float32"), + B: T.Tensor((4,), dtype="float32"), ): for bx in T.thread_binding(1, thread="blockIdx.x"): for tx in T.thread_binding(1, thread="threadIdx.x"): @@ -170,7 +170,7 @@ def test_export_load_with_fallback(monkeypatch, tmp_path): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), "float32"), B: T.Buffer((n,), "float32")): + def main(A: T.Tensor((n,), "float32"), B: T.Tensor((n,), "float32")): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(n // 32, thread="blockIdx.x"): for i_1 in T.thread_binding(32, thread="threadIdx.x"): diff --git a/tests/python/codegen/test_target_codegen_static_init.py b/tests/python/codegen/test_target_codegen_static_init.py index c5bb62c630fd..a754715e1f46 100644 --- a/tests/python/codegen/test_target_codegen_static_init.py +++ b/tests/python/codegen/test_target_codegen_static_init.py @@ -34,7 +34,7 @@ def test_cb(sh, A): @I.ir_module class Module: @T.prim_func - def ramp(Ab: T.Buffer((n,), "int64")): + def ramp(Ab: T.Tensor((n,), "int64")): T.func_attr({"global_symbol": "ramp"}) T.call_packed( diff --git a/tests/python/codegen/test_target_codegen_vulkan.py b/tests/python/codegen/test_target_codegen_vulkan.py index 154050054be9..b77009b9e636 100644 --- a/tests/python/codegen/test_target_codegen_vulkan.py +++ b/tests/python/codegen/test_target_codegen_vulkan.py @@ -66,7 +66,7 @@ def test_vector_comparison(dtype): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1024,), dtype), B: T.Buffer((1024,), dtype)): + def main(A: T.Tensor((1024,), dtype), B: T.Tensor((1024,), dtype)): for i_0 in T.thread_binding(8, thread="blockIdx.x"): for i_1 in T.thread_binding(32, thread="threadIdx.x"): for i_2 in T.vectorized(4): @@ -137,7 +137,7 @@ def test_array_vectorize_add(dtype): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((64,), vec_dtype), B: T.Buffer((64,), vec_dtype)): + def main(A: T.Tensor((64,), vec_dtype), B: T.Tensor((64,), vec_dtype)): for i_0 in T.thread_binding(16, thread="blockIdx.x"): for i_1 in T.thread_binding(4, thread="threadIdx.x"): B[i_0 * 4 + i_1] = A[i_0 * 4 + i_1] + one @@ -169,7 +169,7 @@ def test_vulkan_bool_load(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((1024,), "bool"), B: T.Buffer((1024,), "int32")): + def main(A: T.Tensor((1024,), "bool"), B: T.Tensor((1024,), "int32")): for i_0 in T.thread_binding(8, thread="blockIdx.x"): for i_1 in T.thread_binding(128, thread="threadIdx.x"): B[i_0 * 128 + i_1] = T.Cast("int32", A[i_0 * 128 + i_1]) @@ -233,8 +233,8 @@ def test_vulkan_constant_passing(vulkan_parameter_impl, vulkan_parameter_dtype): v = T_builder.arg_(f"scale{i}", tvm.tirx.Var("", dtype)) scalar_vars.append(v) n_var = T_builder.int32() - A = T_builder.arg_("var_A", T_builder.Buffer((n_var,), dtype)) - B = T_builder.arg_("var_B", T_builder.Buffer((n_var,), dtype)) + A = T_builder.arg_("var_A", T_builder.Tensor((n_var,), dtype)) + B = T_builder.arg_("var_B", T_builder.Tensor((n_var,), dtype)) T_builder.func_attr({"tirx.noalias": True}) scalar_sum = scalar_vars[0] @@ -276,9 +276,9 @@ def test_vulkan_while_if(): dtype = "int32" @T.prim_func - def while_if_gpu(A: T.Buffer((1,), "int32"), B: T.Buffer((1,), "int32")): + def while_if_gpu(A: T.Tensor((1,), "int32"), B: T.Tensor((1,), "int32")): for bx in T.thread_binding(1, thread="blockIdx.x"): - iterations = T.decl_buffer((1,), "int32", scope="local") + iterations = T.decl_tensor((1,), "int32", scope="local") iterations[0] = 0 B[0] = 0 while iterations[0] < T.if_then_else(A[0] > 0, 10, 20): @@ -310,7 +310,7 @@ def test_vulkan_local_threadidx(): n = 32 @T.prim_func - def local_threadidx_func(A: T.Buffer((32,), "int32"), B: T.Buffer((32,), "int32")): + def local_threadidx_func(A: T.Tensor((32,), "int32"), B: T.Tensor((32,), "int32")): # First block with thread extent 16 for _ in range(1): for tx in T.thread_binding(16, thread="threadIdx.x"): @@ -351,7 +351,7 @@ def test_vectorized_index_ramp(): class Module: @T.prim_func def main( - A: T.Buffer((n,), "int32", offset_factor=1), B: T.Buffer((n,), "int32", offset_factor=1) + A: T.Tensor((n,), "int32", offset_factor=1), B: T.Tensor((n,), "int32", offset_factor=1) ): T.func_attr({"tirx.noalias": True}) @@ -389,7 +389,7 @@ def test_vectorized_index_broadcast(): class Module: @T.prim_func def main( - A: T.Buffer((n,), "int32", offset_factor=1), B: T.Buffer((n,), "int32", offset_factor=1) + A: T.Tensor((n,), "int32", offset_factor=1), B: T.Tensor((n,), "int32", offset_factor=1) ): T.func_attr({"tirx.noalias": True}) @@ -437,7 +437,7 @@ def test_negative_operand_divmod(): divisor = 5 @T.prim_func - def func(A: T.Buffer((N, 2), "int32")): + def func(A: T.Tensor((N, 2), "int32")): for i in T.thread_binding(N, thread="threadIdx.x"): A[i, 0] = T.floordiv(i - offset, divisor) A[i, 1] = T.floormod(i - offset, divisor) @@ -463,13 +463,13 @@ def test_cooperative_matrix(out_dtype): @I.ir_module class Module: @T.prim_func - def main(X: T.Buffer((16, 32), "float16"), W: T.Buffer((32, 16), "float16"), compute: T.Buffer((16, 16), out_dtype)): + def main(X: T.Tensor((16, 32), "float16"), W: T.Tensor((32, 16), "float16"), compute: T.Tensor((16, 16), out_dtype)): T.func_attr({"tirx.noalias": True}) - X_shared = T.alloc_buffer((16, 32), "float16", scope="shared") - W_shared = T.alloc_buffer((32, 16), "float16", scope="shared") - X_shared_wmma_matrix_a = T.alloc_buffer((16, 32), "float16", scope="wmma.matrix_a") - W_shared_wmma_matrix_b = T.alloc_buffer((32, 16), "float16", scope="wmma.matrix_b") - compute_wmma_accumulator = T.alloc_buffer((16, 16), out_dtype, scope="wmma.accumulator") + X_shared = T.alloc_tensor((16, 32), "float16", scope="shared") + W_shared = T.alloc_tensor((32, 16), "float16", scope="shared") + X_shared_wmma_matrix_a = T.alloc_tensor((16, 32), "float16", scope="wmma.matrix_a") + W_shared_wmma_matrix_b = T.alloc_tensor((32, 16), "float16", scope="wmma.matrix_b") + compute_wmma_accumulator = T.alloc_tensor((16, 16), out_dtype, scope="wmma.accumulator") for i_0_j_0_fused in T.thread_binding(1, thread="blockIdx.x"): T.tvm_fill_fragment(compute_wmma_accumulator.data, 16, 16, 16, 0, T.float32(0.0)) for k_0 in range(2): @@ -514,15 +514,15 @@ def run_and_check(): @pytest.mark.gpu @pytest.mark.skipif(not env.has_vulkan(), reason="need vulkan") def test_codegen_decl_buffer(): - """DeclBuffer aliases should retain their backing storage metadata.""" + """DeclTensor aliases should retain their backing storage metadata.""" @I.ir_module class AllocationBacked: @T.prim_func def kernel(): T.func_attr({"calling_conv": 2, "global_symbol": "kernel", "tirx.noalias": True}) - A = T.alloc_buffer((256,), dtype="float32", scope="local") - A_buf = T.decl_buffer([256], dtype="float32", scope="local", data=A.data) + A = T.alloc_tensor((256,), dtype="float32", scope="local") + A_buf = T.decl_tensor([256], dtype="float32", scope="local", data=A.data) A_buf[0] = T.float32(1) T.evaluate(A_buf[0]) @@ -533,9 +533,9 @@ def kernel(): @I.ir_module class ParameterBacked: @T.prim_func - def main(A: T.Buffer((1,), "float32"), B: T.Buffer((1,), "float32")): - A_buf = T.decl_buffer([1], dtype="float32", data=A.data) - B_buf = T.decl_buffer([1], dtype="float32", data=B.data) + def main(A: T.Tensor((1,), "float32"), B: T.Tensor((1,), "float32")): + A_buf = T.decl_tensor([1], dtype="float32", data=A.data) + B_buf = T.decl_tensor([1], dtype="float32", data=B.data) for tx in T.thread_binding(1, thread="threadIdx.x"): B_buf[tx] = A_buf[tx] @@ -550,8 +550,8 @@ def test_codegen_static_shared_memory(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): - A_shared = T.alloc_buffer((128,), dtype="float32", scope="shared") + def main(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): + A_shared = T.alloc_tensor((128,), dtype="float32", scope="shared") for bx in T.thread_binding(1, thread="blockIdx.x"): for tx in T.thread_binding(128, thread="threadIdx.x"): @@ -587,7 +587,7 @@ def run_test(tvm_intrin, np_func): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((m,), "float32"), B: T.Buffer((m,), "float32")): + def main(A: T.Tensor((m,), "float32"), B: T.Tensor((m,), "float32")): for i_0 in T.thread_binding((m + 63) // 64, thread="blockIdx.x"): for i_1 in T.thread_binding(64, thread="threadIdx.x"): if i_0 * 64 + i_1 < m: @@ -627,7 +627,7 @@ def test_export_load_with_fallback(monkeypatch, tmp_path): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), "float32"), B: T.Buffer((n,), "float32")): + def main(A: T.Tensor((n,), "float32"), B: T.Tensor((n,), "float32")): T.func_attr({"tirx.noalias": True}) for i_0 in T.thread_binding(n // 32, thread="blockIdx.x"): for i_1 in T.thread_binding(32, thread="threadIdx.x"): diff --git a/tests/python/codegen/test_target_codegen_webgpu.py b/tests/python/codegen/test_target_codegen_webgpu.py index eaca6e0ce77c..e10af436b0ff 100644 --- a/tests/python/codegen/test_target_codegen_webgpu.py +++ b/tests/python/codegen/test_target_codegen_webgpu.py @@ -31,7 +31,7 @@ def test_codegen_buffer_access_modes(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((8,), "float32"), B: T.Buffer((8,), "float32")): + def main(A: T.Tensor((8,), "float32"), B: T.Tensor((8,), "float32")): for tx in T.thread_binding(8, thread="threadIdx.x"): B[tx] = A[tx] @@ -60,7 +60,7 @@ def main(n: T.int32): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer((T.min(n, 64), 2), "float32", scope="local") + scratch = T.alloc_tensor((T.min(n, 64), 2), "float32", scope="local") T.evaluate(scratch.data) source = _build_webgpu(Module).inspect_source() @@ -86,9 +86,9 @@ def main(n: T.int32): ) # Common subexpression elimination can hoist the bounded extent. extent: T.let[T.int32] = T.min(n, limit) - first = T.alloc_buffer((extent * 2,), "float32", scope=scope) + first = T.alloc_tensor((extent * 2,), "float32", scope=scope) elements: T.let[T.int32] = extent * 2 - second = T.alloc_buffer((elements,), "float32", scope=scope) + second = T.alloc_tensor((elements,), "float32", scope=scope) first[0] = 1.0 second[0] = first[0] @@ -119,14 +119,14 @@ def main(): "tirx.is_global_func": True, } ) - state = T.alloc_buffer((1,), "int32", scope="local") + state = T.alloc_tensor((1,), "int32", scope="local") state[0] = 0 snapshot: T.let[T.int32] = state[0] state[0] = 32 difference: T.let[T.int32] = state[0] - snapshot # The snapshot is immutable, but the buffer it read has changed. # Substituting the load would incorrectly reduce this extent to 1. - scratch = T.alloc_buffer( + scratch = T.alloc_tensor( (T.min(T.max(difference, 1), 32 if bounded else 2147483647),), "float32", scope=scope, @@ -162,7 +162,7 @@ def main(n: T.int32): ) extent: T.let[T.int32] = T.min(n, 64) elements: T.let[T.int32] = extent * 2 - scratch = T.alloc_buffer((elements,), "float32", scope="shared") + scratch = T.alloc_tensor((elements,), "float32", scope="shared") scratch[0] = 1.0 target = {"kind": "webgpu", "max_shared_memory_per_block": target_limit} @@ -190,7 +190,7 @@ def main(n: T.int32): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer((n,), "float32", scope="local") + scratch = T.alloc_tensor((n,), "float32", scope="local") scratch[0] = 1.0 T.evaluate(scratch[0]) @@ -215,7 +215,7 @@ def main(): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer((extent,), "float32", scope="local") + scratch = T.alloc_tensor((extent,), "float32", scope="local") T.evaluate(scratch.data) with pytest.raises( @@ -238,7 +238,7 @@ def main(n: T.int32, m: T.int32, k: T.int32): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer( + scratch = T.alloc_tensor( (T.min(n, 1 << 30), T.min(m, 1 << 30), T.min(k, 1 << 30)), "uint8", scope="local", @@ -264,7 +264,7 @@ def main(n: T.int32, m: T.int32): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer( + scratch = T.alloc_tensor( (T.min(n, 1 << 30), T.min(m, 1 << 30), 4), "float32", scope="local" ) T.evaluate(scratch.data) @@ -288,7 +288,7 @@ def main(): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer((8192,), "float32", scope="shared") + scratch = T.alloc_tensor((8192,), "float32", scope="shared") scratch[0] = 1.0 source = _build_webgpu(Module).inspect_source() @@ -308,8 +308,8 @@ def main(): "tirx.is_global_func": True, } ) - first = T.alloc_buffer((4096,), "float32", scope="shared") - second = T.alloc_buffer((4097,), "float32", scope="shared") + first = T.alloc_tensor((4096,), "float32", scope="shared") + second = T.alloc_tensor((4097,), "float32", scope="shared") first[0] = 1.0 second[0] = 2.0 @@ -333,8 +333,8 @@ def main(): "tirx.is_global_func": True, } ) - first = T.alloc_buffer((1,), "float32", scope="shared") - second = T.alloc_buffer((1,), "float32", scope="shared") + first = T.alloc_tensor((1,), "float32", scope="shared") + second = T.alloc_tensor((1,), "float32", scope="shared") first[0] = 1.0 second[0] = 2.0 @@ -358,7 +358,7 @@ def main(): "tirx.is_global_func": True, } ) - scratch = T.alloc_buffer((16384,), "float32", scope="shared") + scratch = T.alloc_tensor((16384,), "float32", scope="shared") scratch[0] = 1.0 _build_webgpu(Module, {"kind": "webgpu", "max_shared_memory_per_block": 65536}) diff --git a/tests/python/codegen/test_target_codegen_x86.py b/tests/python/codegen/test_target_codegen_x86.py index 234551dc6126..8dd8e9e64602 100644 --- a/tests/python/codegen/test_target_codegen_x86.py +++ b/tests/python/codegen/test_target_codegen_x86.py @@ -42,8 +42,8 @@ def fp16_to_fp32(target, width, match=None, not_match=None): class Module: @T.prim_func def main( - A: T.Buffer((elements, width), "float16"), - B: T.Buffer((elements, width), "float32"), + A: T.Tensor((elements, width), "float16"), + B: T.Tensor((elements, width), "float32"), ): T.func_attr({"tirx.noalias": True}) for i0 in range(elements): diff --git a/tests/python/contrib/test_android/test_meta_schedule.py b/tests/python/contrib/test_android/test_meta_schedule.py index d88dcdf2fd3c..8a29244c6d41 100644 --- a/tests/python/contrib/test_android/test_meta_schedule.py +++ b/tests/python/contrib/test_android/test_meta_schedule.py @@ -36,7 +36,7 @@ @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) diff --git a/tests/python/contrib/test_tir_triton_integration.py b/tests/python/contrib/test_tir_triton_integration.py index 89b384122f76..2f7f37121253 100644 --- a/tests/python/contrib/test_tir_triton_integration.py +++ b/tests/python/contrib/test_tir_triton_integration.py @@ -71,9 +71,9 @@ def add_kernel( class Module: @Ts.prim_func def add( - x: T.Buffer((add_m,), "float32"), - y: T.Buffer((add_m,), "float32"), - output: T.Buffer((add_m,), "float32"), + x: T.Tensor((add_m,), "float32"), + y: T.Tensor((add_m,), "float32"), + output: T.Tensor((add_m,), "float32"), ) -> None: T.func_attr({"global_symbol": "add"}) @@ -110,7 +110,7 @@ def main(x: R.Tensor((main_m,), "float32"), y: R.Tensor((main_m,), "float32")): @I.ir_module class Parsed: @Ts.prim_func - def add(x: T.Buffer((m,)), y: T.Buffer((m,)), output: T.Buffer((m,))): + def add(x: T.Tensor((m,)), y: T.Tensor((m,)), output: T.Tensor((m,))): with Ts.sblock("root"): Ts.reads(x[0:m], y[0:m]) Ts.writes(output[0:m]) diff --git a/tests/python/disco/test_callback.py b/tests/python/disco/test_callback.py index e4cb3430bb65..ebeec07e98c7 100644 --- a/tests/python/disco/test_callback.py +++ b/tests/python/disco/test_callback.py @@ -47,9 +47,9 @@ def test_callback(): class Module: @Ts.prim_func(private=True) def slice_A( - A: T.Buffer((4, 4), "int32"), + A: T.Tensor((4, 4), "int32"), rank: T.int64, - A_sharded: T.Buffer((2, 4), "int32"), + A_sharded: T.Tensor((2, 4), "int32"), ): for i, j in T.grid(2, 4): with Ts.sblock("slice_A"): @@ -58,9 +58,9 @@ def slice_A( @Ts.prim_func(private=True) def slice_B( - B: T.Buffer((2, 2), "float32"), + B: T.Tensor((2, 2), "float32"), rank: T.int64, - B_sharded: T.Buffer((2, 1), "float32"), + B_sharded: T.Tensor((2, 1), "float32"), ): for i in range(2): with Ts.sblock("slice_B"): diff --git a/tests/python/disco/test_nvshmem.py b/tests/python/disco/test_nvshmem.py index 028af969e773..24a2bbeeff82 100644 --- a/tests/python/disco/test_nvshmem.py +++ b/tests/python/disco/test_nvshmem.py @@ -232,7 +232,7 @@ def _compile(): sess.sync_worker_0() @Ts.prim_func - def main(A: T.Buffer((8, 16), "float32"), B: T.Buffer((16, 8), "float32")): + def main(A: T.Tensor((8, 16), "float32"), B: T.Tensor((16, 8), "float32")): for i in T.thread_binding(T.int64(8), thread="threadIdx.y"): for j in T.thread_binding(T.int64(16), thread="threadIdx.x"): with Ts.sblock("T_transpose"): @@ -308,8 +308,8 @@ def _kernel_compile(compile_mode): class NvshmemQueryModule: @Ts.prim_func def query_pe( - my_pe_out: T.Buffer((1,), "int32"), - n_pes_out: T.Buffer((1,), "int32"), + my_pe_out: T.Tensor((1,), "int32"), + n_pes_out: T.Tensor((1,), "int32"), ): with Ts.sblock("root"): Ts.reads() diff --git a/tests/python/disco/test_session.py b/tests/python/disco/test_session.py index 6f9116b84f85..8aa7f8f01c27 100644 --- a/tests/python/disco/test_session.py +++ b/tests/python/disco/test_session.py @@ -232,7 +232,7 @@ def test_vm_module(session_kind): @I.ir_module class TestMod: @Ts.prim_func - def transpose(A: T.Buffer((8, 16), "float32"), B: T.Buffer((16, 8), "float32")): + def transpose(A: T.Tensor((8, 16), "float32"), B: T.Tensor((16, 8), "float32")): for i, j in T.grid(16, 8): with Ts.sblock("transpose"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -276,14 +276,14 @@ def test_vm_multi_func(session_kind): @I.ir_module class TestMod: @Ts.prim_func - def t1(A: T.Buffer((8, 16), "float32"), B: T.Buffer((16, 8), "float32")): + def t1(A: T.Tensor((8, 16), "float32"), B: T.Tensor((16, 8), "float32")): for i, j in T.grid(16, 8): with Ts.sblock("t1"): vi, vj = Ts.axis.remap("SS", [i, j]) B[vi, vj] = A[vj, vi] @Ts.prim_func - def t2(A: T.Buffer((16, 8), "float32"), B: T.Buffer((8, 16), "float32")): + def t2(A: T.Tensor((16, 8), "float32"), B: T.Tensor((8, 16), "float32")): for i, j in T.grid(8, 16): with Ts.sblock("t2"): vi, vj = Ts.axis.remap("SS", [i, j]) diff --git a/tests/python/driver/test_compile.py b/tests/python/driver/test_compile.py index 6a9ba7f6e575..0ca6817aab47 100644 --- a/tests/python/driver/test_compile.py +++ b/tests/python/driver/test_compile.py @@ -91,7 +91,7 @@ def test_compile_mixed_module(): @tvm.script.ir_module class MyModule: @Ts.prim_func - def add_one(X: T.Buffer((4,), "float32"), Y: T.Buffer((4,), "float32")): + def add_one(X: T.Tensor((4,), "float32"), Y: T.Tensor((4,), "float32")): for i in range(4): Y[i] = X[i] + 1 diff --git a/tests/python/ir/test_datatype_nv_fp4.py b/tests/python/ir/test_datatype_nv_fp4.py index 521bdd13a635..af6bd7203510 100644 --- a/tests/python/ir/test_datatype_nv_fp4.py +++ b/tests/python/ir/test_datatype_nv_fp4.py @@ -48,7 +48,7 @@ def test_create_nv_fp4_nd_array(np_dtype, dtype_str): def test_nv_fp4_buffer(np_dtype, dtype_str): m = te.var("m") n = te.var("n") - A = tvm.tirx.decl_buffer((m, n), dtype_str) + A = tvm.tirx.decl_tensor((m, n), dtype_str) assert A.dtype == dtype_str diff --git a/tests/python/ir/test_datatype_nv_fp8.py b/tests/python/ir/test_datatype_nv_fp8.py index 059e7f78eb2b..638f7158df33 100644 --- a/tests/python/ir/test_datatype_nv_fp8.py +++ b/tests/python/ir/test_datatype_nv_fp8.py @@ -44,13 +44,13 @@ def fp8_unary(dtype: str): @Ts.prim_func def func( - A: T.Buffer([128], dtype=dtype), - B: T.Buffer([128], dtype=dtype), - A_add_B: T.Buffer([128], dtype=dtype), - A_sub_B: T.Buffer([128], dtype=dtype), - A_mul_B: T.Buffer([128], dtype=dtype), - A_fp32: T.Buffer([128], dtype="float32"), - A_roundtrip: T.Buffer([128], dtype=dtype), + A: T.Tensor([128], dtype=dtype), + B: T.Tensor([128], dtype=dtype), + A_add_B: T.Tensor([128], dtype=dtype), + A_sub_B: T.Tensor([128], dtype=dtype), + A_mul_B: T.Tensor([128], dtype=dtype), + A_fp32: T.Tensor([128], dtype="float32"), + A_roundtrip: T.Tensor([128], dtype=dtype), ) -> None: for i in range(128): with Ts.sblock("fp8_unary"): @@ -120,7 +120,7 @@ def test_fp8_unary_op(np_dtype, dtype_str): def test_nv_fp8_buffer(np_dtype, dtype_str): m = te.var("m") n = te.var("n") - A = tvm.tirx.decl_buffer((m, n), dtype_str) + A = tvm.tirx.decl_tensor((m, n), dtype_str) assert A.dtype == dtype_str diff --git a/tests/python/ir/test_pass_instrument.py b/tests/python/ir/test_pass_instrument.py index 3548d910ec26..22d15930ed36 100644 --- a/tests/python/ir/test_pass_instrument.py +++ b/tests/python/ir/test_pass_instrument.py @@ -30,7 +30,7 @@ def test_tir_print_all_passes(capsys): @Ts.prim_func - def func(A: T.Buffer((128, 128, 128, 128)), B: T.Buffer((128, 128, 128, 128))) -> None: + def func(A: T.Tensor((128, 128, 128, 128)), B: T.Tensor((128, 128, 128, 128))) -> None: for i, j, k, l in T.grid(128, 128, 128, 128): with Ts.sblock("B"): vi, vj, vk, vl = Ts.axis.remap("SSSS", [i, j, k, l]) diff --git a/tests/python/relax/backend/adreno/mod_utils.py b/tests/python/relax/backend/adreno/mod_utils.py index 616cef9bd887..830d5c554ed9 100644 --- a/tests/python/relax/backend/adreno/mod_utils.py +++ b/tests/python/relax/backend/adreno/mod_utils.py @@ -752,14 +752,14 @@ def main( @Ts.prim_func def dequantize( - lm_head_q_weight1: T.Buffer((T.int64(K // 8), T.int64(N)), "uint32"), - lm_head_q_scale1: T.Buffer((T.int64(K // 32), T.int64(N)), "float16"), - dequantize: T.Buffer((T.int64(K), T.int64(N)), "float16"), + lm_head_q_weight1: T.Tensor((T.int64(K // 8), T.int64(N)), "uint32"), + lm_head_q_scale1: T.Tensor((T.int64(K // 32), T.int64(N)), "float16"), + dequantize: T.Tensor((T.int64(K), T.int64(N)), "float16"), ): T.func_attr({"tirx.noalias": T.bool(True)}) # with Ts.sblock("root"): - compute = T.alloc_buffer((T.int64(K), T.int64(N)), "float16") + compute = T.alloc_tensor((T.int64(K), T.int64(N)), "float16") for i0, i1 in T.grid(T.int64(K), T.int64(N)): with Ts.sblock("compute"): v_i0, v_i1 = Ts.axis.remap("SS", [i0, i1]) @@ -814,14 +814,14 @@ def main( @Ts.prim_func def dequantize( - lm_head_q_weight1: T.Buffer((T.int64(K // 8), vocab_size_dequantize), "uint32"), - lm_head_q_scale1: T.Buffer((T.int64(K // 32), vocab_size_dequantize), "float16"), - dequantize: T.Buffer((T.int64(K), vocab_size_dequantize), "float16"), + lm_head_q_weight1: T.Tensor((T.int64(K // 8), vocab_size_dequantize), "uint32"), + lm_head_q_scale1: T.Tensor((T.int64(K // 32), vocab_size_dequantize), "float16"), + dequantize: T.Tensor((T.int64(K), vocab_size_dequantize), "float16"), ): T.func_attr({"tirx.noalias": T.bool(True)}) # with Ts.sblock("root"): - compute = T.alloc_buffer((T.int64(K), vocab_size_dequantize), "float16") + compute = T.alloc_tensor((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]) diff --git a/tests/python/relax/backend/adreno/test_transform_fold_vdevice_scope_change.py b/tests/python/relax/backend/adreno/test_transform_fold_vdevice_scope_change.py index d88e59cf1c35..4cfd73586348 100644 --- a/tests/python/relax/backend/adreno/test_transform_fold_vdevice_scope_change.py +++ b/tests/python/relax/backend/adreno/test_transform_fold_vdevice_scope_change.py @@ -47,8 +47,8 @@ class Input: @Ts.prim_func(private=True) def max_pool2d_opencl( - gv: T.Buffer((T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32"), - pool_max: T.Buffer( + gv: T.Tensor((T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32"), + pool_max: T.Tensor( (T.int64(2), T.int64(1), T.int64(13), T.int64(13), T.int64(4)), "float32" ), ): @@ -88,8 +88,8 @@ def max_pool2d_opencl( @Ts.prim_func(private=True) def te_layout_transform( - x: T.Buffer((T.int64(2), T.int64(4), T.int64(26), T.int64(26)), "float32"), - te_layout_transform: T.Buffer( + x: T.Tensor((T.int64(2), T.int64(4), T.int64(26), T.int64(26)), "float32"), + te_layout_transform: T.Tensor( (T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32" ), ): @@ -109,10 +109,10 @@ def te_layout_transform( @Ts.prim_func(private=True) def te_layout_transform2( - lv2: T.Buffer( + lv2: T.Tensor( (T.int64(2), T.int64(1), T.int64(13), T.int64(13), T.int64(4)), "float32" ), - te_layout_transform: T.Buffer( + te_layout_transform: T.Tensor( (T.int64(2), T.int64(4), T.int64(13), T.int64(13)), "float32" ), ): @@ -172,8 +172,8 @@ class Expected: @Ts.prim_func(private=True) def max_pool2d_opencl( - gv: T.Buffer((T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32"), - pool_max: T.Buffer( + gv: T.Tensor((T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32"), + pool_max: T.Tensor( (T.int64(2), T.int64(1), T.int64(13), T.int64(13), T.int64(4)), "float32" ), ): @@ -213,8 +213,8 @@ def max_pool2d_opencl( @Ts.prim_func(private=True) def te_layout_transform( - x: T.Buffer((T.int64(2), T.int64(4), T.int64(26), T.int64(26)), "float32"), - te_layout_transform: T.Buffer( + x: T.Tensor((T.int64(2), T.int64(4), T.int64(26), T.int64(26)), "float32"), + te_layout_transform: T.Tensor( (T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32" ), ): @@ -234,10 +234,10 @@ def te_layout_transform( @Ts.prim_func(private=True) def te_layout_transform2( - lv2: T.Buffer( + lv2: T.Tensor( (T.int64(2), T.int64(1), T.int64(13), T.int64(13), T.int64(4)), "float32" ), - te_layout_transform: T.Buffer( + te_layout_transform: T.Tensor( (T.int64(2), T.int64(4), T.int64(13), T.int64(13)), "float32" ), ): diff --git a/tests/python/relax/distributed/test_distributed_transform_lower_distir.py b/tests/python/relax/distributed/test_distributed_transform_lower_distir.py index fec676fa6824..44178145b744 100644 --- a/tests/python/relax/distributed/test_distributed_transform_lower_distir.py +++ b/tests/python/relax/distributed/test_distributed_transform_lower_distir.py @@ -39,8 +39,8 @@ class MLP: @Ts.prim_func(private=True) def gelu1( - A: T.Buffer((T.int64(128), T.int64(64)), "float32"), - T_multiply: T.Buffer((T.int64(128), T.int64(64)), "float32"), + A: T.Tensor((T.int64(128), T.int64(64)), "float32"), + T_multiply: T.Tensor((T.int64(128), T.int64(64)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -81,9 +81,9 @@ def gelu1( @Ts.prim_func(private=True) def matmul1( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(64)), "float32"), - matmul_1: T.Buffer((T.int64(128), T.int64(64)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(64)), "float32"), + matmul_1: T.Tensor((T.int64(128), T.int64(64)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -98,9 +98,9 @@ def matmul1( @Ts.prim_func(private=True) def matmul2( - A: T.Buffer((T.int64(128), T.int64(64)), "float32"), - B: T.Buffer((T.int64(64), T.int64(128)), "float32"), - matmul_1: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(64)), "float32"), + B: T.Tensor((T.int64(64), T.int64(128)), "float32"), + matmul_1: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -198,8 +198,8 @@ class MLPWithTuple: @Ts.prim_func(private=True) def gelu1( - A: T.Buffer((T.int64(128), T.int64(64)), "float32"), - T_multiply: T.Buffer((T.int64(128), T.int64(64)), "float32"), + A: T.Tensor((T.int64(128), T.int64(64)), "float32"), + T_multiply: T.Tensor((T.int64(128), T.int64(64)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -240,9 +240,9 @@ def gelu1( @Ts.prim_func(private=True) def matmul11( - A: T.Buffer((T.int64(64), T.int64(64)), "float32"), - B: T.Buffer((T.int64(64), T.int64(128)), "float32"), - matmul: T.Buffer((T.int64(64), T.int64(128)), "float32"), + A: T.Tensor((T.int64(64), T.int64(64)), "float32"), + B: T.Tensor((T.int64(64), T.int64(128)), "float32"), + matmul: T.Tensor((T.int64(64), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -257,9 +257,9 @@ def matmul11( @Ts.prim_func(private=True) def matmul2( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(64)), "float32"), - matmul: T.Buffer((T.int64(128), T.int64(64)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(64)), "float32"), + matmul: T.Tensor((T.int64(128), T.int64(64)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -274,9 +274,9 @@ def matmul2( @Ts.prim_func(private=True) def split11( - A: T.Buffer((128, 64), "float32"), - T_split: T.Buffer((64, 64), "float32"), - T_split_1: T.Buffer((64, 64), "float32"), + A: T.Tensor((128, 64), "float32"), + T_split: T.Tensor((64, 64), "float32"), + T_split_1: T.Tensor((64, 64), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): diff --git a/tests/python/relax/distributed/test_distributed_transform_lower_global_to_local_view.py b/tests/python/relax/distributed/test_distributed_transform_lower_global_to_local_view.py index d453bc7b81af..d8fe0f96cfdd 100644 --- a/tests/python/relax/distributed/test_distributed_transform_lower_global_to_local_view.py +++ b/tests/python/relax/distributed/test_distributed_transform_lower_global_to_local_view.py @@ -39,8 +39,8 @@ class MLP: @Ts.prim_func(private=True) def gelu( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - T_multiply: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + T_multiply: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -81,9 +81,9 @@ def gelu( @Ts.prim_func(private=True) def matmul( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(128)), "float32"), - matmul_1: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(128)), "float32"), + matmul_1: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -128,8 +128,8 @@ class Expected: @Ts.prim_func(private=True) def gelu1( - A: T.Buffer((T.int64(128), T.int64(64)), "float32"), - T_multiply: T.Buffer((T.int64(128), T.int64(64)), "float32"), + A: T.Tensor((T.int64(128), T.int64(64)), "float32"), + T_multiply: T.Tensor((T.int64(128), T.int64(64)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -170,9 +170,9 @@ def gelu1( @Ts.prim_func(private=True) def matmul1( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(64)), "float32"), - matmul_1: T.Buffer((T.int64(128), T.int64(64)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(64)), "float32"), + matmul_1: T.Tensor((T.int64(128), T.int64(64)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -187,9 +187,9 @@ def matmul1( @Ts.prim_func(private=True) def matmul2( - A: T.Buffer((T.int64(128), T.int64(64)), "float32"), - B: T.Buffer((T.int64(64), T.int64(128)), "float32"), - matmul_1: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(64)), "float32"), + B: T.Tensor((T.int64(64), T.int64(128)), "float32"), + matmul_1: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -244,9 +244,9 @@ class LlamaAttentionLayer: @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), - B: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), - T_add: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + B: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + T_add: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -259,9 +259,9 @@ def add( @Ts.prim_func(private=True) def divide( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - T_divide: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + T_divide: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -276,9 +276,9 @@ def divide( @Ts.prim_func(private=True) def matmul( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), - B: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - matmul: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + B: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + matmul: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -295,9 +295,9 @@ def matmul( @Ts.prim_func(private=True) def matmul1( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), - B: T.Buffer((T.int64(1), T.int64(32), T.int64(128), T.int64(256)), "float16"), - matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + B: T.Tensor((T.int64(1), T.int64(32), T.int64(128), T.int64(256)), "float16"), + matmul: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -317,9 +317,9 @@ def matmul1( @Ts.prim_func(private=True) def matmul2( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), - matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + matmul: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -339,9 +339,9 @@ def matmul2( @Ts.prim_func(private=True) def maximum( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - T_maximum: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + T_maximum: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -356,9 +356,9 @@ def maximum( @Ts.prim_func(private=True) def minimum( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(1), T.int64(256), T.int64(256)), "float16"), - T_minimum: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(1), T.int64(256), T.int64(256)), "float16"), + T_minimum: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -373,8 +373,8 @@ def minimum( @Ts.prim_func(private=True) def reshape( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), - T_reshape: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + T_reshape: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -398,8 +398,8 @@ def reshape( @Ts.prim_func(private=True) def reshape1( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), - T_reshape: T.Buffer((T.int64(256), T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + T_reshape: T.Tensor((T.int64(256), T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -424,8 +424,8 @@ def reshape1( @Ts.prim_func(private=True) def reshape2( - A: T.Buffer((T.int64(256), T.int64(32), T.int64(128)), "float16"), - T_reshape: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(256), T.int64(32), T.int64(128)), "float16"), + T_reshape: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -448,8 +448,8 @@ def reshape2( @Ts.prim_func(private=True) def reshape3( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), - T_reshape: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + T_reshape: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -474,9 +474,9 @@ def reshape3( @Ts.prim_func def rms_norm( - A: T.Buffer((T.int64(1), 256, T.int64(4096)), "float16"), - B: T.Buffer((T.int64(4096),), "float16"), - rms_norm_1: T.Buffer((T.int64(1), 256, T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), 256, T.int64(4096)), "float16"), + B: T.Tensor((T.int64(4096),), "float16"), + rms_norm_1: T.Tensor((T.int64(1), 256, T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -510,10 +510,10 @@ def rms_norm( @Ts.prim_func def rotary_embedding( - A: T.Buffer((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), - B: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - C: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - rotary: T.Buffer((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), + B: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + C: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + rotary: T.Tensor((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -536,8 +536,8 @@ def rotary_embedding( @Ts.prim_func(private=True) def softmax( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - T_softmax_norm: T.Buffer( + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + T_softmax_norm: T.Tensor( (T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16" ), ): @@ -594,8 +594,8 @@ def softmax( @Ts.prim_func(private=True) def transpose( - A: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - T_transpose: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + T_transpose: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -608,8 +608,8 @@ def transpose( @Ts.prim_func(private=True) def transpose1( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), - T_transpose: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + T_transpose: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -622,8 +622,8 @@ def transpose1( @Ts.prim_func(private=True) def transpose2( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), - T_transpose: T.Buffer((T.int64(1), T.int64(32), T.int64(128), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + T_transpose: T.Tensor((T.int64(1), T.int64(32), T.int64(128), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -636,8 +636,8 @@ def transpose2( @Ts.prim_func(private=True) def transpose3( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), - T_transpose: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + T_transpose: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -857,9 +857,9 @@ class Expected: @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), - B: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), - T_add: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + B: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + T_add: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -872,9 +872,9 @@ def add( @Ts.prim_func(private=True) def divide1( - A: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), - T_divide: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), + T_divide: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -889,9 +889,9 @@ def divide1( @Ts.prim_func(private=True) def matmul11( - A: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), - B: T.Buffer((T.int64(1), T.int64(16), T.int64(128), T.int64(256)), "float16"), - matmul: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), + B: T.Tensor((T.int64(1), T.int64(16), T.int64(128), T.int64(256)), "float16"), + matmul: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -911,9 +911,9 @@ def matmul11( @Ts.prim_func(private=True) def matmul21( - A: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), - matmul: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), + matmul: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -933,9 +933,9 @@ def matmul21( @Ts.prim_func(private=True) def matmul3( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), - B: T.Buffer((T.int64(4096), T.int64(2048)), "float16"), - matmul: T.Buffer((T.int64(1), T.int64(256), T.int64(2048)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + B: T.Tensor((T.int64(4096), T.int64(2048)), "float16"), + matmul: T.Tensor((T.int64(1), T.int64(256), T.int64(2048)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -952,9 +952,9 @@ def matmul3( @Ts.prim_func(private=True) def matmul4( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(2048)), "float16"), - B: T.Buffer((T.int64(2048), T.int64(4096)), "float16"), - matmul: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(2048)), "float16"), + B: T.Tensor((T.int64(2048), T.int64(4096)), "float16"), + matmul: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -971,9 +971,9 @@ def matmul4( @Ts.prim_func(private=True) def maximum1( - A: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), - T_maximum: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), + T_maximum: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -988,9 +988,9 @@ def maximum1( @Ts.prim_func(private=True) def minimum1( - A: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(1), T.int64(256), T.int64(256)), "float16"), - T_minimum: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(1), T.int64(256), T.int64(256)), "float16"), + T_minimum: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1005,8 +1005,8 @@ def minimum1( @Ts.prim_func(private=True) def reshape11( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), - T_reshape: T.Buffer((T.int64(256), T.int64(16), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), + T_reshape: T.Tensor((T.int64(256), T.int64(16), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1031,8 +1031,8 @@ def reshape11( @Ts.prim_func(private=True) def reshape21( - A: T.Buffer((T.int64(256), T.int64(16), T.int64(128)), "float16"), - T_reshape: T.Buffer((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), + A: T.Tensor((T.int64(256), T.int64(16), T.int64(128)), "float16"), + T_reshape: T.Tensor((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1055,8 +1055,8 @@ def reshape21( @Ts.prim_func(private=True) def reshape31( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), - T_reshape: T.Buffer((T.int64(1), T.int64(256), T.int64(2048)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), + T_reshape: T.Tensor((T.int64(1), T.int64(256), T.int64(2048)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1081,8 +1081,8 @@ def reshape31( @Ts.prim_func(private=True) def reshape4( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(2048)), "float16"), - T_reshape: T.Buffer((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(2048)), "float16"), + T_reshape: T.Tensor((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1106,9 +1106,9 @@ def reshape4( @Ts.prim_func def rms_norm( - A: T.Buffer((T.int64(1), 256, T.int64(4096)), "float16"), - B: T.Buffer((T.int64(4096),), "float16"), - rms_norm_1: T.Buffer((T.int64(1), 256, T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), 256, T.int64(4096)), "float16"), + B: T.Tensor((T.int64(4096),), "float16"), + rms_norm_1: T.Tensor((T.int64(1), 256, T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1142,10 +1142,10 @@ def rms_norm( @Ts.prim_func def rotary_embedding( - A: T.Buffer((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), - B: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - C: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - rotary: T.Buffer((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), + B: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + C: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + rotary: T.Tensor((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1168,10 +1168,10 @@ def rotary_embedding( @Ts.prim_func def rotary_embedding1( - A: T.Buffer((T.int64(1), 256, T.int64(16), T.int64(128)), "float16"), - B: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - C: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - rotary: T.Buffer((T.int64(1), 256, T.int64(16), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), 256, T.int64(16), T.int64(128)), "float16"), + B: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + C: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + rotary: T.Tensor((T.int64(1), 256, T.int64(16), T.int64(128)), "float16"), ): T.func_attr({"global_symbol": "rotary_embedding", "tirx.noalias": True}) # with Ts.sblock("root"): @@ -1194,8 +1194,8 @@ def rotary_embedding1( @Ts.prim_func(private=True) def softmax1( - A: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), - T_softmax_norm: T.Buffer( + A: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16"), + T_softmax_norm: T.Tensor( (T.int64(1), T.int64(16), T.int64(256), T.int64(256)), "float16" ), ): @@ -1252,8 +1252,8 @@ def softmax1( @Ts.prim_func(private=True) def transpose11( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), - T_transpose: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), + T_transpose: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1266,8 +1266,8 @@ def transpose11( @Ts.prim_func(private=True) def transpose21( - A: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), - T_transpose: T.Buffer((T.int64(1), T.int64(16), T.int64(128), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), + T_transpose: T.Tensor((T.int64(1), T.int64(16), T.int64(128), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1280,8 +1280,8 @@ def transpose21( @Ts.prim_func(private=True) def transpose31( - A: T.Buffer((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), - T_transpose: T.Buffer((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(16), T.int64(256), T.int64(128)), "float16"), + T_transpose: T.Tensor((T.int64(1), T.int64(256), T.int64(16), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1294,8 +1294,8 @@ def transpose31( @Ts.prim_func(private=True) def transpose4( - A: T.Buffer((T.int64(2048), T.int64(4096)), "float16"), - T_transpose: T.Buffer((T.int64(4096), T.int64(2048)), "float16"), + A: T.Tensor((T.int64(2048), T.int64(4096)), "float16"), + T_transpose: T.Tensor((T.int64(4096), T.int64(2048)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1308,8 +1308,8 @@ def transpose4( @Ts.prim_func(private=True) def transpose5( - A: T.Buffer((T.int64(4096), T.int64(2048)), "float16"), - T_transpose: T.Buffer((T.int64(2048), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(4096), T.int64(2048)), "float16"), + T_transpose: T.Tensor((T.int64(2048), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): 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 64c12da8a66b..58b5e022bfe6 100644 --- a/tests/python/relax/distributed/test_distributed_transform_propagate_sharding.py +++ b/tests/python/relax/distributed/test_distributed_transform_propagate_sharding.py @@ -92,9 +92,9 @@ class MLPWithTuple: @Ts.prim_func(private=True) def split1( - A: T.Buffer((128, 128), "float32"), - T_split: T.Buffer((64, 128), "float32"), - T_split_1: T.Buffer((64, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + T_split: T.Tensor((64, 128), "float32"), + T_split_1: T.Tensor((64, 128), "float32"), ): T.func_attr({"tirx.noalias": True}) @@ -139,9 +139,9 @@ class ShardedMLPWithTuple: @Ts.prim_func(private=True) def split1( - A: T.Buffer((128, 128), "float32"), - T_split: T.Buffer((64, 128), "float32"), - T_split_1: T.Buffer((64, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + T_split: T.Tensor((64, 128), "float32"), + T_split_1: T.Tensor((64, 128), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -375,9 +375,9 @@ class LlamaAttentionLayer: @Ts.prim_func def rms_norm( - A: T.Buffer((T.int64(1), 256, T.int64(4096)), "float16"), - B: T.Buffer((T.int64(4096),), "float16"), - rms_norm_1: T.Buffer((T.int64(1), 256, T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), 256, T.int64(4096)), "float16"), + B: T.Tensor((T.int64(4096),), "float16"), + rms_norm_1: T.Tensor((T.int64(1), 256, T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) @@ -412,10 +412,10 @@ def rms_norm( @Ts.prim_func def rotary_embedding( - A: T.Buffer((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), - B: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - C: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - rotary: T.Buffer((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), + B: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + C: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + rotary: T.Tensor((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) @@ -572,9 +572,9 @@ class ShardedLlamaAttentionLayer: @Ts.prim_func def rms_norm( - A: T.Buffer((T.int64(1), 256, T.int64(4096)), "float16"), - B: T.Buffer((T.int64(4096),), "float16"), - rms_norm_1: T.Buffer((T.int64(1), 256, T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), 256, T.int64(4096)), "float16"), + B: T.Tensor((T.int64(4096),), "float16"), + rms_norm_1: T.Tensor((T.int64(1), 256, T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -608,10 +608,10 @@ def rms_norm( @Ts.prim_func def rotary_embedding( - A: T.Buffer((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), - B: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - C: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - rotary: T.Buffer((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), + B: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + C: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + rotary: T.Tensor((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -776,9 +776,9 @@ class LlamaAttentionLayerTIR: @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), - B: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), - T_add: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + B: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + T_add: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -794,9 +794,9 @@ def add( @Ts.prim_func(private=True) def divide( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - T_divide: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + T_divide: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -812,9 +812,9 @@ def divide( @Ts.prim_func(private=True) def matmul( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), - B: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - matmul: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + B: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + matmul: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -832,9 +832,9 @@ def matmul( @Ts.prim_func(private=True) def matmul1( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), - B: T.Buffer((T.int64(1), T.int64(32), T.int64(128), T.int64(256)), "float16"), - matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + B: T.Tensor((T.int64(1), T.int64(32), T.int64(128), T.int64(256)), "float16"), + matmul: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -855,9 +855,9 @@ def matmul1( @Ts.prim_func(private=True) def matmul2( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), - matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + matmul: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -878,9 +878,9 @@ def matmul2( @Ts.prim_func(private=True) def maximum( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - T_maximum: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + T_maximum: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -896,9 +896,9 @@ def maximum( @Ts.prim_func(private=True) def minimum( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - B: T.Buffer((T.int64(1), T.int64(1), T.int64(256), T.int64(256)), "float16"), - T_minimum: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + B: T.Tensor((T.int64(1), T.int64(1), T.int64(256), T.int64(256)), "float16"), + T_minimum: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -916,8 +916,8 @@ def minimum( @Ts.prim_func(private=True) def reshape( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), - T_reshape: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + T_reshape: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -933,8 +933,8 @@ def reshape( @Ts.prim_func(private=True) def reshape1( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), - T_reshape: T.Buffer((T.int64(256), T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + T_reshape: T.Tensor((T.int64(256), T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -947,8 +947,8 @@ def reshape1( @Ts.prim_func(private=True) def reshape2( - A: T.Buffer((T.int64(256), T.int64(32), T.int64(128)), "float16"), - T_reshape: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(256), T.int64(32), T.int64(128)), "float16"), + T_reshape: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -962,8 +962,8 @@ def reshape2( @Ts.prim_func(private=True) def reshape3( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), - T_reshape: T.Buffer((T.int64(1), T.int64(256), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + T_reshape: T.Tensor((T.int64(1), T.int64(256), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -979,9 +979,9 @@ def reshape3( @Ts.prim_func def rms_norm( - A: T.Buffer((T.int64(1), 256, T.int64(4096)), "float16"), - B: T.Buffer((T.int64(4096),), "float16"), - rms_norm_1: T.Buffer((T.int64(1), 256, T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), 256, T.int64(4096)), "float16"), + B: T.Tensor((T.int64(4096),), "float16"), + rms_norm_1: T.Tensor((T.int64(1), 256, T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1017,10 +1017,10 @@ def rms_norm( @Ts.prim_func def rotary_embedding( - A: T.Buffer((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), - B: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - C: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - rotary: T.Buffer((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), + B: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + C: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + rotary: T.Tensor((T.int64(1), 256, T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1049,8 +1049,8 @@ def rotary_embedding( @Ts.prim_func(private=True) def softmax( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), - T_softmax_norm: T.Buffer( + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16"), + T_softmax_norm: T.Tensor( (T.int64(1), T.int64(32), T.int64(256), T.int64(256)), "float16" ), ): @@ -1116,8 +1116,8 @@ def softmax( @Ts.prim_func(private=True) def transpose( - A: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - T_transpose: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + T_transpose: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1130,8 +1130,8 @@ def transpose( @Ts.prim_func(private=True) def transpose1( - A: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), - T_transpose: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + T_transpose: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1147,8 +1147,8 @@ def transpose1( @Ts.prim_func(private=True) def transpose2( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), - T_transpose: T.Buffer((T.int64(1), T.int64(32), T.int64(128), T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + T_transpose: T.Tensor((T.int64(1), T.int64(32), T.int64(128), T.int64(256)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1164,8 +1164,8 @@ def transpose2( @Ts.prim_func(private=True) def transpose3( - A: T.Buffer((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), - T_transpose: T.Buffer((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), T.int64(32), T.int64(256), T.int64(128)), "float16"), + T_transpose: T.Tensor((T.int64(1), T.int64(256), T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1567,9 +1567,9 @@ class LlamaAttentionLayerDynamicShape: @Ts.prim_func def rms_norm( - A: T.Buffer((T.int64(1), rms_norm_n, T.int64(4096)), "float16"), - B: T.Buffer((T.int64(4096),), "float16"), - rms_norm_1: T.Buffer((T.int64(1), rms_norm_n, T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), rms_norm_n, T.int64(4096)), "float16"), + B: T.Tensor((T.int64(4096),), "float16"), + rms_norm_1: T.Tensor((T.int64(1), rms_norm_n, T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) @@ -1604,11 +1604,11 @@ def rms_norm( @Ts.prim_func def rotary_embedding( - A: T.Buffer((T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)), "float16"), - B: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - C: T.Buffer((T.int64(2048), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)), "float16"), + B: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + C: T.Tensor((T.int64(2048), T.int64(128)), "float16"), m: T.int64, - rotary: T.Buffer( + rotary: T.Tensor( (T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)), "float16" ), ): @@ -1774,9 +1774,9 @@ class ShardedLlamaAttentionLayerDynamicShape: @Ts.prim_func def rms_norm( - A: T.Buffer((T.int64(1), rms_norm_n, T.int64(4096)), "float16"), - B: T.Buffer((T.int64(4096),), "float16"), - rms_norm_1: T.Buffer((T.int64(1), rms_norm_n, T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), rms_norm_n, T.int64(4096)), "float16"), + B: T.Tensor((T.int64(4096),), "float16"), + rms_norm_1: T.Tensor((T.int64(1), rms_norm_n, T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) @@ -1811,11 +1811,11 @@ def rms_norm( @Ts.prim_func def rotary_embedding( - A: T.Buffer((T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)), "float16"), - B: T.Buffer((T.int64(2048), T.int64(128)), "float16"), - C: T.Buffer((T.int64(2048), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)), "float16"), + B: T.Tensor((T.int64(2048), T.int64(128)), "float16"), + C: T.Tensor((T.int64(2048), T.int64(128)), "float16"), m: T.int64, - rotary: T.Buffer( + rotary: T.Tensor( (T.int64(1), rotary_embedding_n, T.int64(32), T.int64(128)), "float16" ), ): diff --git a/tests/python/relax/script/test_relax_script_basic_usage.py b/tests/python/relax/script/test_relax_script_basic_usage.py index df22596adef5..237600f075e9 100644 --- a/tests/python/relax/script/test_relax_script_basic_usage.py +++ b/tests/python/relax/script/test_relax_script_basic_usage.py @@ -81,8 +81,8 @@ def test_simple_module(): class TestModule: @Ts.prim_func(private=True) def tir_func( - x: T.Buffer((T.int64(128), T.int64(128)), "float32"), - y: T.Buffer((T.int64(128), T.int64(128)), "float32"), + x: T.Tensor((T.int64(128), T.int64(128)), "float32"), + y: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j in T.grid(T.int64(128), T.int64(128)): @@ -120,8 +120,8 @@ class TestModule: @Ts.prim_func(private=True) def tir_func( - x: T.Buffer((T.int64(128), T.int64(128)), "float32"), - y: T.Buffer((T.int64(128), T.int64(128)), "float32"), + x: T.Tensor((T.int64(128), T.int64(128)), "float32"), + y: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j in T.grid(T.int64(128), T.int64(128)): @@ -170,8 +170,8 @@ class TestModule: @Ts.prim_func(private=True) def tir_func( - x: T.Buffer((T.int64(128), T.int64(128)), "float32"), - y: T.Buffer((T.int64(128), T.int64(128)), "float32"), + x: T.Tensor((T.int64(128), T.int64(128)), "float32"), + y: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j in T.grid(T.int64(128), T.int64(128)): @@ -616,7 +616,7 @@ def test_call_tir_with_grad(): @I.ir_module class Module: @Ts.prim_func - def identity_tir(A: T.Buffer([54, 96]), B: T.Buffer([54, 96])) -> None: + def identity_tir(A: T.Tensor([54, 96]), B: T.Tensor([54, 96])) -> None: for i, j in T.grid(54, 96): with Ts.sblock("compute"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -642,9 +642,9 @@ def test_call_tir_inplace(): class Module: @Ts.prim_func def copy( - A: T.Buffer((2, 3), "int32"), - B: T.Buffer((2, 3), "int32"), - out1: T.Buffer((2, 3), "int32"), + A: T.Tensor((2, 3), "int32"), + B: T.Tensor((2, 3), "int32"), + out1: T.Tensor((2, 3), "int32"), ): # copies the contents of B into A and out1 T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/relax/script/test_relax_script_distributed.py b/tests/python/relax/script/test_relax_script_distributed.py index 6c67dd11c933..ccb6e0a79916 100644 --- a/tests/python/relax/script/test_relax_script_distributed.py +++ b/tests/python/relax/script/test_relax_script_distributed.py @@ -65,8 +65,8 @@ class TestModule: @Ts.prim_func def tir_func( - x: T.Buffer((T.int64(128), T.int64(128)), "float32"), - y: T.Buffer((T.int64(128), T.int64(128)), "float32"), + x: T.Tensor((T.int64(128), T.int64(128)), "float32"), + y: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j in T.grid(T.int64(128), T.int64(128)): @@ -128,8 +128,8 @@ class TestModule: @Ts.prim_func def tir_func( - x: T.Buffer((T.int64(128), T.int64(128)), "float32"), - y: T.Buffer((T.int64(128), T.int64(128)), "float32"), + x: T.Tensor((T.int64(128), T.int64(128)), "float32"), + y: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j in T.grid(T.int64(128), T.int64(128)): @@ -168,8 +168,8 @@ class TestModule: @Ts.prim_func def tir_func( - x: T.Buffer((T.int64(128), T.int64(128)), "float32"), - y: T.Buffer((T.int64(128), T.int64(128)), "float32"), + x: T.Tensor((T.int64(128), T.int64(128)), "float32"), + y: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j in T.grid(T.int64(128), T.int64(128)): diff --git a/tests/python/relax/script/test_relax_script_dynamic_shape.py b/tests/python/relax/script/test_relax_script_dynamic_shape.py index 1de3b59cbeab..6033d249fbbc 100644 --- a/tests/python/relax/script/test_relax_script_dynamic_shape.py +++ b/tests/python/relax/script/test_relax_script_dynamic_shape.py @@ -216,9 +216,9 @@ def main( @Ts.prim_func def copy( - X: T.Buffer((copy_n * 2,), dtype="float32"), + X: T.Tensor((copy_n * 2,), dtype="float32"), n: copy_n, - Y: T.Buffer((copy_n * 2,), dtype="float32"), + Y: T.Tensor((copy_n * 2,), dtype="float32"), ): for i in T.grid(n * 2): with Ts.sblock("block"): diff --git a/tests/python/relax/script/test_relax_script_error_handling.py b/tests/python/relax/script/test_relax_script_error_handling.py index 961a1a985ba2..f33023dcb671 100644 --- a/tests/python/relax/script/test_relax_script_error_handling.py +++ b/tests/python/relax/script/test_relax_script_error_handling.py @@ -117,7 +117,7 @@ def test_unexpected_tir_args(): @tvm.script.ir_module class TestWellCallTIR: @Ts.prim_func - def tir_addone(A: T.Buffer((16, 16), "int32"), B: T.Buffer((16, 16), "int32")) -> None: + def tir_addone(A: T.Tensor((16, 16), "int32"), B: T.Tensor((16, 16), "int32")) -> None: T.func_attr({"global_symbol": "tir_addone"}) for i, j in T.grid(16, 16): with Ts.sblock("tir_addone"): @@ -332,9 +332,9 @@ def main(x: R.Tensor((2, 3), "int32"), y: R.Tensor((2, 3), "int32")): @Ts.prim_func def copy( - A: T.Buffer((2, 3), "int32"), - B: T.Buffer((2, 3), "int32"), - out1: T.Buffer((2, 3), "int32"), + A: T.Tensor((2, 3), "int32"), + B: T.Tensor((2, 3), "int32"), + out1: T.Tensor((2, 3), "int32"), ): # copies the contents of B into A and out1 T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/relax/script/test_relax_script_meta_programming.py b/tests/python/relax/script/test_relax_script_meta_programming.py index 147eb1791b71..9c8b9f4ad65d 100644 --- a/tests/python/relax/script/test_relax_script_meta_programming.py +++ b/tests/python/relax/script/test_relax_script_meta_programming.py @@ -61,8 +61,8 @@ def test_emit_te_primfunc_attrs(): class TestModule: @Ts.prim_func(private=True) def plus_one( - x: T.Buffer((T.int64(128), T.int64(128)), "float32"), - y: T.Buffer((T.int64(128), T.int64(128)), "float32"), + x: T.Tensor((T.int64(128), T.int64(128)), "float32"), + y: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"some_attr": "foo", "another_attr": True, "tirx.noalias": True}) for i, j in T.grid(T.int64(128), T.int64(128)): @@ -173,7 +173,7 @@ class TestModule: def f(x: R.Tensor((128, 128), "float32"), y: R.Tensor((128, 128), "float32")): @Ts.prim_func def my_matmul( - A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), B: T.Tensor((128, 128)), C: T.Tensor((128, 128)) ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock(): @@ -242,9 +242,9 @@ def test_context_aware_parsing(monkeypatch): class Module: @Ts.prim_func def add( - X: T.Buffer([T.int64(2), T.int64(4)], "float32"), - Y: T.Buffer((), "float32"), - Z: T.Buffer([T.int64(2), T.int64(4)], "float32"), + X: T.Tensor([T.int64(2), T.int64(4)], "float32"), + Y: T.Tensor((), "float32"), + Z: T.Tensor([T.int64(2), T.int64(4)], "float32"), ): T.evaluate(0) diff --git a/tests/python/relax/script/test_relax_script_printer.py b/tests/python/relax/script/test_relax_script_printer.py index 4a8ab35f9584..1d3a507127b5 100644 --- a/tests/python/relax/script/test_relax_script_printer.py +++ b/tests/python/relax/script/test_relax_script_printer.py @@ -89,8 +89,8 @@ class TestModule: @Ts.prim_func def tir_func( - x: T.Buffer((T.int64(128), T.int64(128)), "float32"), - y: T.Buffer((T.int64(128), T.int64(128)), "float32"), + x: T.Tensor((T.int64(128), T.int64(128)), "float32"), + y: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j in T.grid(T.int64(128), T.int64(128)): @@ -144,7 +144,7 @@ class Module: I.module_attrs({"device_num": 10}) I.module_global_infos({"mesh": [R.device_mesh((2, 2), I.Range(0, 4)), R.device_mesh((1,), I.Range(4, 5))]}) @Ts.prim_func - def tir_func(x: T.Buffer((T.int64(128), T.int64(128)), "float32"), y: T.Buffer((T.int64(128), T.int64(128)), "float32")): + def tir_func(x: T.Tensor((T.int64(128), T.int64(128)), "float32"), y: T.Tensor((T.int64(128), T.int64(128)), "float32")): T.func_attr({"tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -869,7 +869,7 @@ def test_module_cross_func_call(): class TestModule: @Ts.prim_func def tir_func( - x: T.Buffer((T.int64(128),), "float32"), y: T.Buffer((T.int64(128),), "float32") + x: T.Tensor((T.int64(128),), "float32"), y: T.Tensor((T.int64(128),), "float32") ): T.evaluate(0) @@ -893,7 +893,7 @@ def foo(x: R.Tensor((128,), "float32")) -> R.Tensor((128,), "float32"): @I.ir_module class Module: @Ts.prim_func - def tir_func(x: T.Buffer((T.int64(128),), "float32"), y: T.Buffer((T.int64(128),), "float32")): + def tir_func(x: T.Tensor((T.int64(128),), "float32"), y: T.Tensor((T.int64(128),), "float32")): T.evaluate(0) @R.function @@ -918,7 +918,7 @@ def foo(x: R.Tensor((128,), dtype="float32")) -> R.Tensor((128,), dtype="float32 @I.ir_module class Module: @Ts.prim_func - def tir_func(x: T.Buffer((T.int64(128),), "float32"), y: T.Buffer((T.int64(128),), "float32")): + def tir_func(x: T.Tensor((T.int64(128),), "float32"), y: T.Tensor((T.int64(128),), "float32")): T.evaluate(0) @R.function diff --git a/tests/python/relax/script/test_relax_script_pyfunc.py b/tests/python/relax/script/test_relax_script_pyfunc.py index 54cb5199eb1a..7b691f8be754 100644 --- a/tests/python/relax/script/test_relax_script_pyfunc.py +++ b/tests/python/relax/script/test_relax_script_pyfunc.py @@ -53,8 +53,8 @@ def pytorch_complex_ops(x: torch.Tensor) -> torch.Tensor: @Ts.prim_func def simple_tir_func( - A: T.Buffer((n,), "float32"), - B: T.Buffer((n,), "float32"), + A: T.Tensor((n,), "float32"), + B: T.Tensor((n,), "float32"), ): T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/relax/test_analysis.py b/tests/python/relax/test_analysis.py index d66f0c617652..d4986d6fc893 100644 --- a/tests/python/relax/test_analysis.py +++ b/tests/python/relax/test_analysis.py @@ -557,8 +557,8 @@ def test_all_global_vars(): def test_reshape_pattern_reshape(): @Ts.prim_func def reshape( - rxplaceholder: T.Buffer((1, 2, 3, 4), "float32"), - T_reshape: T.Buffer((8, 3), "float32"), + rxplaceholder: T.Tensor((1, 2, 3, 4), "float32"), + T_reshape: T.Tensor((8, 3), "float32"), ): for i0, i1 in T.grid(8, 3): with Ts.sblock("T_reshape"): @@ -585,8 +585,8 @@ def reshape( def test_reshape_pattern_reshape_scheduled(): @Ts.prim_func def reshape_scheduled( - rxplaceholder: T.Buffer((1, 2, 3, 4), "float32"), - T_reshape: T.Buffer((8, 3), "float32"), + rxplaceholder: T.Tensor((1, 2, 3, 4), "float32"), + T_reshape: T.Tensor((8, 3), "float32"), ): for i0_i1_fused_0 in T.thread_binding(1, thread="blockIdx.x"): for i0_i1_fused_1 in T.thread_binding(24, thread="threadIdx.x"): @@ -615,8 +615,8 @@ def reshape_scheduled( def test_reshape_pattern_zero_extent(): @Ts.prim_func def transpose_zero( - rxplaceholder: T.Buffer((3, 0, 4), "float32"), - T_transpose: T.Buffer((0, 3, 4), "float32"), + rxplaceholder: T.Tensor((3, 0, 4), "float32"), + T_transpose: T.Tensor((0, 3, 4), "float32"), ): for i0, i1, i2 in T.grid(0, 3, 4): with Ts.sblock("T_transpose"): @@ -631,8 +631,8 @@ def transpose_zero( def test_reshape_pattern_expand_dims(): @Ts.prim_func def expand_dims( - rxplaceholder: T.Buffer((2, 3, 4), "float32"), - expand_dims: T.Buffer((2, 1, 1, 1, 3, 1, 4, 1), "float32"), + rxplaceholder: T.Tensor((2, 3, 4), "float32"), + expand_dims: T.Tensor((2, 1, 1, 1, 3, 1, 4, 1), "float32"), ): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3, i4, i5, i6, i7 in T.grid(2, 1, 1, 1, 3, 1, 4, 1): @@ -654,8 +654,8 @@ def test_reshape_pattern_dyn_1(): @Ts.prim_func def reshape( - A: T.Buffer((n, T.int64(32), T.int64(128)), "float16"), - T_reshape: T.Buffer((T.int64(1), n, T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((n, T.int64(32), T.int64(128)), "float16"), + T_reshape: T.Tensor((T.int64(1), n, T.int64(32), T.int64(128)), "float16"), ): for ax0, ax1, ax2, ax3 in T.grid(T.int64(1), n, T.int64(32), T.int64(128)): with Ts.sblock("T_reshape"): @@ -681,7 +681,7 @@ def test_reshape_pattern_dyn_2(): n = T.dynamic("n") @Ts.prim_func - def reshape(A: T.Buffer((T.int64(1), n), "int32"), T_reshape: T.Buffer((n,), "int32")): + def reshape(A: T.Tensor((T.int64(1), n), "int32"), T_reshape: T.Tensor((n,), "int32")): for ax0 in range(n): with Ts.sblock("T_reshape"): v_ax0 = Ts.axis.spatial(n, ax0) @@ -697,8 +697,8 @@ def test_reshape_pattern_dyn_3(): @Ts.prim_func def reshape( - A: T.Buffer((n, T.int64(4096)), "float16"), - T_reshape: T.Buffer((T.int64(1), n, T.int64(4096)), "float16"), + A: T.Tensor((n, T.int64(4096)), "float16"), + T_reshape: T.Tensor((T.int64(1), n, T.int64(4096)), "float16"), ): T.func_attr({"op_pattern": 8, "tirx.noalias": True}) @@ -719,8 +719,8 @@ def test_reshape_pattern_dyn_4(): @Ts.prim_func def reshape( - A: T.Buffer((T.int64(1), n, T.int64(4096)), "float16"), - T_reshape: T.Buffer((T.int64(1), n, T.int64(32), T.int64(128)), "float16"), + A: T.Tensor((T.int64(1), n, T.int64(4096)), "float16"), + T_reshape: T.Tensor((T.int64(1), n, T.int64(32), T.int64(128)), "float16"), ): T.func_attr({"op_pattern": 8, "tirx.noalias": True}) @@ -749,8 +749,8 @@ def test_reshape_pattern_dyn_5(): @Ts.prim_func def reshape( - A: T.Buffer((T.int64(1), n, T.int64(32), T.int64(128)), "float16"), - T_reshape: T.Buffer((T.int64(1), n, T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), n, T.int64(32), T.int64(128)), "float16"), + T_reshape: T.Tensor((T.int64(1), n, T.int64(4096)), "float16"), ): T.func_attr({"op_pattern": 8, "tirx.noalias": True}) @@ -780,9 +780,9 @@ def reshape( def test_reshape_pattern_with_raggedness(): @Ts.prim_func def reshape_raggedness( - A: T.Buffer((100, 768), "float32"), - src_indptr: T.Buffer((9,), "int32"), - B: T.Buffer((100, 12, 64), "float32"), + A: T.Tensor((100, 768), "float32"), + src_indptr: T.Tensor((9,), "int32"), + B: T.Tensor((100, 12, 64), "float32"), ): for b in T.serial(8): with Ts.sblock("block0"): @@ -801,7 +801,7 @@ def reshape_raggedness( def test_reshape_pattern_reject_seqstmt(): @Ts.prim_func - def identity_bias(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4, 4), "float32")): + def identity_bias(A: T.Tensor((4, 4), "float32"), B: T.Tensor((4, 4), "float32")): C = Ts.sblock_alloc_buffer((128, 128), "float32") for i0, i1 in T.grid(4, 4): with Ts.sblock("identity"): @@ -813,7 +813,7 @@ def identity_bias(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4, 4), "float32") B[vi0, vi1] = C[vi0, vi1] + T.float32(1) @Ts.prim_func - def identity_identity(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4, 4), "float32")): + def identity_identity(A: T.Tensor((4, 4), "float32"), B: T.Tensor((4, 4), "float32")): C = Ts.sblock_alloc_buffer((128, 128), "float32") for i0, i1 in T.grid(4, 4): with Ts.sblock("identity"): @@ -830,7 +830,7 @@ def identity_identity(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4, 4), "float def test_reshape_pattern_reject_reduction(): @Ts.prim_func - def reduction(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4,), "float32")): + def reduction(A: T.Tensor((4, 4), "float32"), B: T.Tensor((4,), "float32")): for i0, i1 in T.grid(4, 4): with Ts.sblock("identity"): vi0, vi1 = Ts.axis.remap("SR", [i0, i1]) @@ -843,7 +843,7 @@ def reduction(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4,), "float32")): def test_reshape_pattern_reject_reduction(): @Ts.prim_func - def reduction(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4,), "float32")): + def reduction(A: T.Tensor((4, 4), "float32"), B: T.Tensor((4,), "float32")): for i0, i1 in T.grid(4, 4): with Ts.sblock("identity"): vi0, vi1 = Ts.axis.remap("SR", [i0, i1]) diff --git a/tests/python/relax/test_analysis_detect_recursion.py b/tests/python/relax/test_analysis_detect_recursion.py index 3d1440f64a0a..f0a0be180f08 100644 --- a/tests/python/relax/test_analysis_detect_recursion.py +++ b/tests/python/relax/test_analysis_detect_recursion.py @@ -422,7 +422,7 @@ def test_disregard_primfuncs(): class CallPrimFunc: # copied from test_analysis.py @Ts.prim_func - def identity_identity(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4, 4), "float32")): + def identity_identity(A: T.Tensor((4, 4), "float32"), B: T.Tensor((4, 4), "float32")): C = Ts.sblock_alloc_buffer((128, 128), "float32") for i0, i1 in T.grid(4, 4): with Ts.sblock("identity"): diff --git a/tests/python/relax/test_analysis_estimate_memory_usage.py b/tests/python/relax/test_analysis_estimate_memory_usage.py index bdcb379e0132..802fae8495c2 100644 --- a/tests/python/relax/test_analysis_estimate_memory_usage.py +++ b/tests/python/relax/test_analysis_estimate_memory_usage.py @@ -29,43 +29,43 @@ def test_basic(): class Module: @Ts.prim_func def add( - rxplaceholder: T.Buffer(T.int64(8), "float32"), - rxplaceholder_1: T.Buffer((), "float32"), - T_add: T.Buffer(T.int64(8), "float32"), + rxplaceholder: T.Tensor(T.int64(8), "float32"), + rxplaceholder_1: T.Tensor((), "float32"), + T_add: T.Tensor(T.int64(8), "float32"), ): T.evaluate(0) @Ts.prim_func def reshape( - rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), - T_reshape: T.Buffer(T.int64(8), "float32"), + rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), + T_reshape: T.Tensor(T.int64(8), "float32"), ): T.evaluate(0) @Ts.prim_func def relu( - rxplaceholder: T.Buffer(T.int64(8), "float32"), compute: T.Buffer(T.int64(8), "float32") + rxplaceholder: T.Tensor(T.int64(8), "float32"), compute: T.Tensor(T.int64(8), "float32") ): T.evaluate(0) @Ts.prim_func def log( - rxplaceholder: T.Buffer(T.int64(10), "float32"), - compute: T.Buffer(T.int64(10), "float32"), + rxplaceholder: T.Tensor(T.int64(10), "float32"), + compute: T.Tensor(T.int64(10), "float32"), ): T.evaluate(0) @Ts.prim_func def exp( - rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), - compute: T.Buffer((T.int64(2), T.int64(4)), "float32"), + rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), + compute: T.Tensor((T.int64(2), T.int64(4)), "float32"), ): T.evaluate(0) @Ts.prim_func def pad( - rxplaceholder: T.Buffer(T.int64(8), "float32"), - PadInput: T.Buffer(T.int64(10), "float32"), + rxplaceholder: T.Tensor(T.int64(8), "float32"), + PadInput: T.Tensor(T.int64(10), "float32"), ): T.evaluate(0) diff --git a/tests/python/relax/test_analysis_suggest_layout_transforms.py b/tests/python/relax/test_analysis_suggest_layout_transforms.py index f33bebb932e0..64a0fa4dfd2e 100644 --- a/tests/python/relax/test_analysis_suggest_layout_transforms.py +++ b/tests/python/relax/test_analysis_suggest_layout_transforms.py @@ -46,8 +46,8 @@ def apply_transformations(func, suggested_transfoms, print_transformation=False) def test_nested_blocks(): @Ts.prim_func(private=True) def nested_block( - arg: T.Buffer((32, 64, 224, 224), "float32"), - relu: T.Buffer((32, 64, 224, 224), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + relu: T.Tensor((32, 64, 224, 224), "float32"), ): for i, j in T.grid(32, 64): with Ts.sblock("outer"): @@ -71,8 +71,8 @@ def nested_block( def test_mismatch_transformations_and_num_params(): @Ts.prim_func(private=True) def elemwise( - arg: T.Buffer((32, 64, 224, 224), "float32"), - relu: T.Buffer((32, 64, 224, 224), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + relu: T.Tensor((32, 64, 224, 224), "float32"), ): for i0, i1, i2, i3 in T.grid(32, 64, 224, 224): with Ts.sblock("compute"): @@ -95,8 +95,8 @@ def elemwise( def test_empty_write_transformations(): @Ts.prim_func(private=True) def elemwise( - arg: T.Buffer((32, 64, 224, 224), "float32"), - relu: T.Buffer((32, 64, 224, 224), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + relu: T.Tensor((32, 64, 224, 224), "float32"), ): for i0, i1, i2, i3 in T.grid(32, 64, 224, 224): with Ts.sblock("compute"): @@ -114,8 +114,8 @@ def elemwise( def test_non_bijective_block_transform(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 64), "float32"), - output: T.Buffer((32, 64), "float32"), + arg: T.Tensor((32, 64), "float32"), + output: T.Tensor((32, 64), "float32"), ): for ax0, ax1 in T.grid(32, 64): with Ts.sblock("compute"): @@ -133,8 +133,8 @@ def before( def test_non_affine_access(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 64), "float32"), - output: T.Buffer((32 * 64, 10), "float32"), + arg: T.Tensor((32, 64), "float32"), + output: T.Tensor((32 * 64, 10), "float32"), ): for ax0, ax1, ax2 in T.grid(32, 64, 10): with Ts.sblock("compute"): @@ -152,8 +152,8 @@ def before( def test_unsupported_write_spatial_layout(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((4, 4), "float32"), - output: T.Buffer((16), "float32"), + arg: T.Tensor((4, 4), "float32"), + output: T.Tensor((16), "float32"), ): for ax0, ax1 in T.grid(4, 4): with Ts.sblock("flatten"): @@ -171,8 +171,8 @@ def before( def test_unpacked_iter_used_in_read_access(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((8, 4), "float32"), - output: T.Buffer((4, 8), "float32"), + arg: T.Tensor((8, 4), "float32"), + output: T.Tensor((4, 8), "float32"), ): for ax0, ax1, ax2 in T.grid(4, 8, 4): with Ts.sblock("compute"): @@ -183,8 +183,8 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((8, 4), "float32"), - output: T.Buffer((32), "float32"), + arg: T.Tensor((8, 4), "float32"), + output: T.Tensor((32), "float32"), ): for ax0, ax2 in T.grid(32, 4): with Ts.sblock("compute"): @@ -203,8 +203,8 @@ def expected( def test_invalid_index_map(): @Ts.prim_func(private=True) def elemwise( - arg: T.Buffer((32, 64, 224, 224), "float32"), - relu: T.Buffer((32, 64, 224, 224), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + relu: T.Tensor((32, 64, 224, 224), "float32"), ): for i0, i1, i2, i3 in T.grid(32, 64, 224, 224): with Ts.sblock("compute"): @@ -224,8 +224,8 @@ def elemwise( def test_SRSR_block(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 224, 64, 224), "float32"), - sum: T.Buffer((32, 64), "float32"), + arg: T.Tensor((32, 224, 64, 224), "float32"), + sum: T.Tensor((32, 64), "float32"), ): for ax0, k2, ax1, k3 in T.grid(32, 224, 64, 224): with Ts.sblock("rxplaceholder_red"): @@ -238,8 +238,8 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((32, 224, 16, 224, 4), "float32"), - sum: T.Buffer((32, 16, 4), "float32"), + arg: T.Tensor((32, 224, 16, 224, 4), "float32"), + sum: T.Tensor((32, 16, 4), "float32"), ): for ax0, ax1, ax2, ax3, ax4 in T.grid(32, 224, 16, 224, 4): with Ts.sblock("rxplaceholder_red"): @@ -264,7 +264,7 @@ def test_op_elemwise_symbolic(): W = T.dynamic("W") @Ts.prim_func(private=True) - def before(Arg: T.Buffer((N, C, H, W)), Relu: T.Buffer((N, C, H, W))): + def before(Arg: T.Tensor((N, C, H, W)), Relu: T.Tensor((N, C, H, W))): for i0, i1, i2, i3 in T.grid(N, C, H, W): with Ts.sblock("compute"): v_i0, v_i1, v_i2, v_i3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) @@ -278,7 +278,7 @@ def before(Arg: T.Buffer((N, C, H, W)), Relu: T.Buffer((N, C, H, W))): W = T.dynamic("W") @Ts.prim_func(private=True) - def expected(Arg: T.Buffer((N, H, W, C)), Relu: T.Buffer((N, H, W, C))): + def expected(Arg: T.Tensor((N, H, W, C)), Relu: T.Tensor((N, H, W, C))): # with Ts.sblock("root"): for ax0, ax1, ax2, ax3 in T.grid(N, H, W, C): with Ts.sblock("compute"): @@ -297,8 +297,8 @@ def expected(Arg: T.Buffer((N, H, W, C)), Relu: T.Buffer((N, H, W, C))): def test_op_elemwise(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 64, 224, 224), "float32"), - relu: T.Buffer((32, 64, 224, 224), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + relu: T.Tensor((32, 64, 224, 224), "float32"), ): for i0, i1, i2, i3 in T.grid(32, 64, 224, 224): with Ts.sblock("compute"): @@ -309,8 +309,8 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((32, 224, 224, 64), "float32"), - relu: T.Buffer((32, 224, 224, 64), "float32"), + arg: T.Tensor((32, 224, 224, 64), "float32"), + relu: T.Tensor((32, 224, 224, 64), "float32"), ): for ax0, ax1, ax2, ax3 in T.grid(32, 224, 224, 64): with Ts.sblock("compute"): @@ -329,8 +329,8 @@ def expected( def test_op_pool_nchw_nhwc(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 64, 224, 224), "float32"), - pool_max: T.Buffer((32, 64, 111, 223), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + pool_max: T.Tensor((32, 64, 111, 223), "float32"), ): for ax0, ax1, ax2, ax3, rv0, rv1 in T.grid(32, 64, 111, 223, 2, 2): with Ts.sblock("pool_max"): @@ -361,8 +361,8 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((32, 224, 224, 64), "float32"), - pool_max: T.Buffer((32, 111, 223, 64), "float32"), + arg: T.Tensor((32, 224, 224, 64), "float32"), + pool_max: T.Tensor((32, 111, 223, 64), "float32"), ): # with Ts.sblock("root"): for ax0, ax1, ax2, ax3, ax4, ax5 in T.grid(32, 111, 223, 64, 2, 2): @@ -389,11 +389,11 @@ def expected( def test_op_pool_nchw16c_nhwc(): @Ts.prim_func(private=True) def before( - arg: T.Buffer( + arg: T.Tensor( (32, 4, 224, 224, 16), "float32", ), - pool_max: T.Buffer( + pool_max: T.Tensor( (32, 4, 110, 220, 16), "float32", ), @@ -415,8 +415,8 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((32, 224, 224, 64), "float32"), - pool_max: T.Buffer((32, 110, 220, 64), "float32"), + arg: T.Tensor((32, 224, 224, 64), "float32"), + pool_max: T.Tensor((32, 110, 220, 64), "float32"), ): for ax0, ax1, ax2, ax3, ax4, ax5 in T.grid(32, 110, 220, 64, 5, 5): with Ts.sblock("pool_max"): @@ -442,8 +442,8 @@ def expected( def test_op_reduce(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 64, 224, 224), "float32"), - sum: T.Buffer((32, 64), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + sum: T.Tensor((32, 64), "float32"), ): for ax0, ax1, k2, k3 in T.grid(32, 64, 224, 224): with Ts.sblock("rxplaceholder_red"): @@ -456,8 +456,8 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((32, 4, 224, 224, 16), "float32"), - sum: T.Buffer((32, 4, 16), "float32"), + arg: T.Tensor((32, 4, 224, 224, 16), "float32"), + sum: T.Tensor((32, 4, 16), "float32"), ): for ax0, ax1, ax2, ax3, ax4 in T.grid(32, 4, 224, 224, 16): with Ts.sblock("rxplaceholder_red"): @@ -479,8 +479,8 @@ def test_op_upsampling(): # relax materializes the layout if H, W or D dimensions are moved or tiled. @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 64, 224, 224), "float32"), - resize: T.Buffer((32, 64, 202, 246), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + resize: T.Tensor((32, 64, 202, 246), "float32"), ): for i0, i1, i2, i3 in T.grid(32, 64, 202, 246): with Ts.sblock("resize"): @@ -520,8 +520,8 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((32, 64, 224, 224), "float32"), - resize: T.Buffer((32, 202, 246, 64), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + resize: T.Tensor((32, 202, 246, 64), "float32"), ): # with Ts.sblock("root"): for ax0, ax1, ax2, ax3 in T.grid(32, 202, 246, 64): @@ -570,8 +570,8 @@ def expected( def test_op_strided_slice(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 64, 224, 224), "float32"), - T_strided_slice_with_axes: T.Buffer((32, 64, 10, 8), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + T_strided_slice_with_axes: T.Tensor((32, 64, 10, 8), "float32"), ): for ax0, ax1, ax2, ax3 in T.grid(32, 64, 10, 8): with Ts.sblock("T_strided_slice_with_axes"): @@ -594,8 +594,8 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((32, 224, 224, 16, 4), "float32"), - T_strided_slice_with_axes: T.Buffer((32, 10, 8, 16, 4), "float32"), + arg: T.Tensor((32, 224, 224, 16, 4), "float32"), + T_strided_slice_with_axes: T.Tensor((32, 10, 8, 16, 4), "float32"), ): # with Ts.sblock("root"): for ax0, ax1, ax2, ax3, ax4 in T.grid(32, 10, 8, 16, 4): @@ -617,9 +617,9 @@ def expected( def test_op_binary_broadcast(): @Ts.prim_func(private=True) def before( - arg0: T.Buffer((32, 64, 224, 224), "float32"), - arg1: T.Buffer((64, 224, 224), "float32"), - T_add: T.Buffer((32, 64, 224, 224), "float32"), + arg0: T.Tensor((32, 64, 224, 224), "float32"), + arg1: T.Tensor((64, 224, 224), "float32"), + T_add: T.Tensor((32, 64, 224, 224), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -637,9 +637,9 @@ def before( @Ts.prim_func(private=True) def expected( - arg0: T.Buffer((32, 224, 224, 16, 4), "float32"), - arg1: T.Buffer((224, 224, 16, 4), "float32"), - T_add: T.Buffer((32, 224, 224, 16, 4), "float32"), + arg0: T.Tensor((32, 224, 224, 16, 4), "float32"), + arg1: T.Tensor((224, 224, 16, 4), "float32"), + T_add: T.Tensor((32, 224, 224, 16, 4), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -660,8 +660,8 @@ def expected( def test_op_transpose(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 64, 224, 224), "float32"), - T_transpose: T.Buffer((32, 224, 224, 64), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + T_transpose: T.Tensor((32, 224, 224, 64), "float32"), ): for ax0, ax1, ax2, ax3 in T.grid(32, 224, 224, 64): with Ts.sblock("T_transpose"): @@ -672,8 +672,8 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((32, 64, 224, 224), "float32"), - T_transpose: T.Buffer((32, 224, 64, 224), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + T_transpose: T.Tensor((32, 224, 64, 224), "float32"), ): for ax0, ax1, ax2, ax3 in T.grid(32, 224, 64, 224): with Ts.sblock("T_transpose"): @@ -692,8 +692,8 @@ def expected( def test_op_pad(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 64, 224, 224), "float32"), - PadInput: T.Buffer((32, 64, 230, 230), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + PadInput: T.Tensor((32, 64, 230, 230), "float32"), ): for i0, i1, i2, i3 in T.grid(32, 64, 230, 230): with Ts.sblock("PadInput"): @@ -708,8 +708,8 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((32, 224, 224, 16, 4), "float32"), - PadInput: T.Buffer((32, 230, 230, 16, 4), "float32"), + arg: T.Tensor((32, 224, 224, 16, 4), "float32"), + PadInput: T.Tensor((32, 230, 230, 16, 4), "float32"), ): for ax0, ax1, ax2, ax3, ax4 in T.grid(32, 230, 230, 16, 4): with Ts.sblock("PadInput"): @@ -732,9 +732,9 @@ def expected( def test_op_split(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 64, 224, 224), "float32"), - split0: T.Buffer((32, 32, 224, 224), "float32"), - split1: T.Buffer((32, 32, 224, 224), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + split0: T.Tensor((32, 32, 224, 224), "float32"), + split1: T.Tensor((32, 32, 224, 224), "float32"), ): for ax0, ax1, ax2, ax3 in T.grid(32, 32, 224, 224): with Ts.sblock("T_split_sections"): @@ -751,9 +751,9 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((32, 224, 224, 64), "float32"), - split0: T.Buffer((32, 224, 224, 32), "float32"), - split1: T.Buffer((32, 224, 224, 32), "float32"), + arg: T.Tensor((32, 224, 224, 64), "float32"), + split0: T.Tensor((32, 224, 224, 32), "float32"), + split1: T.Tensor((32, 224, 224, 32), "float32"), ): for ax0, ax1, ax2, ax3 in T.grid(32, 224, 224, 32): with Ts.sblock("T_split_sections"): @@ -780,9 +780,9 @@ def expected( def test_op_split_tiling_split_dim(): @Ts.prim_func(private=True) def before( - arg: T.Buffer((32, 64, 224, 224), "float32"), - split0: T.Buffer((32, 32, 224, 224), "float32"), - split1: T.Buffer((32, 32, 224, 224), "float32"), + arg: T.Tensor((32, 64, 224, 224), "float32"), + split0: T.Tensor((32, 32, 224, 224), "float32"), + split1: T.Tensor((32, 32, 224, 224), "float32"), ): for ax0, ax1, ax2, ax3 in T.grid(32, 32, 224, 224): with Ts.sblock("T_split_sections"): @@ -799,9 +799,9 @@ def before( @Ts.prim_func(private=True) def expected( - arg: T.Buffer((32, 224, 224, 16, 4), "float32"), - split0: T.Buffer((32, 224, 224, 8, 4), "float32"), - split1: T.Buffer((32, 224, 224, 8, 4), "float32"), + arg: T.Tensor((32, 224, 224, 16, 4), "float32"), + split0: T.Tensor((32, 224, 224, 8, 4), "float32"), + split1: T.Tensor((32, 224, 224, 8, 4), "float32"), ): # with Ts.sblock("root"): for ax0, ax1, ax2, ax3, ax4 in T.grid(32, 224, 224, 8, 4): diff --git a/tests/python/relax/test_analysis_well_formed.py b/tests/python/relax/test_analysis_well_formed.py index 4e9e8d829ef7..84c190557230 100644 --- a/tests/python/relax/test_analysis_well_formed.py +++ b/tests/python/relax/test_analysis_well_formed.py @@ -761,7 +761,7 @@ def main(A: R.Tensor([16], "float16")): return B @Ts.prim_func - def add_one(A: T.Buffer(16, "float16"), B: T.Buffer(16, "float16")): + def add_one(A: T.Tensor(16, "float16"), B: T.Tensor(16, "float16")): for i in range(16): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -786,10 +786,10 @@ def main( @Ts.prim_func def add_scaled( - A: T.Buffer([T.int64(16)], "float16"), + A: T.Tensor([T.int64(16)], "float16"), scale: T.float32, - C: T.Buffer([T.int64(16)], "float16"), - B: T.Buffer([T.int64(16)], "float16"), + C: T.Tensor([T.int64(16)], "float16"), + B: T.Tensor([T.int64(16)], "float16"), ): for i in range(T.int64(16)): B[i] = A[i] + T.Cast("float16", scale) * C[i] @@ -809,9 +809,9 @@ def main(A: R.Tensor([16], "float16"), scale: T.int64): @Ts.prim_func def scale( - A: T.Buffer([T.int64(16)], "float16"), + A: T.Tensor([T.int64(16)], "float16"), scale: T.float32, - B: T.Buffer([T.int64(16)], "float16"), + B: T.Tensor([T.int64(16)], "float16"), ): for i in range(T.int64(16)): B[i] = A[i] * T.Cast("float16", scale) @@ -836,7 +836,7 @@ def main(): return B @Ts.prim_func - def make_tensor(m: T.int64, n: T.int64, B: T.Buffer([T.int64(1)], "float32")): + def make_tensor(m: T.int64, n: T.int64, B: T.Tensor([T.int64(1)], "float32")): B[0] = T.Cast("float32", m + n) assert not rx.analysis.check_well_formed(Module) @@ -858,7 +858,7 @@ def main(A: R.Tensor([4, 4], "float16")): return B @Ts.prim_func - def add_one(A: T.Buffer(16, "float16"), B: T.Buffer(16, "float16")): + def add_one(A: T.Tensor(16, "float16"), B: T.Tensor(16, "float16")): for i in range(16): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -882,7 +882,7 @@ def main(A: R.Tensor([16], "float16")): return B @Ts.prim_func - def add_one(A: T.Buffer(16, "float16"), B: T.Buffer(16, "float16")): + def add_one(A: T.Tensor(16, "float16"), B: T.Tensor(16, "float16")): for i in range(16): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -907,7 +907,7 @@ def main(A: R.Tensor([32], "float16")): return B @Ts.prim_func - def add_one(A: T.Buffer(16, "float16"), B: T.Buffer(16, "float16")): + def add_one(A: T.Tensor(16, "float16"), B: T.Tensor(16, "float16")): for i in range(16): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -931,7 +931,7 @@ def main(A: R.Tensor([16], "float16")): return B @Ts.prim_func - def add_one(A: T.Buffer(16, "float16"), B: T.Buffer(16, "float16")): + def add_one(A: T.Tensor(16, "float16"), B: T.Tensor(16, "float16")): for i in range(16): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -957,7 +957,7 @@ def main(A: R.Tensor([16], "float32")): return B @Ts.prim_func - def add_one(A: T.Buffer(16, "float16"), B: T.Buffer(16, "float16")): + def add_one(A: T.Tensor(16, "float16"), B: T.Tensor(16, "float16")): for i in range(16): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -983,7 +983,7 @@ def main(A: R.Tensor([16], "float16")): return B @Ts.prim_func - def add_one(A: T.Buffer(16, "float16"), B: T.Buffer(16, "float16")): + def add_one(A: T.Tensor(16, "float16"), B: T.Tensor(16, "float16")): for i in range(16): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -1015,7 +1015,7 @@ def main(A: R.Tensor([16], "float16")): return B @Ts.prim_func - def reshape(A: T.Buffer(16, "float16"), B: T.Buffer([M, N], dtype="float16")): + def reshape(A: T.Tensor(16, "float16"), B: T.Tensor([M, N], dtype="float16")): for i, j in T.grid(M, N): with Ts.sblock("compute"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -1047,7 +1047,7 @@ def main(A: R.Tensor([16], "float16")): return B @Ts.prim_func - def reshape(A: T.Buffer(16, "float16"), B: T.Buffer([M, N], dtype="float16")): + def reshape(A: T.Tensor(16, "float16"), B: T.Tensor([M, N], dtype="float16")): for i, j in T.grid(M, N): with Ts.sblock("compute"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -1081,7 +1081,7 @@ def main(A: R.Tensor([16], "float16")): return B @Ts.prim_func - def reshape(A: T.Buffer(16, "float16"), B: T.Buffer([M, N], dtype="float16")): + def reshape(A: T.Tensor(16, "float16"), B: T.Tensor([M, N], dtype="float16")): for i, j in T.grid(M, N): with Ts.sblock("compute"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -1118,7 +1118,7 @@ def main(A: R.Tensor([256], "float16")): return B @Ts.prim_func - def reshape(A: T.Buffer(256, "float16"), B: T.Buffer([16, M, N], dtype="float16")): + def reshape(A: T.Tensor(256, "float16"), B: T.Tensor([16, M, N], dtype="float16")): for i, j, k in T.grid(16, M, N): with Ts.sblock("compute"): vi, vj, vk = Ts.axis.remap("SSS", [i, j, k]) @@ -1149,7 +1149,7 @@ def main(A: R.Tensor([8, 4], "float16")): return B @Ts.prim_func - def flatten(A: T.Buffer([M, N], dtype="float16"), B: T.Buffer([M * N], dtype="float16")): + def flatten(A: T.Tensor([M, N], dtype="float16"), B: T.Tensor([M * N], dtype="float16")): for i in T.grid(M * N): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -1185,7 +1185,7 @@ def main(A: R.Tensor([8, 4], "float16")): return B @Ts.prim_func - def flatten(A: T.Buffer([M, N], dtype="float16"), B: T.Buffer([M * N], dtype="float16")): + def flatten(A: T.Tensor([M, N], dtype="float16"), B: T.Tensor([M * N], dtype="float16")): for i in T.grid(M * N): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -1222,7 +1222,7 @@ def main(A: R.dist.DTensor([8, 4], "float16", "mesh[0]", "S[0]")): return B @Ts.prim_func - def flatten(A: T.Buffer([M, N], dtype="float16"), B: T.Buffer([M * N], dtype="float16")): + def flatten(A: T.Tensor([M, N], dtype="float16"), B: T.Tensor([M * N], dtype="float16")): for i in T.grid(M * N): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -1247,7 +1247,7 @@ def main(A: R.Tensor([16], "float16")): return B @Ts.prim_func - def add_one(A: T.Buffer(16, "float16")): + def add_one(A: T.Tensor(16, "float16")): for i in range(16): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -1272,7 +1272,7 @@ def main(A: R.Tensor([16], "float16")): return B @Ts.prim_func - def add_one(A: T.Buffer(16, "float16")): + def add_one(A: T.Tensor(16, "float16")): for i in range(16): with Ts.sblock("compute"): vi = Ts.axis.remap("S", [i]) @@ -1301,9 +1301,9 @@ def main(A: R.Tensor([16], "float16"), B: R.Tensor([32], "float16")): @Ts.prim_func def add_one( - A: T.Buffer(16, "float16"), - B: T.Buffer(32, "float16"), - C: T.Buffer(16, "float16"), + A: T.Tensor(16, "float16"), + B: T.Tensor(32, "float16"), + C: T.Tensor(16, "float16"), ): for i in range(32): with Ts.sblock("inplace_B"): diff --git a/tests/python/relax/test_ast_printer.py b/tests/python/relax/test_ast_printer.py index efea39ea82ab..1a1d2954d56f 100644 --- a/tests/python/relax/test_ast_printer.py +++ b/tests/python/relax/test_ast_printer.py @@ -448,8 +448,8 @@ def test_call_tir(): class TestCallTIR: @Ts.prim_func def addone( - A: T.Buffer((m_addone, n_addone), "float32"), - B: T.Buffer((m_addone, n_addone), "float32"), + A: T.Tensor((m_addone, n_addone), "float32"), + B: T.Tensor((m_addone, n_addone), "float32"), ) -> None: T.func_attr({"global_symbol": "addone"}) for i, j in T.grid(m_addone, n_addone): diff --git a/tests/python/relax/test_backend_dispatch_sampling.py b/tests/python/relax/test_backend_dispatch_sampling.py index b14b21ad5f31..677025653558 100644 --- a/tests/python/relax/test_backend_dispatch_sampling.py +++ b/tests/python/relax/test_backend_dispatch_sampling.py @@ -51,7 +51,7 @@ def test_dispatch_multinomial_from_uniform_generic(): @I.ir_module class Expected: @Ts.prim_func(private=True) - def get_sample_index(prob: T.Buffer((batch, vocab_size)), usample: T.Buffer((out_batch, 1)), sample_indices: T.Buffer((out_batch, 1), 'int64'), output_index: T.Buffer((out_batch, 1), 'int64')): + def get_sample_index(prob: T.Tensor((batch, vocab_size)), usample: T.Tensor((out_batch, 1)), sample_indices: T.Tensor((out_batch, 1), 'int64'), output_index: T.Tensor((out_batch, 1), 'int64')): # with Ts.sblock("root"): for ax0, ax1 in T.grid(out_batch, vocab_size): @@ -89,7 +89,7 @@ def test_dispatch_multinomial_from_uniform_gpu(): @I.ir_module class Expected: @Ts.prim_func - def parallel_sampling_from_prob(prob: T.Buffer((n, vocab_size)), uniform_samples: T.Buffer((batch_size, 1)), row_indices: T.Buffer((batch_size, 1), 'int64'), token_ids: T.Buffer((batch_size, 1), 'int64')): + def parallel_sampling_from_prob(prob: T.Tensor((n, vocab_size)), uniform_samples: T.Tensor((batch_size, 1)), row_indices: T.Tensor((batch_size, 1), 'int64'), token_ids: T.Tensor((batch_size, 1), 'int64')): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): diff --git a/tests/python/relax/test_backend_transform_shape_lower.py b/tests/python/relax/test_backend_transform_shape_lower.py index 7748187dd06f..f1d605e5fdef 100644 --- a/tests/python/relax/test_backend_transform_shape_lower.py +++ b/tests/python/relax/test_backend_transform_shape_lower.py @@ -40,7 +40,7 @@ def main(x: R.Shape([1, 2]), y: R.Shape): return x @Ts.prim_func - def extra_func(H: T.Buffer(T.int64(4), "int64")): + def extra_func(H: T.Tensor(T.int64(4), "int64")): """Extra function, checks if the pass preserves it.""" H[T.int64(1)] = H[T.int64(0)] + T.int64(1) @@ -67,7 +67,7 @@ def main(x: R.Shape([1, 2]), y: R.Shape): return x @Ts.prim_func - def extra_func(H: T.Buffer(T.int64(4), "int64")): + def extra_func(H: T.Tensor(T.int64(4), "int64")): H[T.int64(1)] = H[T.int64(0)] + T.int64(1) before = Before @@ -205,7 +205,7 @@ def main(x: R.Tensor([n, m], "float32"), y: R.Tensor(ndim=3, dtype=None)) -> R.S @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def shape_func(H: T.Buffer(T.int64(4), "int64")): + def shape_func(H: T.Tensor(T.int64(4), "int64")): # generated compute function T.func_attr({"tirx.is_host_func": True}) H[T.int64(sindex["k+1"])] = H[T.int64(sindex["k"])] + T.int64(1) @@ -546,7 +546,7 @@ def main(x: R.Tensor([n, n], "float32")) -> R.Tensor([n * n], "float32"): return out @Ts.prim_func(private=True) - def shape_func(H: T.Buffer(T.int64(2), "int64")): + def shape_func(H: T.Tensor(T.int64(2), "int64")): # generated compute function T.func_attr({"tirx.is_host_func": True}) H[T.int64(sindex["n * n"])] = H[T.int64(sindex["n"])] * H[T.int64(sindex["n"])] diff --git a/tests/python/relax/test_base_py_module_printer.py b/tests/python/relax/test_base_py_module_printer.py index c93f626b1616..29ee974a248c 100644 --- a/tests/python/relax/test_base_py_module_printer.py +++ b/tests/python/relax/test_base_py_module_printer.py @@ -50,14 +50,14 @@ def multiply(self, x, y): @Ts.prim_func def add_tir( - x: T.Buffer((5,), "float32"), y: T.Buffer((5,), "float32"), out: T.Buffer((5,), "float32") + x: T.Tensor((5,), "float32"), y: T.Tensor((5,), "float32"), out: T.Tensor((5,), "float32") ): for i in range(5): out[i] = x[i] + y[i] @Ts.prim_func def multiply_tir( - x: T.Buffer((5,), "float32"), y: T.Buffer((5,), "float32"), out: T.Buffer((5,), "float32") + x: T.Tensor((5,), "float32"), y: T.Tensor((5,), "float32"), out: T.Tensor((5,), "float32") ): for i in range(5): out[i] = x[i] * y[i] @@ -124,7 +124,7 @@ def data_preprocessing(self, raw_data): return self._convert_tvm_to_pytorch(result) @Ts.prim_func - def extract_features(Data: T.Buffer((10,), "float32"), Features: T.Buffer((10,), "float32")): + def extract_features(Data: T.Tensor((10,), "float32"), Features: T.Tensor((10,), "float32")): T.func_attr({"tirx.noalias": True}) for i in range(10): @@ -132,9 +132,9 @@ def extract_features(Data: T.Buffer((10,), "float32"), Features: T.Buffer((10,), @Ts.prim_func def ml_inference( - Features: T.Buffer((10,), "float32"), - Params: T.Buffer((10,), "float32"), - Output: T.Buffer((5,), "float32"), + Features: T.Tensor((10,), "float32"), + Params: T.Tensor((10,), "float32"), + Output: T.Tensor((5,), "float32"), ): T.func_attr({"tirx.noalias": True}) @@ -142,14 +142,14 @@ def ml_inference( Output[i] = Features[i] * Params[i] + Features[i + 5] * Params[i + 5] @Ts.prim_func - def post_process(Predictions: T.Buffer((5,), "float32"), Final: T.Buffer((5,), "float32")): + def post_process(Predictions: T.Tensor((5,), "float32"), Final: T.Tensor((5,), "float32")): T.func_attr({"tirx.noalias": True}) for i in range(5): Final[i] = T.max(Predictions[i], 0.0) @Ts.prim_func - def normalize_data(Data: T.Buffer((10,), "float32"), Normalized: T.Buffer((10,), "float32")): + def normalize_data(Data: T.Tensor((10,), "float32"), Normalized: T.Tensor((10,), "float32")): T.func_attr({"tirx.noalias": True}) for i in range(10): @@ -203,7 +203,7 @@ def loop_with_break(self, data, max_iter): return result @Ts.prim_func - def dummy_tir(Data: T.Buffer((1,), "float32"), Output: T.Buffer((1,), "float32")): + def dummy_tir(Data: T.Tensor((1,), "float32"), Output: T.Tensor((1,), "float32")): T.func_attr({"tirx.noalias": True}) Output[0] = Data[0] @@ -263,7 +263,7 @@ def memory_efficient_transform(self, large_tensor): @Ts.prim_func def vectorized_add( - A: T.Buffer((10,), "float32"), B: T.Buffer((10,), "float32"), C: T.Buffer((10,), "float32") + A: T.Tensor((10,), "float32"), B: T.Tensor((10,), "float32"), C: T.Tensor((10,), "float32") ): T.func_attr({"tirx.noalias": True}) @@ -333,7 +333,7 @@ def multi_stage_pipeline(self, raw_input): return final_result @Ts.prim_func - def final_transform(Data: T.Buffer((10, 10), "float32"), Output: T.Buffer((10, 10), "float32")): + def final_transform(Data: T.Tensor((10, 10), "float32"), Output: T.Tensor((10, 10), "float32")): T.func_attr({"tirx.noalias": True}) for i in range(10): @@ -396,7 +396,7 @@ def graceful_degradation(self, primary_input, fallback_input): return self._get_safe_default() @Ts.prim_func - def safe_transform(Data: T.Buffer((5,), "float32"), Output: T.Buffer((5,), "float32")): + def safe_transform(Data: T.Tensor((5,), "float32"), Output: T.Tensor((5,), "float32")): T.func_attr({"tirx.noalias": True}) for i in range(5): 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 23e376c62f76..7007766342ca 100644 --- a/tests/python/relax/test_base_py_module_symbolic_shape.py +++ b/tests/python/relax/test_base_py_module_symbolic_shape.py @@ -71,9 +71,9 @@ def test_infer_concrete_shape_error_when_uninferrable(): class AddModuleSymbolic(BasePyModule): @Ts.prim_func def add_tir( - x: T.Buffer((n_add_tir,), dtype="float32"), - y: T.Buffer((n_add_tir,), dtype="float32"), - out: T.Buffer((n_add_tir,), dtype="float32"), + x: T.Tensor((n_add_tir,), dtype="float32"), + y: T.Tensor((n_add_tir,), dtype="float32"), + out: T.Tensor((n_add_tir,), dtype="float32"), ): T.func_attr({"global_symbol": "add_tir"}) @@ -209,9 +209,9 @@ def test_infer_concrete_shape_wrong_ndim(): class MatrixModuleSymbolic(BasePyModule): @Ts.prim_func def matmul_tir( - a: T.Buffer((m_matmul_tir, k_matmul_tir), dtype="float32"), - b: T.Buffer((k_matmul_tir, n_matmul_tir), dtype="float32"), - c: T.Buffer((m_matmul_tir, n_matmul_tir), dtype="float32"), + a: T.Tensor((m_matmul_tir, k_matmul_tir), dtype="float32"), + b: T.Tensor((k_matmul_tir, n_matmul_tir), dtype="float32"), + c: T.Tensor((m_matmul_tir, n_matmul_tir), dtype="float32"), ): T.func_attr({"global_symbol": "matmul_tir"}) diff --git a/tests/python/relax/test_blockbuilder_emit_te.py b/tests/python/relax/test_blockbuilder_emit_te.py index 1cbeea03c6fa..75c8affada3a 100644 --- a/tests/python/relax/test_blockbuilder_emit_te.py +++ b/tests/python/relax/test_blockbuilder_emit_te.py @@ -49,9 +49,9 @@ def te_func(A, offset): class Expected: @Ts.prim_func(private=True) def te_func( - A: T.Buffer((T.int64(10),), "float32"), + A: T.Tensor((T.int64(10),), "float32"), m: T.int64, - B: T.Buffer((T.int64(10),), "float32"), + B: T.Tensor((T.int64(10),), "float32"), ): T.func_attr({"tirx.noalias": True}) for i in range(T.int64(10)): @@ -96,9 +96,9 @@ def from_builder(): class Expected: @Ts.prim_func(private=True) def te_slice( - A: T.Buffer([T.int64(16), T.int64(16)], "float32"), + A: T.Tensor([T.int64(16), T.int64(16)], "float32"), row_index: T.int64, - Output: T.Buffer(T.int64(16), "float32"), + Output: T.Tensor(T.int64(16), "float32"), ): T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/relax/test_codegen_cutlass.py b/tests/python/relax/test_codegen_cutlass.py index 7858663a2d84..b40f9ccc1d61 100644 --- a/tests/python/relax/test_codegen_cutlass.py +++ b/tests/python/relax/test_codegen_cutlass.py @@ -1269,9 +1269,9 @@ def test_fp16A_int4B_gemm(): class Module: @Ts.prim_func def decode( - A: T.Buffer((T.int64(64), T.int64(64)), "int8"), - B: T.Buffer((T.int64(128),), "float16"), - decode_1: T.Buffer((T.int64(64), T.int64(128)), "float16"), + A: T.Tensor((T.int64(64), T.int64(64)), "int8"), + B: T.Tensor((T.int64(128),), "float16"), + decode_1: T.Tensor((T.int64(64), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1302,9 +1302,9 @@ def decode( @Ts.prim_func def encode( - A: T.Buffer((T.int64(128), T.int64(64)), "float16"), - w_gathered: T.Buffer((T.int64(64), T.int64(64)), "int8"), - compute: T.Buffer((T.int64(128),), "float16"), + A: T.Tensor((T.int64(128), T.int64(64)), "float16"), + w_gathered: T.Tensor((T.int64(64), T.int64(64)), "int8"), + compute: T.Tensor((T.int64(128),), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1526,9 +1526,9 @@ def test_fp16A_int8B_gemm(): class Module: @Ts.prim_func def decode( - A: T.Buffer((T.int64(64), T.int64(64)), "int8"), - B: T.Buffer((T.int64(64),), "float16"), - decode_1: T.Buffer((T.int64(64), T.int64(64)), "float16"), + A: T.Tensor((T.int64(64), T.int64(64)), "int8"), + B: T.Tensor((T.int64(64),), "float16"), + decode_1: T.Tensor((T.int64(64), T.int64(64)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1541,9 +1541,9 @@ def decode( @Ts.prim_func def encode( - A: T.Buffer((T.int64(64), T.int64(64)), "float16"), - w_gathered: T.Buffer((T.int64(64), T.int64(64)), "int8"), - compute: T.Buffer((T.int64(64),), "float16"), + A: T.Tensor((T.int64(64), T.int64(64)), "float16"), + w_gathered: T.Tensor((T.int64(64), T.int64(64)), "int8"), + compute: T.Tensor((T.int64(64),), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1672,9 +1672,9 @@ def test_rms_norm(): class Module: @Ts.prim_func def rms_norm( - A: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), - B: T.Buffer((T.int64(4096),), "float16"), - rms_norm: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), + B: T.Tensor((T.int64(4096),), "float16"), + rms_norm: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1807,9 +1807,9 @@ def test_fp16A_int8B_gemm_batched(): class Module: @Ts.prim_func def decode( - A: T.Buffer((T.int64(64), T.int64(64)), "int8"), - B: T.Buffer((T.int64(64),), "float16"), - decode_1: T.Buffer((T.int64(64), T.int64(64)), "float16"), + A: T.Tensor((T.int64(64), T.int64(64)), "int8"), + B: T.Tensor((T.int64(64),), "float16"), + decode_1: T.Tensor((T.int64(64), T.int64(64)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1822,9 +1822,9 @@ def decode( @Ts.prim_func def encode( - A: T.Buffer((T.int64(64), T.int64(64)), "float16"), - w_gathered: T.Buffer((T.int64(64), T.int64(64)), "int8"), - compute: T.Buffer((T.int64(64),), "float16"), + A: T.Tensor((T.int64(64), T.int64(64)), "float16"), + w_gathered: T.Tensor((T.int64(64), T.int64(64)), "int8"), + compute: T.Tensor((T.int64(64),), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1942,9 +1942,9 @@ def test_fp16A_int8B_gemm_batched_finegrained(): class Module: @Ts.prim_func def decode( - A: T.Buffer((T.int64(128), T.int64(128)), "int8"), - B: T.Buffer((T.int64(2), T.int64(128)), "float16"), - decode_1: T.Buffer((T.int64(128), T.int64(128)), "float16"), + A: T.Tensor((T.int64(128), T.int64(128)), "int8"), + B: T.Tensor((T.int64(2), T.int64(128)), "float16"), + decode_1: T.Tensor((T.int64(128), T.int64(128)), "float16"), ): T.func_attr({"tirx.noalias": True}) for i, j in T.grid(T.int64(128), T.int64(128)): @@ -1956,9 +1956,9 @@ def decode( @Ts.prim_func def encode( - A: T.Buffer((T.int64(128), T.int64(128)), "float16"), - w_gathered: T.Buffer((T.int64(128), T.int64(128)), "int8"), - compute: T.Buffer( + A: T.Tensor((T.int64(128), T.int64(128)), "float16"), + w_gathered: T.Tensor((T.int64(128), T.int64(128)), "int8"), + compute: T.Tensor( ( T.int64(2), T.int64(128), diff --git a/tests/python/relax/test_dataflow_inplace.py b/tests/python/relax/test_dataflow_inplace.py index fcc85594d42e..0e2f34590371 100644 --- a/tests/python/relax/test_dataflow_inplace.py +++ b/tests/python/relax/test_dataflow_inplace.py @@ -180,7 +180,7 @@ def test_alias_call_tir(): class AliasCallTir: @Ts.prim_func def tir_id( - A: T.Buffer((m_tir_id, n_tir_id), "int32"), B: T.Buffer((m_tir_id, n_tir_id), "int32") + A: T.Tensor((m_tir_id, n_tir_id), "int32"), B: T.Tensor((m_tir_id, n_tir_id), "int32") ) -> None: T.func_attr({"global_symbol": "tir_id"}) @@ -191,9 +191,9 @@ def tir_id( @Ts.prim_func def tir_id2( - A: T.Buffer((m_tir_id2, n_tir_id2), "int32"), - B: T.Buffer((m_tir_id2, n_tir_id2), "int32"), - C: T.Buffer((m_tir_id2, n_tir_id2), "int32"), + A: T.Tensor((m_tir_id2, n_tir_id2), "int32"), + B: T.Tensor((m_tir_id2, n_tir_id2), "int32"), + C: T.Tensor((m_tir_id2, n_tir_id2), "int32"), ) -> None: T.func_attr({"global_symbol": "tir_id"}) @@ -382,8 +382,8 @@ def main( @Ts.prim_func(private=True) def expected_add( - A: T.Buffer((T.int64(2), T.int64(3)), "float32"), - B: T.Buffer((T.int64(2), T.int64(3)), "float32"), + A: T.Tensor((T.int64(2), T.int64(3)), "float32"), + B: T.Tensor((T.int64(2), T.int64(3)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(2), T.int64(3)): @@ -401,7 +401,7 @@ def expected_add( new_add.attrs.inplace_indices == [0] @Ts.prim_func(private=True) - def expected_silu(A: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def expected_silu(A: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) compute = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3))) for i0, i1 in T.grid(T.int64(2), T.int64(3)): @@ -450,8 +450,8 @@ def main( class Expected: @Ts.prim_func(private=True) def add_inplace( - A: T.Buffer((T.int64(2), T.int64(3)), "float32"), - B: T.Buffer((T.int64(1), T.int64(3)), "float32"), + A: T.Tensor((T.int64(2), T.int64(3)), "float32"), + B: T.Tensor((T.int64(1), T.int64(3)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(2), T.int64(3)): @@ -463,8 +463,8 @@ def add_inplace( @Ts.prim_func(private=True) def multiply_inplace( - A: T.Buffer((T.int64(2), T.int64(3)), "float32"), - B: T.Buffer((T.int64(1), T.int64(3)), "float32"), + A: T.Tensor((T.int64(2), T.int64(3)), "float32"), + B: T.Tensor((T.int64(1), T.int64(3)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(2), T.int64(3)): @@ -476,8 +476,8 @@ def multiply_inplace( @Ts.prim_func(private=True) def subtract_inplace( - A: T.Buffer((T.int64(1), T.int64(3)), "float32"), - B: T.Buffer((T.int64(1), T.int64(3)), "float32"), + A: T.Tensor((T.int64(1), T.int64(3)), "float32"), + B: T.Tensor((T.int64(1), T.int64(3)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(1), T.int64(3)): @@ -578,7 +578,7 @@ def main( class Expected: @Ts.prim_func(private=True) def add_inplace( - A: T.Buffer((a_add_inplace, b_add_inplace)), B: T.Buffer((a_add_inplace, b_add_inplace)) + A: T.Tensor((a_add_inplace, b_add_inplace)), B: T.Tensor((a_add_inplace, b_add_inplace)) ): T.func_attr({"tirx.noalias": True}) @@ -591,8 +591,8 @@ def add_inplace( @Ts.prim_func(private=True) def subtract_inplace( - A: T.Buffer((a_subtract_inplace, b_subtract_inplace)), - B: T.Buffer((a_subtract_inplace, b_subtract_inplace)), + A: T.Tensor((a_subtract_inplace, b_subtract_inplace)), + B: T.Tensor((a_subtract_inplace, b_subtract_inplace)), ): T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/relax/test_dataflow_pattern.py b/tests/python/relax/test_dataflow_pattern.py index 81cdd36b2231..c7b1f56006e6 100644 --- a/tests/python/relax/test_dataflow_pattern.py +++ b/tests/python/relax/test_dataflow_pattern.py @@ -37,7 +37,7 @@ @tvm.script.ir_module class Module: @Ts.prim_func - def tir_matmul(A: T.Buffer((32, 32)), B: T.Buffer((32, 32)), C: T.Buffer((32, 32))) -> None: + def tir_matmul(A: T.Tensor((32, 32)), B: T.Tensor((32, 32)), C: T.Tensor((32, 32))) -> None: T.func_attr({"global_symbol": "tir_matmul"}) for i0, j0, k0 in T.grid(32, 32, 32): @@ -48,7 +48,7 @@ def tir_matmul(A: T.Buffer((32, 32)), B: T.Buffer((32, 32)), C: T.Buffer((32, 32 C[i, j] += A[i, k] * B[j, k] @Ts.prim_func - def tir_relu(A: T.Buffer((32, 32)), B: T.Buffer((32, 32))): + def tir_relu(A: T.Tensor((32, 32)), B: T.Tensor((32, 32))): T.func_attr({"global_symbol": "tir_relu"}) for i, j in T.grid(32, 32): @@ -57,7 +57,7 @@ def tir_relu(A: T.Buffer((32, 32)), B: T.Buffer((32, 32))): B[vi, vj] = T.max(A[vi, vj], 0.0) @Ts.prim_func - def tir_zeros(n: T.int64, A: T.Buffer([n])): + def tir_zeros(n: T.int64, A: T.Tensor([n])): T.func_attr({"global_symbol": "tir_zeros"}) for i in range(n): diff --git a/tests/python/relax/test_dataflow_rewriter.py b/tests/python/relax/test_dataflow_rewriter.py index db9d46b201ba..4b3c5a99df05 100644 --- a/tests/python/relax/test_dataflow_rewriter.py +++ b/tests/python/relax/test_dataflow_rewriter.py @@ -441,7 +441,7 @@ def replacement(A: R.Tensor([16], "float32")): return R.call_tir(RewriteMul.subroutine_mul, [A], out_ty=R.Tensor([16], "float32")) @Ts.prim_func(private=True) - def subroutine_mul(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): + def subroutine_mul(A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32")): for i in range(16): B[i] = A[i] * A[i] @@ -519,7 +519,7 @@ def replacement(A: R.Tensor([16], "float32")): return R.call_tir(RewriteMul.subroutine, [A], out_ty=R.Tensor([16], "float32")) @Ts.prim_func(private=True) - def subroutine(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): + def subroutine(A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32")): for i in range(16): B[i] = A[i] * A[i] @@ -544,7 +544,7 @@ def subroutine(A: R.Tensor([16], "float32")) -> R.Tensor([16], "float32"): return A * R.const(2.0, "float32") @Ts.prim_func(private=True) - def subroutine_1(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): + def subroutine_1(A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32")): for i in range(16): B[i] = A[i] * A[i] diff --git a/tests/python/relax/test_dlpack_integration.py b/tests/python/relax/test_dlpack_integration.py index ddfc0064738e..429069f5c64f 100644 --- a/tests/python/relax/test_dlpack_integration.py +++ b/tests/python/relax/test_dlpack_integration.py @@ -203,12 +203,12 @@ def test_dlpack_with_base_py_module(self): # Create a simple IRModule @Ts.prim_func - def identity_func(A: T.Buffer((3,), "float32"), B: T.Buffer((3,), "float32")): + def identity_func(A: T.Tensor((3,), "float32"), B: T.Tensor((3,), "float32")): for i in T.grid(3): B[i] = A[i] @Ts.prim_func - def constant_func(B: T.Buffer((2,), "float32")): + def constant_func(B: T.Tensor((2,), "float32")): for i in T.grid(2): B[i] = T.float32(5.0) diff --git a/tests/python/relax/test_eliminate_pad_branch_using_buffer_assumption.py b/tests/python/relax/test_eliminate_pad_branch_using_buffer_assumption.py index f3d29294a87e..cd9574f7d142 100644 --- a/tests/python/relax/test_eliminate_pad_branch_using_buffer_assumption.py +++ b/tests/python/relax/test_eliminate_pad_branch_using_buffer_assumption.py @@ -34,15 +34,15 @@ class AddBefore: @Ts.prim_func(private=True) def add( - a: T.Buffer( + a: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - b: T.Buffer( + b: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - compute: T.Buffer( + compute: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), @@ -131,15 +131,15 @@ def main( class AddExpected: @Ts.prim_func(private=True) def add( - a: T.Buffer( + a: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - b: T.Buffer( + b: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - compute: T.Buffer( + compute: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), @@ -233,15 +233,15 @@ def main( class SubBefore: @Ts.prim_func(private=True) def sub( - a: T.Buffer( + a: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - b: T.Buffer( + b: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - compute: T.Buffer( + compute: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), @@ -330,15 +330,15 @@ def main( class SubExpected: @Ts.prim_func(private=True) def sub( - a: T.Buffer( + a: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - b: T.Buffer( + b: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - compute: T.Buffer( + compute: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), @@ -432,15 +432,15 @@ def main( class MulBefore: @Ts.prim_func(private=True) def mul( - a: T.Buffer( + a: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - b: T.Buffer( + b: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - compute: T.Buffer( + compute: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), @@ -529,15 +529,15 @@ def main( class MulExpected: @Ts.prim_func(private=True) def mul( - a: T.Buffer( + a: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - b: T.Buffer( + b: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), - compute: T.Buffer( + compute: T.Tensor( (T.int64(1), T.int64(4), T.int64(4), T.int64(16), T.int64(8), T.int64(8), T.int64(32)), "uint8", ), diff --git a/tests/python/relax/test_frontend_common.py b/tests/python/relax/test_frontend_common.py index 79e9a09a4821..0e6105c4b125 100644 --- a/tests/python/relax/test_frontend_common.py +++ b/tests/python/relax/test_frontend_common.py @@ -71,8 +71,8 @@ def test_constant(self): class expected: @Ts.prim_func(private=True) def pad( - x: T.Buffer((T.int64(1), T.int64(1), T.int64(4), T.int64(4)), "float32"), - PadInput: T.Buffer((T.int64(1), T.int64(1), T.int64(5), T.int64(5)), "float32"), + x: T.Tensor((T.int64(1), T.int64(1), T.int64(4), T.int64(4)), "float32"), + PadInput: T.Tensor((T.int64(1), T.int64(1), T.int64(5), T.int64(5)), "float32"), ): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(1), T.int64(1), T.int64(5), T.int64(5)): @@ -107,8 +107,8 @@ def test_edge(self): class expected: @Ts.prim_func(private=True) def replicate_pad( - x: T.Buffer((T.int64(1), T.int64(1), T.int64(4), T.int64(4)), "float32"), - ReplicatePadInput: T.Buffer( + x: T.Tensor((T.int64(1), T.int64(1), T.int64(4), T.int64(4)), "float32"), + ReplicatePadInput: T.Tensor( (T.int64(1), T.int64(1), T.int64(5), T.int64(5)), "float32" ), ): @@ -152,8 +152,8 @@ def test_reflect(self): class expected: @Ts.prim_func(private=True) def mirror_pad( - x: T.Buffer((T.int64(1), T.int64(1), T.int64(4), T.int64(4)), "float32"), - MirrorPadInput: T.Buffer( + x: T.Tensor((T.int64(1), T.int64(1), T.int64(4), T.int64(4)), "float32"), + MirrorPadInput: T.Tensor( (T.int64(1), T.int64(1), T.int64(5), T.int64(5)), "float32" ), ): diff --git a/tests/python/relax/test_frontend_dynamo.py b/tests/python/relax/test_frontend_dynamo.py index 6b140ca9e5ed..bb7e4bd3d99b 100644 --- a/tests/python/relax/test_frontend_dynamo.py +++ b/tests/python/relax/test_frontend_dynamo.py @@ -54,10 +54,10 @@ def forward(self, x): class Input1_ir: @Ts.prim_func def main( - inp_0: T.Buffer((T.int64(10), T.int64(100)), "float32"), - param_0: T.Buffer((T.int64(100), T.int64(10)), "float32"), - param_1: T.Buffer(T.int64(10), "float32"), - compute: T.Buffer((T.int64(10), T.int64(10)), "float32"), + inp_0: T.Tensor((T.int64(10), T.int64(100)), "float32"), + param_0: T.Tensor((T.int64(100), T.int64(10)), "float32"), + param_1: T.Tensor(T.int64(10), "float32"), + compute: T.Tensor((T.int64(10), T.int64(10)), "float32"), ): # function attr dict T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) diff --git a/tests/python/relax/test_frontend_nn_op.py b/tests/python/relax/test_frontend_nn_op.py index 20884054c151..a6552e318cbb 100644 --- a/tests/python/relax/test_frontend_nn_op.py +++ b/tests/python/relax/test_frontend_nn_op.py @@ -595,7 +595,7 @@ def test(self, x: Tensor): @I.ir_module class Expected: @Ts.prim_func(private=True) - def add_one(A: T.Buffer((T.int64(10), T.int64(10)), "float32"), T_add: T.Buffer((T.int64(10), T.int64(10)), "float32")): + def add_one(A: T.Tensor((T.int64(10), T.int64(10)), "float32"), T_add: T.Tensor((T.int64(10), T.int64(10)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1 in T.grid(T.int64(10), T.int64(10)): @@ -641,11 +641,11 @@ def test_tensor_ir_op(): @Ts.prim_func(private=True) def fused_rope( # pylint: disable=too-many-locals - qkv: T.Buffer((batch_size, seq_len, fused_heads, head_dim), dtype), + qkv: T.Tensor((batch_size, seq_len, fused_heads, head_dim), dtype), offset: T.int64, - q: T.Buffer((batch_size, seq_len, num_q_heads, head_dim), dtype), - k: T.Buffer((batch_size, seq_len, num_kv_heads, head_dim), dtype), - v: T.Buffer((batch_size, seq_len, num_kv_heads, head_dim), dtype), + q: T.Tensor((batch_size, seq_len, num_q_heads, head_dim), dtype), + k: T.Tensor((batch_size, seq_len, num_kv_heads, head_dim), dtype), + v: T.Tensor((batch_size, seq_len, num_kv_heads, head_dim), dtype), ): T.evaluate(offset) @@ -671,7 +671,7 @@ def test(self, qkv: Tensor, offset: tirx.Var): @I.ir_module class Expected: @Ts.prim_func(private=True) - def llama_fused_rope(qkv: T.Buffer((batch_size, seq_len, 24, 16), 'float16'), offset: T.int64, q: T.Buffer((batch_size, seq_len, 8, 16), 'float16'), k: T.Buffer((batch_size, seq_len, 8, 16), 'float16'), v: T.Buffer((batch_size, seq_len, 8, 16), 'float16')): + def llama_fused_rope(qkv: T.Tensor((batch_size, seq_len, 24, 16), 'float16'), offset: T.int64, q: T.Tensor((batch_size, seq_len, 8, 16), 'float16'), k: T.Tensor((batch_size, seq_len, 8, 16), 'float16'), v: T.Tensor((batch_size, seq_len, 8, 16), 'float16')): T.evaluate(offset) @@ -718,9 +718,9 @@ def test_tensor_ir_inplace_op(): @Ts.prim_func def inplace_take( - weight: T.Buffer((vocab_size, hidden_size), dtype), - pos: T.Buffer((seq_len,), "int32"), - embeddings: T.Buffer((total_seq_len, hidden_size), dtype), + weight: T.Tensor((vocab_size, hidden_size), dtype), + pos: T.Tensor((seq_len,), "int32"), + embeddings: T.Tensor((total_seq_len, hidden_size), dtype), offset: T.int64, ): T.func_attr({"tirx.noalias": True}) @@ -757,9 +757,9 @@ def test( class Expected: @Ts.prim_func def inplace_take( - weight: T.Buffer((vocab_size_inplace_take, hidden_size), dtype), - pos: T.Buffer((seq_len_inplace_take,), "int32"), - embeddings: T.Buffer((total_seq_len_inplace_take, hidden_size), dtype), + weight: T.Tensor((vocab_size_inplace_take, hidden_size), dtype), + pos: T.Tensor((seq_len_inplace_take,), "int32"), + embeddings: T.Tensor((total_seq_len_inplace_take, hidden_size), dtype), offset: T.int64, ): T.func_attr({"tirx.noalias": True}) @@ -822,7 +822,7 @@ def test( def test_tensor_ir_op_no_tir_var(): @Ts.prim_func(private=True) - def tir_func(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")): + def tir_func(A: T.Tensor((16, 16), "float32"), B: T.Tensor((16, 16), "float32")): T.evaluate(0) class Model(Module): @@ -838,7 +838,7 @@ def test(self, A: Tensor): @I.ir_module class Expected: @Ts.prim_func(private=True) - def tir_func(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")): + def tir_func(A: T.Tensor((16, 16), "float32"), B: T.Tensor((16, 16), "float32")): T.evaluate(0) @R.function @@ -1027,7 +1027,7 @@ def foo( @I.ir_module class Expected: @Ts.prim_func(private=True) - def get_index_from_sorted(cumsum_sorted: T.Buffer((batch_get_index_from_sorted, vocab_size_get_index_from_sorted)), indices: T.Buffer((batch_get_index_from_sorted, vocab_size_get_index_from_sorted), 'int64'), renorm_prob: T.Buffer((batch_get_index_from_sorted, 1)), usample: T.Buffer((out_batch, 1)), sample_indices: T.Buffer((out_batch, 1), 'int64'), output_index: T.Buffer((out_batch, 1), 'int64')): + def get_index_from_sorted(cumsum_sorted: T.Tensor((batch_get_index_from_sorted, vocab_size_get_index_from_sorted)), indices: T.Tensor((batch_get_index_from_sorted, vocab_size_get_index_from_sorted), 'int64'), renorm_prob: T.Tensor((batch_get_index_from_sorted, 1)), usample: T.Tensor((out_batch, 1)), sample_indices: T.Tensor((out_batch, 1), 'int64'), output_index: T.Tensor((out_batch, 1), 'int64')): # with Ts.sblock("root"): for ax0, ax1 in T.grid(out_batch, vocab_size_get_index_from_sorted): @@ -1043,7 +1043,7 @@ def get_index_from_sorted(cumsum_sorted: T.Buffer((batch_get_index_from_sorted, output_index[v_ax0, 0] = indices[sample_indices[v_ax0, T.int64(0)], v_ax1] @Ts.prim_func(private=True) - def get_renorm_prob(cumsum_sorted: T.Buffer((batch_get_renorm_prob, vocab_size_get_renorm_prob)), top_p: T.Buffer((batch_get_renorm_prob, 1)), top_k: T.Buffer((batch_get_renorm_prob, 1), 'int64'), renorm_prob: T.Buffer((batch_get_renorm_prob, 1))): + def get_renorm_prob(cumsum_sorted: T.Tensor((batch_get_renorm_prob, vocab_size_get_renorm_prob)), top_p: T.Tensor((batch_get_renorm_prob, 1)), top_k: T.Tensor((batch_get_renorm_prob, 1), 'int64'), renorm_prob: T.Tensor((batch_get_renorm_prob, 1))): # with Ts.sblock("root"): for ax0, ax1 in T.grid(batch_get_renorm_prob, vocab_size_get_renorm_prob): @@ -1149,7 +1149,7 @@ def foo( @I.ir_module class Expected: @Ts.prim_func(private=True) - def filter_with_top_p_top_k(A: T.Buffer((T.int64(2), T.int64(3)), "float32"), B: T.Buffer((T.int64(2), T.int64(1)), "float32"), filter_with_top_p_top_k: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def filter_with_top_p_top_k(A: T.Tensor((T.int64(2), T.int64(3)), "float32"), B: T.Tensor((T.int64(2), T.int64(1)), "float32"), filter_with_top_p_top_k: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i, j in T.grid(T.int64(2), T.int64(3)): @@ -1160,7 +1160,7 @@ def filter_with_top_p_top_k(A: T.Buffer((T.int64(2), T.int64(3)), "float32"), B: filter_with_top_p_top_k[v_i, v_j] = T.Select(B[v_i, T.int64(0)] <= A[v_i, v_j], A[v_i, v_j], T.float32(0)) @Ts.prim_func(private=True) - def get_renorm_cutoff(sorted_prob: T.Buffer((batch, vocab_size)), cumsum_sorted: T.Buffer((batch, vocab_size)), top_p: T.Buffer((batch, 1)), top_k: T.Buffer((batch, 1), 'int64'), cutoff: T.Buffer((batch, 1))): + def get_renorm_cutoff(sorted_prob: T.Tensor((batch, vocab_size)), cumsum_sorted: T.Tensor((batch, vocab_size)), top_p: T.Tensor((batch, 1)), top_k: T.Tensor((batch, 1), 'int64'), cutoff: T.Tensor((batch, 1))): # with Ts.sblock("root"): for ax0, ax1 in T.grid(batch, vocab_size): diff --git a/tests/python/relax/test_op_index.py b/tests/python/relax/test_op_index.py index 5020ae4b3b0a..85666db36f94 100644 --- a/tests/python/relax/test_op_index.py +++ b/tests/python/relax/test_op_index.py @@ -921,9 +921,9 @@ def main(A: R.Tensor((16, 16), "float32"), B: R.Shape([index])) -> R.Tensor((1, @Ts.prim_func(private=True) def strided_slice( - A: T.Buffer((T.int64(16), T.int64(16))), + A: T.Tensor((T.int64(16), T.int64(16))), index: T.int64, - B: T.Buffer((T.int64(1), T.int64(16))), + B: T.Tensor((T.int64(1), T.int64(16))), ): T.func_attr({"tirx.noalias": True}) for (*iters,) in T.grid(*B.shape): @@ -956,7 +956,7 @@ class expected: strided_slice_index = T.int64() @Ts.prim_func(private=True) - def strided_slice(A: T.Buffer((T.int64(16), T.int64(16)), "float32"), index: strided_slice_index, T_dynamic_strided_slice_with_axes: T.Buffer((T.max(T.int64(16) - T.max(T.if_then_else(strided_slice_index < T.int64(0), strided_slice_index + T.int64(16), strided_slice_index), T.int64(0)), T.int64(0)), T.int64(16)))): + def strided_slice(A: T.Tensor((T.int64(16), T.int64(16)), "float32"), index: strided_slice_index, T_dynamic_strided_slice_with_axes: T.Tensor((T.max(T.int64(16) - T.max(T.if_then_else(strided_slice_index < T.int64(0), strided_slice_index + T.int64(16), strided_slice_index), T.int64(0)), T.int64(0)), T.int64(16)))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): diff --git a/tests/python/relax/test_op_misc.py b/tests/python/relax/test_op_misc.py index 68fdb50e5aa5..beef44cd55a2 100644 --- a/tests/python/relax/test_op_misc.py +++ b/tests/python/relax/test_op_misc.py @@ -31,7 +31,7 @@ def identity_packed(a): @Ts.prim_func -def identity_tir(A: T.Buffer([54, 96]), B: T.Buffer([54, 96])) -> None: +def identity_tir(A: T.Tensor([54, 96]), B: T.Tensor([54, 96])) -> None: for i, j in T.grid(54, 96): with Ts.sblock("compute"): vi, vj = Ts.axis.remap("SS", [i, j]) diff --git a/tests/python/relax/test_optimize_layout_transform.py b/tests/python/relax/test_optimize_layout_transform.py index 5838a2fb42f9..90c372157ece 100644 --- a/tests/python/relax/test_optimize_layout_transform.py +++ b/tests/python/relax/test_optimize_layout_transform.py @@ -47,9 +47,9 @@ def test_optimize_transform_layout_pass_one_arg(): class Before: @Ts.prim_func(private=True) def relax_add_replacement( - arg0: T.Buffer((4, 4), "float32"), - arg1: T.Buffer((4, 4), "float32"), - output: T.Buffer((4, 4), "float32"), + arg0: T.Tensor((4, 4), "float32"), + arg1: T.Tensor((4, 4), "float32"), + output: T.Tensor((4, 4), "float32"), ): T.func_attr({"operator_name": "relax.add"}) # with Ts.sblock("root"): @@ -101,9 +101,9 @@ def main( class Expected: @Ts.prim_func(private=True) def relax_add_replacement( - arg0: T.Buffer((4, 4), "float32"), - arg1: T.Buffer((4, 4), "float32"), - output: T.Buffer((4, 4), "float32"), + arg0: T.Tensor((4, 4), "float32"), + arg1: T.Tensor((4, 4), "float32"), + output: T.Tensor((4, 4), "float32"), ): T.func_attr({"operator_name": "relax.add"}) # with Ts.sblock("root"): @@ -149,9 +149,9 @@ def test_optimize_transform_layout_pass_two_args(): class Before: @Ts.prim_func(private=True) def relax_add_replacement( - arg0: T.Buffer((4, 4), "float32"), - arg1: T.Buffer((4, 4), "float32"), - output: T.Buffer((4, 4), "float32"), + arg0: T.Tensor((4, 4), "float32"), + arg1: T.Tensor((4, 4), "float32"), + output: T.Tensor((4, 4), "float32"), ): T.func_attr({"operator_name": "relax.add"}) # with Ts.sblock("root"): @@ -216,9 +216,9 @@ def main( class Expected: @Ts.prim_func(private=True) def relax_add_replacement( - arg0: T.Buffer((4, 4), "float32"), - arg1: T.Buffer((4, 4), "float32"), - output: T.Buffer((4, 4), "float32"), + arg0: T.Tensor((4, 4), "float32"), + arg1: T.Tensor((4, 4), "float32"), + output: T.Tensor((4, 4), "float32"), ): T.func_attr({"operator_name": "relax.add"}) # with Ts.sblock("root"): @@ -277,7 +277,7 @@ def test_tranform_layout_tir_remove_pad_transform_layout(): class Before: @Ts.prim_func(private=True) def relax_relu_replacement( - arg0: T.Buffer((16,), "float32"), output: T.Buffer((16,), "float32") + arg0: T.Tensor((16,), "float32"), output: T.Tensor((16,), "float32") ): T.func_attr({"operator_name": "relax.relu"}) # with Ts.sblock("root"): @@ -289,7 +289,7 @@ def relax_relu_replacement( output[v_ax0] = T.max(arg0[v_ax0], T.float32(0)) @Ts.prim_func(private=True) - def remove_pad(input: T.Buffer((p0,)), output: T.Buffer((i0,))): + def remove_pad(input: T.Tensor((p0,)), output: T.Tensor((i0,))): T.func_attr({"operator_name": "remove_pad", "tirx.noalias": True}) # with Ts.sblock("root"): @@ -350,7 +350,7 @@ def main(x: R.Tensor((14,), dtype="float32")) -> R.Tensor((14,), dtype="float32" class Expected: @Ts.prim_func(private=True) def relax_relu_replacement( - arg0: T.Buffer((16,), "float32"), output: T.Buffer((16,), "float32") + arg0: T.Tensor((16,), "float32"), output: T.Tensor((16,), "float32") ): T.func_attr({"operator_name": "relax.relu"}) # with Ts.sblock("root"): @@ -362,7 +362,7 @@ def relax_relu_replacement( output[v_ax0] = T.max(arg0[v_ax0], T.float32(0)) @Ts.prim_func(private=True) - def remove_pad(input: T.Buffer((p0,)), output: T.Buffer((i0,))): + def remove_pad(input: T.Tensor((p0,)), output: T.Tensor((i0,))): T.func_attr({"operator_name": "remove_pad", "tirx.noalias": True}) # with Ts.sblock("root"): diff --git a/tests/python/relax/test_pytorch_integration.py b/tests/python/relax/test_pytorch_integration.py index 5ecdc339b018..8447bb386954 100644 --- a/tests/python/relax/test_pytorch_integration.py +++ b/tests/python/relax/test_pytorch_integration.py @@ -64,9 +64,9 @@ def main(self, x: torch.Tensor, w: torch.Tensor) -> torch.Tensor: @Ts.prim_func def matmul( - A: T.Buffer((n_matmul, 16), "float32"), - B: T.Buffer((16, 20), "float32"), - C: T.Buffer((n_matmul, 20), "float32"), + A: T.Tensor((n_matmul, 16), "float32"), + B: T.Tensor((16, 20), "float32"), + C: T.Tensor((n_matmul, 20), "float32"), ): """TIR function for matrix multiplication.""" diff --git a/tests/python/relax/test_relax_to_pyfunc_converter.py b/tests/python/relax/test_relax_to_pyfunc_converter.py index c437d956035d..c1b387968995 100644 --- a/tests/python/relax/test_relax_to_pyfunc_converter.py +++ b/tests/python/relax/test_relax_to_pyfunc_converter.py @@ -48,7 +48,7 @@ class ComprehensiveTestModule: @Ts.prim_func def add_tir( - x: T.Buffer((5,), "float32"), y: T.Buffer((5,), "float32"), out: T.Buffer((5,), "float32") + x: T.Tensor((5,), "float32"), y: T.Tensor((5,), "float32"), out: T.Tensor((5,), "float32") ): """TIR function for addition.""" @@ -57,9 +57,9 @@ def add_tir( @Ts.prim_func def mul_tir( - x: T.Buffer((3, 4), "float32"), - y: T.Buffer((3, 4), "float32"), - out: T.Buffer((3, 4), "float32"), + x: T.Tensor((3, 4), "float32"), + y: T.Tensor((3, 4), "float32"), + out: T.Tensor((3, 4), "float32"), ): """TIR function for multiplication.""" @@ -883,9 +883,9 @@ def test_dlpack_conversion_fallback(self): class DLPackTestModule: @Ts.prim_func def test_tir( - x: T.Buffer((4,), "float32"), - y: T.Buffer((4,), "float32"), - out: T.Buffer((4,), "float32"), + x: T.Tensor((4,), "float32"), + y: T.Tensor((4,), "float32"), + out: T.Tensor((4,), "float32"), ): for i in range(4): out[i] = x[i] + y[i] @@ -937,9 +937,9 @@ def test_tvm_runtime_api_compatibility(self): class RuntimeAPITestModule: @Ts.prim_func def test_tir( - x: T.Buffer((3,), "float32"), - y: T.Buffer((3,), "float32"), - out: T.Buffer((3,), "float32"), + x: T.Tensor((3,), "float32"), + y: T.Tensor((3,), "float32"), + out: T.Tensor((3,), "float32"), ): for i in range(3): out[i] = x[i] * y[i] @@ -996,9 +996,9 @@ def test_mixed_tir_and_relax_operations(self): class MixedOpsTestModule: @Ts.prim_func def add_tir( - x: T.Buffer((4,), "float32"), - y: T.Buffer((4,), "float32"), - out: T.Buffer((4,), "float32"), + x: T.Tensor((4,), "float32"), + y: T.Tensor((4,), "float32"), + out: T.Tensor((4,), "float32"), ): for i in range(4): out[i] = x[i] + y[i] diff --git a/tests/python/relax/test_runtime_builtin_rnn_state.py b/tests/python/relax/test_runtime_builtin_rnn_state.py index 3072202f6368..3a5c5687fd72 100644 --- a/tests/python/relax/test_runtime_builtin_rnn_state.py +++ b/tests/python/relax/test_runtime_builtin_rnn_state.py @@ -216,10 +216,10 @@ def rnn_state_get( @Ts.prim_func def _rnn_state_get( - storage: T.Buffer((reserved_nseq, max_history, *shape), dtype), - seq_slot_ids: T.Buffer((batch_size,), 'int32'), - history_slot_ids: T.Buffer((batch_size,), 'int32'), - output: T.Buffer((batch_size, *shape), dtype), + storage: T.Tensor((reserved_nseq, max_history, *shape), dtype), + seq_slot_ids: T.Tensor((batch_size,), 'int32'), + history_slot_ids: T.Tensor((batch_size,), 'int32'), + output: T.Tensor((batch_size, *shape), dtype), ): for i in range(batch_size): @@ -247,10 +247,10 @@ def rnn_state_set( @Ts.prim_func def _rnn_state_set( - storage: T.Buffer((reserved_nseq, max_history, *shape), dtype), - seq_slot_ids: T.Buffer((batch_size,), 'int32'), - history_slot_ids: T.Buffer((batch_size,), 'int32'), - data: T.Buffer((batch_size, *shape), dtype), + storage: T.Tensor((reserved_nseq, max_history, *shape), dtype), + seq_slot_ids: T.Tensor((batch_size,), 'int32'), + history_slot_ids: T.Tensor((batch_size,), 'int32'), + data: T.Tensor((batch_size, *shape), dtype), ): for i in range(batch_size): diff --git a/tests/python/relax/test_tir_call_source_kernel.py b/tests/python/relax/test_tir_call_source_kernel.py index 3d1c9b4ae656..29c8c3c6b383 100644 --- a/tests/python/relax/test_tir_call_source_kernel.py +++ b/tests/python/relax/test_tir_call_source_kernel.py @@ -49,9 +49,9 @@ def test_tir_call_source_kernel(): class Module: @Ts.prim_func def add( - x: T.Buffer((m_add,), "float32"), - y: T.Buffer((m_add,), "float32"), - output: T.Buffer((m_add,), "float32"), + x: T.Tensor((m_add,), "float32"), + y: T.Tensor((m_add,), "float32"), + output: T.Tensor((m_add,), "float32"), ) -> None: T.func_attr({"global_symbol": "add"}) @@ -80,7 +80,7 @@ def main(x: R.Tensor((m_main,), "float32"), y: R.Tensor((m_main,), "float32")): @I.ir_module class Parsed: @Ts.prim_func - def add(x: T.Buffer((m,)), y: T.Buffer((m,)), output: T.Buffer((m,))): + def add(x: T.Tensor((m,)), y: T.Tensor((m,)), output: T.Tensor((m,))): with Ts.sblock("root"): Ts.reads(x[0:m], y[0:m]) Ts.writes(output[0:m]) diff --git a/tests/python/relax/test_transform.py b/tests/python/relax/test_transform.py index 3dabe404c900..26dff362a017 100644 --- a/tests/python/relax/test_transform.py +++ b/tests/python/relax/test_transform.py @@ -146,7 +146,7 @@ def test_call_tir_rewrite(): @tvm.script.ir_module class TestCallTIRRewrite: @Ts.prim_func - def exp(A: T.Buffer((m_exp, n_exp), "float32"), B: T.Buffer((m_exp, n_exp), "float32")): + def exp(A: T.Tensor((m_exp, n_exp), "float32"), B: T.Tensor((m_exp, n_exp), "float32")): T.evaluate(0) @R.function @@ -188,10 +188,10 @@ def test_call_tir_rewrite_with_interspersed_primitive_argument(): class Module: @Ts.prim_func def scale_add( - A: T.Buffer((16,), "float32"), + A: T.Tensor((16,), "float32"), scale: T.float32, - C: T.Buffer((16,), "float32"), - B: T.Buffer((16,), "float32"), + C: T.Tensor((16,), "float32"), + B: T.Tensor((16,), "float32"), ): for i in range(16): B[i] = A[i] + scale * C[i] @@ -412,7 +412,7 @@ def test_call_tir_inplace_simple(): @tvm.script.ir_module class Input: @Ts.prim_func - def zeros(A: T.Buffer((2, 3), "int32")): + def zeros(A: T.Tensor((2, 3), "int32")): # just overwrites A with 0s T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): @@ -431,7 +431,7 @@ def foo(x: R.Tensor((2, 3), "int32")) -> R.Tensor((2, 3), "int32"): @tvm.script.ir_module class Expected: @Ts.prim_func - def zeros(A: T.Buffer((2, 3), "int32")): + def zeros(A: T.Tensor((2, 3), "int32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_zeros"): @@ -455,7 +455,7 @@ def test_call_tir_inplace_multiple_args(): class Input: @Ts.prim_func def copy( - A: T.Buffer((2, 3), "int32"), B: T.Buffer((2, 3), "int32"), C: T.Buffer((2, 3), "int32") + A: T.Tensor((2, 3), "int32"), B: T.Tensor((2, 3), "int32"), C: T.Tensor((2, 3), "int32") ): # copies the contents of C into A and B T.func_attr({"tirx.noalias": True}) @@ -484,7 +484,7 @@ def foo( class Expected: @Ts.prim_func def copy( - A: T.Buffer((2, 3), "int32"), B: T.Buffer((2, 3), "int32"), C: T.Buffer((2, 3), "int32") + A: T.Tensor((2, 3), "int32"), B: T.Tensor((2, 3), "int32"), C: T.Tensor((2, 3), "int32") ): # copies the contents of C into A and B T.func_attr({"tirx.noalias": True}) @@ -514,11 +514,11 @@ def test_call_tir_inplace_some_new(): class Input: @Ts.prim_func def copy( - A: T.Buffer((2, 3), "int32"), - B: T.Buffer((2, 3), "int32"), - C: T.Buffer((2, 3), "int32"), - out1: T.Buffer((2, 3), "int32"), - out2: T.Buffer((2, 3), "int32"), + A: T.Tensor((2, 3), "int32"), + B: T.Tensor((2, 3), "int32"), + C: T.Tensor((2, 3), "int32"), + out1: T.Tensor((2, 3), "int32"), + out2: T.Tensor((2, 3), "int32"), ): # copies the contents of C into A, out1, and out2 T.func_attr({"tirx.noalias": True}) @@ -554,11 +554,11 @@ def foo( class Expected: @Ts.prim_func def copy( - A: T.Buffer((2, 3), "int32"), - B: T.Buffer((2, 3), "int32"), - C: T.Buffer((2, 3), "int32"), - out1: T.Buffer((2, 3), "int32"), - out2: T.Buffer((2, 3), "int32"), + A: T.Tensor((2, 3), "int32"), + B: T.Tensor((2, 3), "int32"), + C: T.Tensor((2, 3), "int32"), + out1: T.Tensor((2, 3), "int32"), + out2: T.Tensor((2, 3), "int32"), ): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): @@ -602,9 +602,9 @@ def test_call_tir_inplace_repeated_input(): class Input: @Ts.prim_func def func( - A: T.Buffer((2, 3), "int32"), - B: T.Buffer((2, 3), "int32"), - C: T.Buffer((2, 3), "int32"), + A: T.Tensor((2, 3), "int32"), + B: T.Tensor((2, 3), "int32"), + C: T.Tensor((2, 3), "int32"), ): T.evaluate(0) @@ -631,7 +631,7 @@ def test_call_tir_inplace_all_new(): @tvm.script.ir_module class Input: @Ts.prim_func - def func(A: T.Buffer((2, 3), "int32")): + def func(A: T.Tensor((2, 3), "int32")): T.evaluate(0) @R.function @@ -671,7 +671,7 @@ def main(A: R.Tensor((16,), dtype="float32")) -> R.Tensor((16,), dtype="float32" return gv1 @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer((16,), "float32")): + def multiply_by_two(A: T.Tensor((16,), "float32")): for i in range(16): A[i] = A[i] * T.float32(2) @@ -702,7 +702,7 @@ def main(A: R.Any): return gv1 @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer((16,), "float32")): + def multiply_by_two(A: T.Tensor((16,), "float32")): for i in range(16): A[i] = A[i] * T.float32(2) @@ -731,7 +731,7 @@ def main(A: R.Tensor([32], dtype="float32")): return gv1 @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer((16,), "float32")): + def multiply_by_two(A: T.Tensor((16,), "float32")): for i in range(16): A[i] = A[i] * T.float32(2) @@ -760,7 +760,7 @@ def main(A: R.Tensor([16], dtype="int32")): return gv1 @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer((16,), "float32")): + def multiply_by_two(A: T.Tensor((16,), "float32")): for i in range(16): A[i] = A[i] * T.float32(2) diff --git a/tests/python/relax/test_transform_alter_op_impl.py b/tests/python/relax/test_transform_alter_op_impl.py index 69d190b8d26c..3ad0adbb6d0e 100644 --- a/tests/python/relax/test_transform_alter_op_impl.py +++ b/tests/python/relax/test_transform_alter_op_impl.py @@ -46,7 +46,7 @@ def test_single_output(): @I.ir_module class Before: @Ts.prim_func(private=True) - def add(arg0: T.Buffer((16,), "float32"), arg1: T.Buffer((16,), "float32"), output: T.Buffer((16,), "float32")): + def add(arg0: T.Tensor((16,), "float32"), arg1: T.Tensor((16,), "float32"), output: T.Tensor((16,), "float32")): T.func_attr({"operator_name": "relax.add"}) for ax0 in range(16): with Ts.sblock("T_add"): @@ -65,7 +65,7 @@ def main(x: R.Tensor((16,), dtype="float32"), y: R.Tensor((16,), dtype="float32" @I.ir_module class Expected: @Ts.prim_func(private=True) - def relax_add_replacement(arg0: T.Buffer((4, 4), "float32"), arg1: T.Buffer((4, 4), "float32"), output: T.Buffer((4, 4), "float32")): + def relax_add_replacement(arg0: T.Tensor((4, 4), "float32"), arg1: T.Tensor((4, 4), "float32"), output: T.Tensor((4, 4), "float32")): T.func_attr({"operator_name": "relax.add"}) for ax0, ax1 in T.grid(4, 4): with Ts.sblock("T_add"): @@ -86,7 +86,7 @@ def main(x: R.Tensor((16,), dtype="float32"), y: R.Tensor((16,), dtype="float32" return gv @Ts.prim_func(private=True) - def add_2d(arg0: T.Buffer((4, 4), "float32"), arg1: T.Buffer((4, 4), "float32"), output: T.Buffer((4, 4), "float32")): + def add_2d(arg0: T.Tensor((4, 4), "float32"), arg1: T.Tensor((4, 4), "float32"), output: T.Tensor((4, 4), "float32")): for ax0, ax1 in T.grid(4, 4): with Ts.sblock("T_add"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) @@ -109,7 +109,7 @@ def test_empty_layout_changes(): @I.ir_module class Before: @Ts.prim_func(private=True) - def mul_by_2(arg0: T.Buffer((16,), "float32"), output: T.Buffer((16,), "float32")): + def mul_by_2(arg0: T.Tensor((16,), "float32"), output: T.Tensor((16,), "float32")): T.func_attr({"operator_name": "relax.mul_by_2"}) for ax0 in range(16): with Ts.sblock("T_add"): @@ -128,7 +128,7 @@ def main(x: R.Tensor((16,), dtype="float32")) -> R.Tensor((16,), dtype="float32" @I.ir_module class Expected: @Ts.prim_func(private=True) - def relax_mul_by_2_replacement(arg0: T.Buffer((16,), "float32"), output: T.Buffer((16,), "float32")): + def relax_mul_by_2_replacement(arg0: T.Tensor((16,), "float32"), output: T.Tensor((16,), "float32")): T.func_attr({"operator_name": "relax.mul_by_2"}) for ax0 in range(16): with Ts.sblock("T_add"): @@ -146,7 +146,7 @@ def main(x: R.Tensor((16,), dtype="float32")) -> R.Tensor((16,), dtype="float32" return gv @Ts.prim_func(private=True) - def add_x_x(arg0: T.Buffer((16,), "float32"), output: T.Buffer((16,), "float32")): + def add_x_x(arg0: T.Tensor((16,), "float32"), output: T.Tensor((16,), "float32")): T.func_attr({"operator_name": "relax.mul_by_2"}) for ax0 in range(16): with Ts.sblock("T_add"): @@ -169,7 +169,7 @@ def test_multiple_outputs(): @I.ir_module class Before: @Ts.prim_func(private=True) - def some_op(arg0: T.Buffer((16,), "float32"), arg1: T.Buffer((16,), "float32"), output0: T.Buffer((16,), "float32"), output1: T.Buffer((16,), "float32")): + def some_op(arg0: T.Tensor((16,), "float32"), arg1: T.Tensor((16,), "float32"), output0: T.Tensor((16,), "float32"), output1: T.Tensor((16,), "float32")): T.func_attr({"operator_name": "relax.some_op"}) for ax0 in range(16): with Ts.sblock("T_add"): @@ -189,7 +189,7 @@ def main(x: R.Tensor((16,), dtype="float32"), y: R.Tensor((16,), dtype="float32" @I.ir_module class Expected: @Ts.prim_func(private=True) - def relax_some_op_replacement(arg0: T.Buffer((4, 4), "float32"), arg1: T.Buffer((4, 4), "float32"), output0: T.Buffer((4, 4), "float32"), output1: T.Buffer((4, 4), "float32")): + def relax_some_op_replacement(arg0: T.Tensor((4, 4), "float32"), arg1: T.Tensor((4, 4), "float32"), output0: T.Tensor((4, 4), "float32"), output1: T.Tensor((4, 4), "float32")): T.func_attr({"operator_name": "relax.some_op"}) for ax0, ax1 in T.grid(4, 4): with Ts.sblock("T_add"): @@ -214,7 +214,7 @@ def main(x: R.Tensor((16,), dtype="float32"), y: R.Tensor((16,), dtype="float32" return gv @Ts.prim_func(private=True) - def some_op_2d(arg0: T.Buffer((4, 4), "float32"), arg1: T.Buffer((4, 4), "float32"), output0: T.Buffer((4, 4), "float32"), output1: T.Buffer((4, 4), "float32")): + def some_op_2d(arg0: T.Tensor((4, 4), "float32"), arg1: T.Tensor((4, 4), "float32"), output0: T.Tensor((4, 4), "float32"), output1: T.Tensor((4, 4), "float32")): for ax0, ax1 in T.grid(4, 4): with Ts.sblock("T_add"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) @@ -246,7 +246,7 @@ def foo(x: R.Tensor((14,), dtype="float32")) -> R.Tensor((14,), dtype="float32") return gv @Ts.prim_func(private=True) - def relu(arg0: T.Buffer((14,), "float32"), output: T.Buffer((14,), "float32")): + def relu(arg0: T.Tensor((14,), "float32"), output: T.Tensor((14,), "float32")): T.func_attr({"operator_name": "relax.relu"}) for ax0 in T.grid(14): with Ts.sblock("T_add"): @@ -287,7 +287,7 @@ def foo(x: R.Tensor((14,), dtype="float32")) -> R.Tensor((14,), dtype="float32") @Ts.prim_func(private=True) def relax_relu_replacement( - arg0: T.Buffer((16,), "float32"), output: T.Buffer((16,), "float32") + arg0: T.Tensor((16,), "float32"), output: T.Tensor((16,), "float32") ): T.func_attr({"operator_name": "relax.relu"}) # with Ts.sblock("root"): @@ -299,7 +299,7 @@ def relax_relu_replacement( output[v_ax0] = T.max(arg0[v_ax0], T.float32(0)) @Ts.prim_func(private=True) - def remove_pad(input: T.Buffer((p0,)), output: T.Buffer((i0,))): + def remove_pad(input: T.Tensor((p0,)), output: T.Tensor((i0,))): T.func_attr({"operator_name": "remove_pad", "tirx.noalias": True}) # with Ts.sblock("root"): @@ -311,7 +311,7 @@ def remove_pad(input: T.Buffer((p0,)), output: T.Buffer((i0,))): output[v_ax0] = input[v_ax0] @Ts.prim_func(private=True) - def relu_pad(arg0: T.Buffer((16,), "float32"), output: T.Buffer((16,), "float32")): + def relu_pad(arg0: T.Tensor((16,), "float32"), output: T.Tensor((16,), "float32")): for ax0 in T.grid(16): with Ts.sblock("T_add"): v_ax0 = Ts.axis.remap("S", [ax0]) @@ -336,7 +336,7 @@ def test_multiple_call_sites(): @I.ir_module class Before: @Ts.prim_func(private=True) - def add(arg0: T.Buffer((16,), "float32"), arg1: T.Buffer((16,), "float32"), output: T.Buffer((16,), "float32")): + def add(arg0: T.Tensor((16,), "float32"), arg1: T.Tensor((16,), "float32"), output: T.Tensor((16,), "float32")): T.func_attr({"operator_name": "relax.add"}) for ax0 in range(16): with Ts.sblock("T_add"): @@ -357,7 +357,7 @@ def main(x: R.Tensor((16,), dtype="float32"), y: R.Tensor((16,), dtype="float32" @I.ir_module class Expected: @Ts.prim_func(private=True) - def relax_add_replacement(arg0: T.Buffer((4, 4), "float32"), arg1: T.Buffer((4, 4), "float32"), output: T.Buffer((4, 4), "float32")): + def relax_add_replacement(arg0: T.Tensor((4, 4), "float32"), arg1: T.Tensor((4, 4), "float32"), output: T.Tensor((4, 4), "float32")): T.func_attr({"operator_name": "relax.add"}) # with Ts.sblock("root"): for ax0, ax1 in T.grid(4, 4): @@ -383,7 +383,7 @@ def main(x: R.Tensor((16,), dtype="float32"), y: R.Tensor((16,), dtype="float32" R.output(gv) return gv @Ts.prim_func(private=True) - def add_2d(arg0: T.Buffer((4, 4), "float32"), arg1: T.Buffer((4, 4), "float32"), output: T.Buffer((4, 4), "float32")): + def add_2d(arg0: T.Tensor((4, 4), "float32"), arg1: T.Tensor((4, 4), "float32"), output: T.Tensor((4, 4), "float32")): for ax0, ax1 in T.grid(4, 4): with Ts.sblock("T_add"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) @@ -406,8 +406,8 @@ def test_reshape(): class Before: @Ts.prim_func(private=True) def reshape( - A: T.Buffer((T.int64(850), T.int64(2048)), "float16"), - T_reshape: T.Buffer((T.int64(850), T.int64(1), T.int64(2048)), "float16"), + A: T.Tensor((T.int64(850), T.int64(2048)), "float16"), + T_reshape: T.Tensor((T.int64(850), T.int64(1), T.int64(2048)), "float16"), ): T.func_attr({"operator_name": "relax.reshape"}) for ax0, ax1, ax2 in T.grid(T.int64(850), T.int64(1), T.int64(2048)): @@ -440,8 +440,8 @@ def main(x: R.Tensor((850, 2048), dtype="float16")) -> R.Tensor( class Expected: @Ts.prim_func(private=True) def relax_reshape_replacement( - A: T.Buffer((T.int64(850), T.int64(2), T.int64(1024)), "float16"), - T_reshape: T.Buffer((T.int64(850), T.int64(1), T.int64(2048)), "float16"), + A: T.Tensor((T.int64(850), T.int64(2), T.int64(1024)), "float16"), + T_reshape: T.Tensor((T.int64(850), T.int64(1), T.int64(2048)), "float16"), ): T.func_attr({"operator_name": "relax.reshape"}) for ax0, ax1, ax2 in T.grid(T.int64(850), T.int64(1), T.int64(2048)): @@ -475,8 +475,8 @@ def main(x: R.Tensor((850, 2048), dtype="float16")) -> R.Tensor( @Ts.prim_func(private=True) def reshape_new( - A: T.Buffer((T.int64(850), T.int64(2), T.int64(1024)), "float16"), - T_reshape: T.Buffer((T.int64(850), T.int64(1), T.int64(2048)), "float16"), + A: T.Tensor((T.int64(850), T.int64(2), T.int64(1024)), "float16"), + T_reshape: T.Tensor((T.int64(850), T.int64(1), T.int64(2048)), "float16"), ): for ax0, ax1, ax2 in T.grid(T.int64(850), T.int64(1), T.int64(2048)): with Ts.sblock("T_reshape"): 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 82b1aa6c6a08..8473b7bf0894 100644 --- a/tests/python/relax/test_transform_annotate_tir_op_pattern.py +++ b/tests/python/relax/test_transform_annotate_tir_op_pattern.py @@ -46,7 +46,7 @@ def test_annotate_opkind_outewisefusable(): @tvm.script.ir_module class InputModule: @Ts.prim_func - def tir_matmul(A: T.Buffer((m, n)), B: T.Buffer((n, k)), C: T.Buffer((m, k))) -> None: + def tir_matmul(A: T.Tensor((m, n)), B: T.Tensor((n, k)), C: T.Tensor((m, k))) -> None: T.func_attr({"global_symbol": "tir_matmul"}) for i, j, k_index in T.grid(m, k, n): @@ -78,9 +78,9 @@ def test_annotate_opkind_outewisefusable_with_cast(cast_pattern): class InputModule: @Ts.prim_func def tir_matmul( - A: T.Buffer((m, n), "float16"), - B: T.Buffer((n, k), "float16"), - C: T.Buffer((m, k), "float32"), + A: T.Tensor((m, n), "float16"), + B: T.Tensor((n, k), "float16"), + C: T.Tensor((m, k), "float32"), ) -> None: T.func_attr({"global_symbol": "tir_matmul"}) @@ -101,9 +101,9 @@ def test_annotate_opkind_outewisefusable_int_var_signature(): class InputModule: @Ts.prim_func def tir_matmul( - A: T.Buffer((m, n)), # noqa: F821 - B: T.Buffer((n, k)), # noqa: F821 - C: T.Buffer((m, k)), # noqa: F821 + A: T.Tensor((m, n)), # noqa: F821 + B: T.Tensor((n, k)), # noqa: F821 + C: T.Tensor((m, k)), # noqa: F821 m: T.int64, n: T.int64, k: T.int64, @@ -126,7 +126,7 @@ def test_annotate_opkind_reduce(): @tvm.script.ir_module class InputModule: @Ts.prim_func - def sum(A: T.Buffer((16, 16)), B: T.Buffer((16,))) -> None: + def sum(A: T.Tensor((16, 16)), B: T.Tensor((16,))) -> None: T.func_attr({"global_symbol": "elemwise"}) for i, j in T.grid(16, 16): @@ -145,7 +145,7 @@ def test_annotate_opkind_ewise(): @tvm.script.ir_module class InputModule: @Ts.prim_func - def elemwise(A: T.Buffer((16, 16)), B: T.Buffer((16, 16))) -> None: + def elemwise(A: T.Tensor((16, 16)), B: T.Tensor((16, 16))) -> None: T.func_attr({"global_symbol": "elemwise"}) for i, j in T.grid(16, 16): @@ -162,7 +162,7 @@ def test_annotate_opkind_broadcast(): @tvm.script.ir_module class InputModule: @Ts.prim_func - def broadcast(A: T.Buffer((16, 16)), B: T.Buffer((16, 16, 16, 16))) -> None: + def broadcast(A: T.Tensor((16, 16)), B: T.Tensor((16, 16, 16, 16))) -> None: T.func_attr({"global_symbol": "elemwise"}) for i0, j0, i1, j1 in T.grid(16, 16, 16, 16): @@ -179,7 +179,7 @@ def test_annotate_opkind_injective(): @tvm.script.ir_module class InputModule: @Ts.prim_func - def injective(A: T.Buffer((4, 4, 4, 4)), B: T.Buffer((16, 16))) -> None: + def injective(A: T.Tensor((4, 4, 4, 4)), B: T.Tensor((16, 16))) -> None: T.func_attr({"global_symbol": "elemwise"}) for i, j in T.grid(16, 16): @@ -197,9 +197,9 @@ def test_annotate_opkind_bias_add(): class InputModule: @Ts.prim_func def tir_bias_add( - A: T.Buffer((1, 1000), "float32"), - B: T.Buffer((1000,), "float32"), - C: T.Buffer((1, 1000), "float32"), + A: T.Tensor((1, 1000), "float32"), + B: T.Tensor((1000,), "float32"), + C: T.Tensor((1, 1000), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "tir_bias_add", "tirx.noalias": True}) @@ -222,9 +222,9 @@ def test_annotate_opkind_add_broadcast_with_unit_shape(): class InputModule: @Ts.prim_func def add_with_unit_dim_len_broadcast( - A: T.Buffer((1, 64, 112, 112), "float32"), - B: T.Buffer((64, 1, 1), "float32"), - C: T.Buffer((1, 64, 112, 112), "float32"), + A: T.Tensor((1, 64, 112, 112), "float32"), + B: T.Tensor((64, 1, 1), "float32"), + C: T.Tensor((1, 64, 112, 112), "float32"), ) -> None: T.func_attr({"global_symbol": "add5", "tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(1, 64, 112, 112): @@ -244,9 +244,9 @@ def test_annotate_opkind_add_zero_dim_element_wise(): class InputModule: @Ts.prim_func def add_zero_dim( - A: T.Buffer((128,), "float32"), - B: T.Buffer((), "float32"), - C: T.Buffer((128,), "float32"), + A: T.Tensor((128,), "float32"), + B: T.Tensor((), "float32"), + C: T.Tensor((128,), "float32"), ) -> None: T.func_attr({"global_symbol": "add8", "tirx.noalias": True}) for i0 in T.serial(128): @@ -266,8 +266,8 @@ def test_annotate_opkind_pooling(): class InputModule: @Ts.prim_func def max_pool2d( - rxplaceholder_1: T.Buffer((1, 64, 112, 112), "float32"), - tensor_1: T.Buffer((1, 64, 56, 56), "float32"), + rxplaceholder_1: T.Tensor((1, 64, 112, 112), "float32"), + tensor_1: T.Tensor((1, 64, 56, 56), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "max_pool2d", "T.noalias": True}) @@ -309,8 +309,8 @@ def test_annotate_opkind_softmax(): class InputModule: @Ts.prim_func def softmax( - rxplaceholder_1: T.Buffer((16, 16), "float32"), - T_softmax_norm_1: T.Buffer((16, 16), "float32"), + rxplaceholder_1: T.Tensor((16, 16), "float32"), + T_softmax_norm_1: T.Tensor((16, 16), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "softmax", "T.noalias": True}) @@ -367,8 +367,8 @@ def test_multiple_bufer_stores_fallback(): class CumsumModule: @Ts.prim_func def cumsum( - rxplaceholder: T.Buffer([10, 16], dtype="float32", offset_factor=1), - out_buf: T.Buffer(160, "float32"), + rxplaceholder: T.Tensor([10, 16], dtype="float32", offset_factor=1), + out_buf: T.Tensor(160, "float32"), ): with Ts.sblock("cumsum_generic"): Ts.reads(rxplaceholder[0:10, 0:16]) @@ -394,9 +394,9 @@ def test_sum_sqsum(): class Module: @Ts.prim_func def sum_sqsum( - A: T.Buffer((32, 64), "float32"), - vsum: T.Buffer((32,), "float32"), - sqsum: T.Buffer((32,), "float32"), + A: T.Tensor((32, 64), "float32"), + vsum: T.Tensor((32,), "float32"), + sqsum: T.Tensor((32,), "float32"), ): for ax0, k0 in T.grid(32, 64): with Ts.sblock("block"): @@ -420,7 +420,7 @@ def test_no_buffer_stores(): @tvm.script.ir_module class Module: @Ts.prim_func - def no_buffer_stores(A: T.Buffer((32, 64), "float32"), vsum: T.Buffer((32,), "float32")): + def no_buffer_stores(A: T.Tensor((32, 64), "float32"), vsum: T.Tensor((32,), "float32")): for ax0, k0 in T.grid(32, 64): with Ts.sblock("block"): v_ax0, v_k0 = Ts.axis.remap("SR", [ax0, k0]) diff --git a/tests/python/relax/test_transform_attach_attr_layout_free_buffers.py b/tests/python/relax/test_transform_attach_attr_layout_free_buffers.py index 61525719d232..be9a62a3e498 100644 --- a/tests/python/relax/test_transform_attach_attr_layout_free_buffers.py +++ b/tests/python/relax/test_transform_attach_attr_layout_free_buffers.py @@ -33,9 +33,9 @@ def test_param(): class Before: @Ts.prim_func(private=True) def matmul( - A: T.Buffer((T.int64(32), T.int64(32)), "float32"), - B: T.Buffer((T.int64(32), T.int64(32)), "float32"), - C: T.Buffer((T.int64(32), T.int64(32)), "float32"), + A: T.Tensor((T.int64(32), T.int64(32)), "float32"), + B: T.Tensor((T.int64(32), T.int64(32)), "float32"), + C: T.Tensor((T.int64(32), T.int64(32)), "float32"), ): for i, j, k in T.grid(T.int64(32), T.int64(32), T.int64(32)): with Ts.sblock("C"): @@ -56,9 +56,9 @@ def main(x: R.Tensor((32, 32), "float32"), y: R.Tensor((32, 32), "float32")): class Expected: @Ts.prim_func(private=True) def matmul1( - A: T.Buffer((T.int64(32), T.int64(32)), "float32"), - B: T.Buffer((T.int64(32), T.int64(32)), "float32"), - C: T.Buffer((T.int64(32), T.int64(32)), "float32"), + A: T.Tensor((T.int64(32), T.int64(32)), "float32"), + B: T.Tensor((T.int64(32), T.int64(32)), "float32"), + C: T.Tensor((T.int64(32), T.int64(32)), "float32"), ): T.func_attr({"layout_free_buffers": [1]}) for i, j, k in T.grid(T.int64(32), T.int64(32), T.int64(32)): @@ -87,9 +87,9 @@ def test_const(): class Before: @Ts.prim_func(private=True) def matmul( - A: T.Buffer((T.int64(32), T.int64(32)), "float32"), - B: T.Buffer((T.int64(32), T.int64(32)), "float32"), - C: T.Buffer((T.int64(32), T.int64(32)), "float32"), + A: T.Tensor((T.int64(32), T.int64(32)), "float32"), + B: T.Tensor((T.int64(32), T.int64(32)), "float32"), + C: T.Tensor((T.int64(32), T.int64(32)), "float32"), ): for i, j, k in T.grid(T.int64(32), T.int64(32), T.int64(32)): with Ts.sblock("C"): @@ -114,9 +114,9 @@ def main(x: R.Tensor((32, 32), "float32")): class Expected: @Ts.prim_func(private=True) def matmul1( - A: T.Buffer((T.int64(32), T.int64(32)), "float32"), - B: T.Buffer((T.int64(32), T.int64(32)), "float32"), - C: T.Buffer((T.int64(32), T.int64(32)), "float32"), + A: T.Tensor((T.int64(32), T.int64(32)), "float32"), + B: T.Tensor((T.int64(32), T.int64(32)), "float32"), + C: T.Tensor((T.int64(32), T.int64(32)), "float32"), ): T.func_attr({"layout_free_buffers": [1]}) for i, j, k in T.grid(T.int64(32), T.int64(32), T.int64(32)): @@ -147,9 +147,9 @@ def test_multiple_same_func(): class Before: @Ts.prim_func(private=True) def matmul( - A: T.Buffer((T.int64(32), T.int64(32)), "float32"), - B: T.Buffer((T.int64(32), T.int64(32)), "float32"), - C: T.Buffer((T.int64(32), T.int64(32)), "float32"), + A: T.Tensor((T.int64(32), T.int64(32)), "float32"), + B: T.Tensor((T.int64(32), T.int64(32)), "float32"), + C: T.Tensor((T.int64(32), T.int64(32)), "float32"), ): for i, j, k in T.grid(T.int64(32), T.int64(32), T.int64(32)): with Ts.sblock("C"): @@ -183,9 +183,9 @@ def main( class Expected: @Ts.prim_func(private=True) def matmul1( - A: T.Buffer((T.int64(32), T.int64(32)), "float32"), - B: T.Buffer((T.int64(32), T.int64(32)), "float32"), - C: T.Buffer((T.int64(32), T.int64(32)), "float32"), + A: T.Tensor((T.int64(32), T.int64(32)), "float32"), + B: T.Tensor((T.int64(32), T.int64(32)), "float32"), + C: T.Tensor((T.int64(32), T.int64(32)), "float32"), ): T.func_attr({"layout_free_buffers": [1]}) for i, j, k in T.grid(T.int64(32), T.int64(32), T.int64(32)): @@ -225,9 +225,9 @@ def test_multiple_same_func_with_different_free_buffers(): class Before: @Ts.prim_func(private=True) def matmul( - A: T.Buffer((T.int64(32), T.int64(32)), "float32"), - B: T.Buffer((T.int64(32), T.int64(32)), "float32"), - C: T.Buffer((T.int64(32), T.int64(32)), "float32"), + A: T.Tensor((T.int64(32), T.int64(32)), "float32"), + B: T.Tensor((T.int64(32), T.int64(32)), "float32"), + C: T.Tensor((T.int64(32), T.int64(32)), "float32"), ): for i, j, k in T.grid(T.int64(32), T.int64(32), T.int64(32)): with Ts.sblock("C"): @@ -261,9 +261,9 @@ def main( class Expected: @Ts.prim_func(private=True) def matmul1( - A: T.Buffer((T.int64(32), T.int64(32)), "float32"), - B: T.Buffer((T.int64(32), T.int64(32)), "float32"), - C: T.Buffer((T.int64(32), T.int64(32)), "float32"), + A: T.Tensor((T.int64(32), T.int64(32)), "float32"), + B: T.Tensor((T.int64(32), T.int64(32)), "float32"), + C: T.Tensor((T.int64(32), T.int64(32)), "float32"), ): T.func_attr({"layout_free_buffers": [1]}) for i, j, k in T.grid(T.int64(32), T.int64(32), T.int64(32)): @@ -274,9 +274,9 @@ def matmul1( @Ts.prim_func(private=True) def matmul2( - A: T.Buffer((T.int64(32), T.int64(32)), "float32"), - B: T.Buffer((T.int64(32), T.int64(32)), "float32"), - C: T.Buffer((T.int64(32), T.int64(32)), "float32"), + A: T.Tensor((T.int64(32), T.int64(32)), "float32"), + B: T.Tensor((T.int64(32), T.int64(32)), "float32"), + C: T.Tensor((T.int64(32), T.int64(32)), "float32"), ): T.func_attr({"layout_free_buffers": [0]}) for i, j, k in T.grid(T.int64(32), T.int64(32), T.int64(32)): diff --git a/tests/python/relax/test_transform_attach_global_symbol.py b/tests/python/relax/test_transform_attach_global_symbol.py index d638732551e1..13c28be4ab0b 100644 --- a/tests/python/relax/test_transform_attach_global_symbol.py +++ b/tests/python/relax/test_transform_attach_global_symbol.py @@ -40,9 +40,9 @@ def test_basic(): class Before: @Ts.prim_func def tir_matmul( - A: T.Buffer((m_tir_matmul, n_tir_matmul)), - B: T.Buffer((n_tir_matmul, k_tir_matmul)), - C: T.Buffer((m_tir_matmul, k_tir_matmul)), + A: T.Tensor((m_tir_matmul, n_tir_matmul)), + B: T.Tensor((n_tir_matmul, k_tir_matmul)), + C: T.Tensor((m_tir_matmul, k_tir_matmul)), ) -> None: for i, j, k_tir_matmul_index in T.grid(m_tir_matmul, k_tir_matmul, n_tir_matmul): with Ts.sblock("matmul"): @@ -69,9 +69,9 @@ def main( class Expected: @Ts.prim_func def tir_matmul( - A: T.Buffer((m_tir_matmul, n_tir_matmul)), - B: T.Buffer((n_tir_matmul, k_tir_matmul)), - C: T.Buffer((m_tir_matmul, k_tir_matmul)), + A: T.Tensor((m_tir_matmul, n_tir_matmul)), + B: T.Tensor((n_tir_matmul, k_tir_matmul)), + C: T.Tensor((m_tir_matmul, k_tir_matmul)), ) -> None: T.func_attr({"global_symbol": "tir_matmul"}) @@ -103,7 +103,7 @@ class Before: I.module_attrs({"system_lib_prefix": "hello_"}) @Ts.prim_func(private=True) - def tir_zeros(x: T.Buffer((2), "float32")) -> None: + def tir_zeros(x: T.Tensor((2), "float32")) -> None: x[0] = T.float32(0) @R.function(private=True) @@ -116,7 +116,7 @@ class Expected: I.module_attrs({"system_lib_prefix": "hello_"}) @Ts.prim_func - def hello_tir_zeros(x: T.Buffer((2), "float32")) -> None: + def hello_tir_zeros(x: T.Tensor((2), "float32")) -> None: T.func_attr({"global_symbol": "hello_tir_zeros"}) x[0] = T.float32(0) diff --git a/tests/python/relax/test_transform_bind_params.py b/tests/python/relax/test_transform_bind_params.py index d499f993004f..0dc684267ef3 100644 --- a/tests/python/relax/test_transform_bind_params.py +++ b/tests/python/relax/test_transform_bind_params.py @@ -33,7 +33,7 @@ def test_bind_params(use_np_array): @tvm.script.ir_module class InputModule: @Ts.prim_func - def tir_matmul(A: T.Buffer((16, 16)), B: T.Buffer((16, 16)), C: T.Buffer((16, 16))) -> None: + def tir_matmul(A: T.Tensor((16, 16)), B: T.Tensor((16, 16)), C: T.Tensor((16, 16))) -> None: T.func_attr({"global_symbol": "tir_matmul"}) for i0, j, k0, i1, k1 in T.grid(4, 16, 4, 4, 4): diff --git a/tests/python/relax/test_transform_codegen_pass.py b/tests/python/relax/test_transform_codegen_pass.py index 166331f61ab1..784ebb73cb19 100644 --- a/tests/python/relax/test_transform_codegen_pass.py +++ b/tests/python/relax/test_transform_codegen_pass.py @@ -395,7 +395,7 @@ def main(x: R.Tensor([4], "int64")): return x @Ts.prim_func(private=True) - def shape_func(H: T.Buffer(T.int64(4), "int64")): + def shape_func(H: T.Tensor(T.int64(4), "int64")): H[T.int64(0)] = H[T.int64(0)] + T.int64(1) Expected = Before diff --git a/tests/python/relax/test_transform_cse.py b/tests/python/relax/test_transform_cse.py index 643c5784a10b..fd787238e552 100644 --- a/tests/python/relax/test_transform_cse.py +++ b/tests/python/relax/test_transform_cse.py @@ -402,9 +402,9 @@ def main(A: R.Tensor([16, 16], "int32"), B: R.Tensor([16, 16], "int32")): @Ts.prim_func(private=True) def product( - A: T.Buffer([16, 16], "int32"), - B: T.Buffer([16, 16], "int32"), - C: T.Buffer([16, 16], "int32"), + A: T.Tensor([16, 16], "int32"), + B: T.Tensor([16, 16], "int32"), + C: T.Tensor([16, 16], "int32"), ): for (*iters,) in T.grid(*A.shape): with Ts.sblock("compute"): @@ -413,9 +413,9 @@ def product( @Ts.prim_func(private=True) def sum( - A: T.Buffer([16, 16], "int32"), - B: T.Buffer([16, 16], "int32"), - C: T.Buffer([16, 16], "int32"), + A: T.Tensor([16, 16], "int32"), + B: T.Tensor([16, 16], "int32"), + C: T.Tensor([16, 16], "int32"), ): for (*iters,) in T.grid(*A.shape): with Ts.sblock("compute"): diff --git a/tests/python/relax/test_transform_dead_code_elimination.py b/tests/python/relax/test_transform_dead_code_elimination.py index a201cf6b12e7..aac33b48aab0 100644 --- a/tests/python/relax/test_transform_dead_code_elimination.py +++ b/tests/python/relax/test_transform_dead_code_elimination.py @@ -165,9 +165,9 @@ def test_unused_relax_func(): class InputModule: @Ts.prim_func def tir_add( - x: T.Buffer((16, 16), "float32"), - y: T.Buffer((16, 16), "float32"), - z: T.Buffer((16, 16), "float32"), + x: T.Tensor((16, 16), "float32"), + y: T.Tensor((16, 16), "float32"), + z: T.Tensor((16, 16), "float32"), ) -> None: for i, j in T.grid(16, 16): with Ts.sblock("add"): @@ -202,9 +202,9 @@ def test_unused_relax_func_custom_entry_func(provide_entry_func_name): class InputModule: @Ts.prim_func(private=True) def tir_add( - x: T.Buffer((16, 16), "float32"), - y: T.Buffer((16, 16), "float32"), - z: T.Buffer((16, 16), "float32"), + x: T.Tensor((16, 16), "float32"), + y: T.Tensor((16, 16), "float32"), + z: T.Tensor((16, 16), "float32"), ) -> None: for i, j in T.grid(16, 16): with Ts.sblock("add"): @@ -243,9 +243,9 @@ def test_tracking_through_externally_exposed_func(provide_entry_func_name): class InputModule: @Ts.prim_func(private=True) def tir_add( - x: T.Buffer((16, 16), "float32"), - y: T.Buffer((16, 16), "float32"), - z: T.Buffer((16, 16), "float32"), + x: T.Tensor((16, 16), "float32"), + y: T.Tensor((16, 16), "float32"), + z: T.Tensor((16, 16), "float32"), ) -> None: for i, j in T.grid(16, 16): with Ts.sblock("add"): @@ -295,9 +295,9 @@ def test_unused_relax_func_symbolic_shape(): class InputModule: @Ts.prim_func def tir_matmul( - x: T.Buffer((m_tir_matmul, n_tir_matmul), "float32"), - y: T.Buffer((n_tir_matmul, k_tir_matmul), "float32"), - z: T.Buffer((m_tir_matmul, k_tir_matmul), "float32"), + x: T.Tensor((m_tir_matmul, n_tir_matmul), "float32"), + y: T.Tensor((n_tir_matmul, k_tir_matmul), "float32"), + z: T.Tensor((m_tir_matmul, k_tir_matmul), "float32"), ) -> None: for i, j, k_tir_matmul_index in T.grid(m_tir_matmul, k_tir_matmul, n_tir_matmul): with Ts.sblock("matmul"): @@ -337,9 +337,9 @@ def test_unused_prim_func(): class InputModule: @Ts.prim_func def unused_func( - x: T.Buffer((16, 16), "float32"), - y: T.Buffer((16, 16), "float32"), - z: T.Buffer((16, 16), "float32"), + x: T.Tensor((16, 16), "float32"), + y: T.Tensor((16, 16), "float32"), + z: T.Tensor((16, 16), "float32"), ) -> None: T.func_attr({"global_symbol": "tir_unused"}) for i, j in T.grid(16, 16): @@ -384,9 +384,9 @@ def main(x: R.Tensor((16, 16), "float32"), w: R.Tensor((16, 16), "float32")) -> @Ts.prim_func(private=True) def tir_add_tensors( - x: T.Buffer((16, 16), "float32"), - y: T.Buffer((16, 16), "float32"), - z: T.Buffer((16, 16), "float32"), + x: T.Tensor((16, 16), "float32"), + y: T.Tensor((16, 16), "float32"), + z: T.Tensor((16, 16), "float32"), ): for i, j in T.grid(16, 16): with Ts.sblock("add"): @@ -409,9 +409,9 @@ def test_multiple_unused_funcs(): class InputModule: @Ts.prim_func def unused_func1( - x: T.Buffer((16, 16), "float32"), - y: T.Buffer((16, 16), "float32"), - z: T.Buffer((16, 16), "float32"), + x: T.Tensor((16, 16), "float32"), + y: T.Tensor((16, 16), "float32"), + z: T.Tensor((16, 16), "float32"), ) -> None: T.func_attr({"global_symbol": "tir_unused"}) for i, j in T.grid(16, 16): diff --git a/tests/python/relax/test_transform_fold_constant.py b/tests/python/relax/test_transform_fold_constant.py index 810873ce4af4..20c6b7dc40bf 100644 --- a/tests/python/relax/test_transform_fold_constant.py +++ b/tests/python/relax/test_transform_fold_constant.py @@ -63,7 +63,7 @@ def test_one_fold_addone(): @tvm.script.ir_module class Module: @Ts.prim_func - def addone(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")) -> None: + def addone(A: T.Tensor((16, 16), "float32"), B: T.Tensor((16, 16), "float32")) -> None: for i, j in T.grid(16, 16): with Ts.sblock("addone"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -93,7 +93,7 @@ def test_one_fold_transpose(): @tvm.script.ir_module class Module: @Ts.prim_func - def func(A: T.Buffer((2, 3), "float32"), B: T.Buffer((3, 2), "float32")) -> None: + def func(A: T.Tensor((2, 3), "float32"), B: T.Tensor((3, 2), "float32")) -> None: for i, j in T.grid(3, 2): with Ts.sblock("transpose"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -122,7 +122,7 @@ def test_two_hop_addone(): @tvm.script.ir_module class Module: @Ts.prim_func - def addone(A: T.Buffer((2, 2), "float32"), B: T.Buffer((2, 2), "float32")) -> None: + def addone(A: T.Tensor((2, 2), "float32"), B: T.Tensor((2, 2), "float32")) -> None: for i, j in T.grid(2, 2): with Ts.sblock("addone"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -153,7 +153,7 @@ def test_dataflow_fold(): @tvm.script.ir_module class Module: @Ts.prim_func - def identity(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")) -> None: + def identity(A: T.Tensor((16, 16), "float32"), B: T.Tensor((16, 16), "float32")) -> None: for i, j in T.grid(16, 16): with Ts.sblock("identity"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -191,7 +191,7 @@ def test_fold_mixed_case(): class Module: # TIR function can handle different cases. @Ts.prim_func - def addone(A: T.Buffer((n_addone, m_addone)), B: T.Buffer((n_addone, m_addone))) -> None: + def addone(A: T.Tensor((n_addone, m_addone)), B: T.Tensor((n_addone, m_addone))) -> None: for i, j in T.grid(n_addone, m_addone): with Ts.sblock("addone"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -199,9 +199,9 @@ def addone(A: T.Buffer((n_addone, m_addone)), B: T.Buffer((n_addone, m_addone))) @Ts.prim_func def sub( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: for i, j in T.grid(16, 16): with Ts.sblock("sub"): @@ -251,7 +251,7 @@ def test_int32_fold(): @tvm.script.ir_module class Module: @Ts.prim_func - def addone(A: T.Buffer((16, 16), "int32"), B: T.Buffer((16, 16), "int32")) -> None: + def addone(A: T.Tensor((16, 16), "int32"), B: T.Tensor((16, 16), "int32")) -> None: for i, j in T.grid(16, 16): with Ts.sblock("addone"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -452,9 +452,9 @@ def test_fold_tuple_output(): class Module: @Ts.prim_func def split( - A: T.Buffer((4, 4), "float32"), - B: T.Buffer((2, 4), "float32"), - C: T.Buffer((2, 4), "float32"), + A: T.Tensor((4, 4), "float32"), + B: T.Tensor((2, 4), "float32"), + C: T.Tensor((2, 4), "float32"), ) -> None: for i, j in T.grid(2, 4): with Ts.sblock("upper"): @@ -563,7 +563,7 @@ def test_fold_large_op_with_tensor_input(): @tvm.script.ir_module class Module: @Ts.prim_func - def addone(A: T.Buffer((2048,), "float32"), B: T.Buffer((2048,), "float32")) -> None: + def addone(A: T.Tensor((2048,), "float32"), B: T.Tensor((2048,), "float32")) -> None: for i in range(2048): with Ts.sblock("addone"): vi = Ts.axis.remap("S", [i]) @@ -595,7 +595,7 @@ def test_call_tir_with_primitive_args_not_folded(): @tvm.script.ir_module class Module: @Ts.prim_func(private=True) - def shape_to_tensor(m: T.int64, out: T.Buffer((T.int64(1),), "int64")): + def shape_to_tensor(m: T.int64, out: T.Tensor((T.int64(1),), "int64")): for i in range(T.int64(1)): with Ts.sblock("out"): vi = Ts.axis.remap("S", [i]) diff --git a/tests/python/relax/test_transform_fuse_ops.py b/tests/python/relax/test_transform_fuse_ops.py index d4afb8b0ad52..3ca2d32ea82e 100644 --- a/tests/python/relax/test_transform_fuse_ops.py +++ b/tests/python/relax/test_transform_fuse_ops.py @@ -870,7 +870,7 @@ def main(x: R.Tensor((2, 3), "float32")): return R.tuple(b, c) @Ts.prim_func(private=True) - def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")): + def exp(A: T.Tensor((2, 3), "float32"), B: T.Tensor((2, 3), "float32")): T.evaluate(0) # FuseOps should does no change to it. @@ -891,7 +891,7 @@ def main(x: R.Tensor((1, 512, 64, 64), "float32"), mean: R.Tensor((64, 64), "flo return gv1 @Ts.prim_func(private=True) - def layer_norm(A: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32"), gamma: T.Buffer((T.int64(64), T.int64(64)), "float32"), beta: T.Buffer((T.int64(64), T.int64(64)), "float32"), T_layer_norm: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32")): + def layer_norm(A: T.Tensor((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32"), gamma: T.Tensor((T.int64(64), T.int64(64)), "float32"), beta: T.Tensor((T.int64(64), T.int64(64)), "float32"), T_layer_norm: T.Tensor((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32")): rxplaceholder_red_temp_v0 = Ts.sblock_alloc_buffer([T.int64(64), T.int64(64)], dtype="float32") rxplaceholder_red_temp_v1 = Ts.sblock_alloc_buffer([T.int64(64), T.int64(64)], dtype="float32") for i0, i1, i2, i3 in T.grid(T.int64(1), T.int64(512), T.int64(64), T.int64(64)): @@ -914,7 +914,7 @@ def layer_norm(A: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), T_layer_norm[ax0, ax1, ax2, ax3] = (A[ax0, ax1, ax2, ax3] - rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.05)) * T.rsqrt(rxplaceholder_red_temp_v1[ax0, ax1] * T.float32(0.05) - rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.05) * (rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.05)) + T.float32(1e-05)) * gamma[ax2, ax3] + beta[ax2, ax3] @Ts.prim_func(private=True) - def relu(A: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32"), B: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32")): + def relu(A: T.Tensor((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32"), B: T.Tensor((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32")): for i0, i1, i2, i3 in T.grid(T.int64(1), T.int64(512), T.int64(64), T.int64(64)): with Ts.sblock("relu"): v_i0, v_i1, v_i2, v_i3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) @@ -925,7 +925,7 @@ def relu(A: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "floa @I.ir_module class Expected: @Ts.prim_func(private=True) - def layer_norm(A: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32"), gamma: T.Buffer((T.int64(64), T.int64(64)), "float32"), beta: T.Buffer((T.int64(64), T.int64(64)), "float32"), T_layer_norm: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32")): + def layer_norm(A: T.Tensor((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32"), gamma: T.Tensor((T.int64(64), T.int64(64)), "float32"), beta: T.Tensor((T.int64(64), T.int64(64)), "float32"), T_layer_norm: T.Tensor((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32")): T.func_attr({"op_pattern": 4}) # with Ts.sblock("root"): rxplaceholder_red_temp_v0 = Ts.sblock_alloc_buffer((T.int64(64), T.int64(64))) @@ -950,7 +950,7 @@ def layer_norm(A: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), T_layer_norm[ax0, ax1, ax2, ax3] = (A[ax0, ax1, ax2, ax3] - rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.050000000000000003)) * T.rsqrt(rxplaceholder_red_temp_v1[ax0, ax1] * T.float32(0.050000000000000003) - rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.050000000000000003) * (rxplaceholder_red_temp_v0[ax0, ax1] * T.float32(0.050000000000000003)) + T.float32(1.0000000000000001e-05)) * gamma[ax2, ax3] + beta[ax2, ax3] @Ts.prim_func(private=True) - def relu(A: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32"), B: T.Buffer((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32")): + def relu(A: T.Tensor((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32"), B: T.Tensor((T.int64(1), T.int64(512), T.int64(64), T.int64(64)), "float32")): T.func_attr({"op_pattern": 0}) # with Ts.sblock("root"): for i0, i1, i2, i3 in T.grid(T.int64(1), T.int64(512), T.int64(64), T.int64(64)): @@ -1012,7 +1012,7 @@ def main( @I.ir_module class Expected: @Ts.prim_func(private=True) - def add(rxplaceholder: T.Buffer((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32"), rxplaceholder_1: T.Buffer((T.int64(1), T.int64(320), T.int64(1), T.int64(1)), "float32"), T_add: T.Buffer((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32")): + def add(rxplaceholder: T.Tensor((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32"), rxplaceholder_1: T.Tensor((T.int64(1), T.int64(320), T.int64(1), T.int64(1)), "float32"), T_add: T.Tensor((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32")): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) for ax0, ax1, ax2, ax3 in T.grid(T.int64(2), T.int64(320), T.int64(64), T.int64(64)): with Ts.sblock("T_add"): @@ -1022,7 +1022,7 @@ def add(rxplaceholder: T.Buffer((T.int64(2), T.int64(320), T.int64(64), T.int64( T_add[v_ax0, v_ax1, v_ax2, v_ax3] = rxplaceholder[v_ax0, v_ax1, v_ax2, v_ax3] + rxplaceholder_1[T.int64(0), v_ax1, T.int64(0), T.int64(0)] @Ts.prim_func(private=True) - def add1(rxplaceholder: T.Buffer((T.int64(2), T.int64(320)), "float32"), rxplaceholder_1: T.Buffer((T.int64(320),), "float32"), T_add: T.Buffer((T.int64(2), T.int64(320)), "float32")): + def add1(rxplaceholder: T.Tensor((T.int64(2), T.int64(320)), "float32"), rxplaceholder_1: T.Tensor((T.int64(320),), "float32"), T_add: T.Tensor((T.int64(2), T.int64(320)), "float32")): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(2), T.int64(320)): with Ts.sblock("T_add"): @@ -1032,7 +1032,7 @@ def add1(rxplaceholder: T.Buffer((T.int64(2), T.int64(320)), "float32"), rxplace T_add[v_ax0, v_ax1] = rxplaceholder[v_ax0, v_ax1] + rxplaceholder_1[v_ax1] @Ts.prim_func(private=True) - def add2(rxplaceholder: T.Buffer((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32"), rxplaceholder_1: T.Buffer((T.int64(2), T.int64(320), T.int64(1), T.int64(1)), "float32"), T_add: T.Buffer((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32")): + def add2(rxplaceholder: T.Tensor((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32"), rxplaceholder_1: T.Tensor((T.int64(2), T.int64(320), T.int64(1), T.int64(1)), "float32"), T_add: T.Tensor((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32")): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) for ax0, ax1, ax2, ax3 in T.grid(T.int64(2), T.int64(320), T.int64(64), T.int64(64)): with Ts.sblock("T_add"): @@ -1042,7 +1042,7 @@ def add2(rxplaceholder: T.Buffer((T.int64(2), T.int64(320), T.int64(64), T.int64 T_add[v_ax0, v_ax1, v_ax2, v_ax3] = rxplaceholder[v_ax0, v_ax1, v_ax2, v_ax3] + rxplaceholder_1[v_ax0, v_ax1, T.int64(0), T.int64(0)] @Ts.prim_func(private=True) - def conv2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32"), rxplaceholder_1: T.Buffer((T.int64(320), T.int64(320), T.int64(3), T.int64(3)), "float32"), conv2d_nchw: T.Buffer((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32")): + def conv2d(rxplaceholder: T.Tensor((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32"), rxplaceholder_1: T.Tensor((T.int64(320), T.int64(320), T.int64(3), T.int64(3)), "float32"), conv2d_nchw: T.Tensor((T.int64(2), T.int64(320), T.int64(64), T.int64(64)), "float32")): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) pad_temp = Ts.sblock_alloc_buffer((T.int64(2), T.int64(320), T.int64(66), T.int64(66))) for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(320), T.int64(66), T.int64(66)): @@ -1061,7 +1061,7 @@ def conv2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(320), T.int64(64), T.int conv2d_nchw[v_nn, v_ff, v_yy, v_xx] = conv2d_nchw[v_nn, v_ff, v_yy, v_xx] + pad_temp[v_nn, v_rc, v_yy + v_ry, v_xx + v_rx] * rxplaceholder_1[v_ff, v_rc, v_ry, v_rx] @Ts.prim_func(private=True) - def matmul(rxplaceholder: T.Buffer((T.int64(2), T.int64(1280)), "float32"), rxplaceholder_1: T.Buffer((T.int64(1280), T.int64(320)), "float32"), matmul: T.Buffer((T.int64(2), T.int64(320)), "float32")): + def matmul(rxplaceholder: T.Tensor((T.int64(2), T.int64(1280)), "float32"), rxplaceholder_1: T.Tensor((T.int64(1280), T.int64(320)), "float32"), matmul: T.Tensor((T.int64(2), T.int64(320)), "float32")): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) for i0, i1, k in T.grid(T.int64(2), T.int64(320), T.int64(1280)): with Ts.sblock("matmul"): @@ -1073,7 +1073,7 @@ def matmul(rxplaceholder: T.Buffer((T.int64(2), T.int64(1280)), "float32"), rxpl matmul[v_i0, v_i1] = matmul[v_i0, v_i1] + rxplaceholder[v_i0, v_k] * rxplaceholder_1[v_k, v_i1] @Ts.prim_func(private=True) - def reshape(rxplaceholder: T.Buffer((T.int64(320),), "float32"), T_reshape: T.Buffer((T.int64(1), T.int64(320), T.int64(1), T.int64(1)), "float32")): + def reshape(rxplaceholder: T.Tensor((T.int64(320),), "float32"), T_reshape: T.Tensor((T.int64(1), T.int64(320), T.int64(1), T.int64(1)), "float32")): T.func_attr({"op_pattern": 2, "tirx.noalias": True}) for ax0, ax1, ax2, ax3 in T.grid(T.int64(1), T.int64(320), T.int64(1), T.int64(1)): with Ts.sblock("T_reshape"): @@ -1083,7 +1083,7 @@ def reshape(rxplaceholder: T.Buffer((T.int64(320),), "float32"), T_reshape: T.Bu T_reshape[v_ax0, v_ax1, v_ax2, v_ax3] = rxplaceholder[(v_ax1 + v_ax2 + v_ax3) % T.int64(320)] @Ts.prim_func(private=True) - def reshape1(rxplaceholder: T.Buffer((T.int64(2), T.int64(320)), "float32"), T_reshape: T.Buffer((T.int64(2), T.int64(320), T.int64(1), T.int64(1)), "float32")): + def reshape1(rxplaceholder: T.Tensor((T.int64(2), T.int64(320)), "float32"), T_reshape: T.Tensor((T.int64(2), T.int64(320), T.int64(1), T.int64(1)), "float32")): T.func_attr({"op_pattern": 2, "tirx.noalias": True}) for ax0, ax1, ax2, ax3 in T.grid(T.int64(2), T.int64(320), T.int64(1), T.int64(1)): with Ts.sblock("T_reshape"): @@ -1093,7 +1093,7 @@ def reshape1(rxplaceholder: T.Buffer((T.int64(2), T.int64(320)), "float32"), T_r T_reshape[v_ax0, v_ax1, v_ax2, v_ax3] = rxplaceholder[((v_ax1 + v_ax2 + v_ax3) // T.int64(320) + v_ax0) % T.int64(2), (v_ax1 + v_ax2 + v_ax3) % T.int64(320)] @Ts.prim_func(private=True) - def transpose(rxplaceholder: T.Buffer((T.int64(320), T.int64(1280)), "float32"), T_transpose: T.Buffer((T.int64(1280), T.int64(320)), "float32")): + def transpose(rxplaceholder: T.Tensor((T.int64(320), T.int64(1280)), "float32"), T_transpose: T.Tensor((T.int64(1280), T.int64(320)), "float32")): T.func_attr({"op_pattern": 2, "tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(1280), T.int64(320)): with Ts.sblock("T_transpose"): @@ -1167,7 +1167,7 @@ def main(inp_0: R.Tensor((1, 784), dtype="float32"), inp_1: R.Tensor((1, 128), d @I.ir_module class Expected: @Ts.prim_func(private=True) - def add(rxplaceholder: T.Buffer((T.int64(1), T.int64(128)), "float32"), rxplaceholder_1: T.Buffer((T.int64(128),), "float32"), T_add: T.Buffer((T.int64(1), T.int64(128)), "float32")): + def add(rxplaceholder: T.Tensor((T.int64(1), T.int64(128)), "float32"), rxplaceholder_1: T.Tensor((T.int64(128),), "float32"), T_add: T.Tensor((T.int64(1), T.int64(128)), "float32")): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1 in T.grid(T.int64(1), T.int64(128)): @@ -1178,7 +1178,7 @@ def add(rxplaceholder: T.Buffer((T.int64(1), T.int64(128)), "float32"), rxplaceh T_add[v_ax0, v_ax1] = rxplaceholder[v_ax0, v_ax1] + rxplaceholder_1[v_ax1] @Ts.prim_func(private=True) - def add1(rxplaceholder: T.Buffer((T.int64(1), T.int64(10)), "float32"), rxplaceholder_1: T.Buffer((T.int64(10),), "float32"), T_add: T.Buffer((T.int64(1), T.int64(10)), "float32")): + def add1(rxplaceholder: T.Tensor((T.int64(1), T.int64(10)), "float32"), rxplaceholder_1: T.Tensor((T.int64(10),), "float32"), T_add: T.Tensor((T.int64(1), T.int64(10)), "float32")): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1 in T.grid(T.int64(1), T.int64(10)): @@ -1189,7 +1189,7 @@ def add1(rxplaceholder: T.Buffer((T.int64(1), T.int64(10)), "float32"), rxplaceh T_add[v_ax0, v_ax1] = rxplaceholder[v_ax0, v_ax1] + rxplaceholder_1[v_ax1] @Ts.prim_func(private=True) - def matmul(rxplaceholder: T.Buffer((T.int64(1), T.int64(784)), "float32"), rxplaceholder_1: T.Buffer((T.int64(784), T.int64(128)), "float32"), matmul_1: T.Buffer((T.int64(1), T.int64(128)), "float32")): + def matmul(rxplaceholder: T.Tensor((T.int64(1), T.int64(784)), "float32"), rxplaceholder_1: T.Tensor((T.int64(784), T.int64(128)), "float32"), matmul_1: T.Tensor((T.int64(1), T.int64(128)), "float32")): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1, k in T.grid(T.int64(1), T.int64(128), T.int64(784)): @@ -1202,7 +1202,7 @@ def matmul(rxplaceholder: T.Buffer((T.int64(1), T.int64(784)), "float32"), rxpla matmul_1[v_i0, v_i1] = matmul_1[v_i0, v_i1] + rxplaceholder[v_i0, v_k] * rxplaceholder_1[v_k, v_i1] @Ts.prim_func(private=True) - def matmul1(rxplaceholder: T.Buffer((T.int64(1), T.int64(128)), "float32"), rxplaceholder_1: T.Buffer((T.int64(128), T.int64(10)), "float32"), matmul: T.Buffer((T.int64(1), T.int64(10)), "float32")): + def matmul1(rxplaceholder: T.Tensor((T.int64(1), T.int64(128)), "float32"), rxplaceholder_1: T.Tensor((T.int64(128), T.int64(10)), "float32"), matmul: T.Tensor((T.int64(1), T.int64(10)), "float32")): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1, k in T.grid(T.int64(1), T.int64(10), T.int64(128)): @@ -1215,7 +1215,7 @@ def matmul1(rxplaceholder: T.Buffer((T.int64(1), T.int64(128)), "float32"), rxpl matmul[v_i0, v_i1] = matmul[v_i0, v_i1] + rxplaceholder[v_i0, v_k] * rxplaceholder_1[v_k, v_i1] @Ts.prim_func(private=True) - def relu(rxplaceholder: T.Buffer((T.int64(1), T.int64(128)), "float32"), compute: T.Buffer((T.int64(1), T.int64(128)), "float32")): + def relu(rxplaceholder: T.Tensor((T.int64(1), T.int64(128)), "float32"), compute: T.Tensor((T.int64(1), T.int64(128)), "float32")): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1 in T.grid(T.int64(1), T.int64(128)): @@ -1226,7 +1226,7 @@ def relu(rxplaceholder: T.Buffer((T.int64(1), T.int64(128)), "float32"), compute compute[v_i0, v_i1] = T.max(rxplaceholder[v_i0, v_i1], T.float32(0)) @Ts.prim_func(private=True) - def transpose(rxplaceholder: T.Buffer((T.int64(128), T.int64(784)), "float32"), T_transpose: T.Buffer((T.int64(784), T.int64(128)), "float32")): + def transpose(rxplaceholder: T.Tensor((T.int64(128), T.int64(784)), "float32"), T_transpose: T.Tensor((T.int64(784), T.int64(128)), "float32")): T.func_attr({"op_pattern": 2, "tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1 in T.grid(T.int64(784), T.int64(128)): @@ -1237,7 +1237,7 @@ def transpose(rxplaceholder: T.Buffer((T.int64(128), T.int64(784)), "float32"), T_transpose[v_ax0, v_ax1] = rxplaceholder[v_ax1, v_ax0] @Ts.prim_func(private=True) - def transpose1(rxplaceholder: T.Buffer((T.int64(10), T.int64(128)), "float32"), T_transpose: T.Buffer((T.int64(128), T.int64(10)), "float32")): + def transpose1(rxplaceholder: T.Tensor((T.int64(10), T.int64(128)), "float32"), T_transpose: T.Tensor((T.int64(128), T.int64(10)), "float32")): T.func_attr({"op_pattern": 2, "tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1 in T.grid(T.int64(128), T.int64(10)): @@ -1387,9 +1387,9 @@ class Before: @Ts.prim_func(private=True) def add_one( - x: T.Buffer((T.int64(1), add_one_n), "float32"), + x: T.Tensor((T.int64(1), add_one_n), "float32"), n: add_one_n, - out: T.Buffer((T.int64(1), add_one_n), "float32"), + out: T.Tensor((T.int64(1), add_one_n), "float32"), ): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) @@ -1402,9 +1402,9 @@ def add_one( @Ts.prim_func(private=True) def exp( - x: T.Buffer((T.int64(1), exp_n), "float32"), + x: T.Tensor((T.int64(1), exp_n), "float32"), n: exp_n, - out: T.Buffer((T.int64(1), exp_n), "float32"), + out: T.Tensor((T.int64(1), exp_n), "float32"), ): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) @@ -1465,8 +1465,8 @@ class Before: @Ts.prim_func(private=True) def add_one( n: add_one_n, - x: T.Buffer((T.int64(1), add_one_n), "float32"), - out: T.Buffer((T.int64(1), add_one_n), "float32"), + x: T.Tensor((T.int64(1), add_one_n), "float32"), + out: T.Tensor((T.int64(1), add_one_n), "float32"), ): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) @@ -1480,8 +1480,8 @@ def add_one( @Ts.prim_func(private=True) def exp( n: exp_n, - x: T.Buffer((T.int64(1), exp_n), "float32"), - out: T.Buffer((T.int64(1), exp_n), "float32"), + x: T.Tensor((T.int64(1), exp_n), "float32"), + out: T.Tensor((T.int64(1), exp_n), "float32"), ): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) @@ -1541,11 +1541,11 @@ class Before: @Ts.prim_func(private=True) def add_one( - x: T.Buffer( + x: T.Tensor( (T.int64(1), (add_one_n - T.int64(1)) // T.int64(4) + T.int64(1)), "float32" ), n: add_one_n, - out: T.Buffer( + out: T.Tensor( (T.int64(1), (add_one_n - T.int64(1)) // T.int64(4) + T.int64(1)), "float32" ), ): @@ -1560,9 +1560,9 @@ def add_one( @Ts.prim_func(private=True) def exp( - x: T.Buffer((T.int64(1), (exp_n - T.int64(1)) // T.int64(4) + T.int64(1)), "float32"), + x: T.Tensor((T.int64(1), (exp_n - T.int64(1)) // T.int64(4) + T.int64(1)), "float32"), n: exp_n, - out: T.Buffer((T.int64(1), (exp_n - T.int64(1)) // T.int64(4) + T.int64(1)), "float32"), + out: T.Tensor((T.int64(1), (exp_n - T.int64(1)) // T.int64(4) + T.int64(1)), "float32"), ): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) @@ -1622,10 +1622,10 @@ class Before: @Ts.prim_func(private=True) def add_one( - x: T.Buffer((T.int64(1), add_one_n + T.int64(1)), "float32"), + x: T.Tensor((T.int64(1), add_one_n + T.int64(1)), "float32"), n: add_one_n, m: T.int64, - out: T.Buffer((T.int64(1), add_one_n + T.int64(1)), "float32"), + out: T.Tensor((T.int64(1), add_one_n + T.int64(1)), "float32"), ): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) @@ -1638,10 +1638,10 @@ def add_one( @Ts.prim_func(private=True) def exp( - x: T.Buffer((T.int64(1), exp_n + T.int64(1)), "float32"), + x: T.Tensor((T.int64(1), exp_n + T.int64(1)), "float32"), n: exp_n, m: T.int64, - out: T.Buffer((T.int64(1), exp_n + T.int64(1)), "float32"), + out: T.Tensor((T.int64(1), exp_n + T.int64(1)), "float32"), ): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) @@ -1688,9 +1688,9 @@ def test_primitive_call_arg_not_inlined(): class Before: @Ts.prim_func(private=True) def add_scalar( - x: T.Buffer((T.int64(4),), "int64"), + x: T.Tensor((T.int64(4),), "int64"), value: T.int64, - out: T.Buffer((T.int64(4),), "int64"), + out: T.Tensor((T.int64(4),), "int64"), ): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) @@ -1700,7 +1700,7 @@ def add_scalar( out[vi] = x[vi] + value @Ts.prim_func(private=True) - def double(x: T.Buffer((T.int64(4),), "int64"), out: T.Buffer((T.int64(4),), "int64")): + def double(x: T.Tensor((T.int64(4),), "int64"), out: T.Tensor((T.int64(4),), "int64")): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) for i in range(4): @@ -1752,7 +1752,7 @@ def test_primitive_call_arg_used_by_output_shape_not_inlined(): @I.ir_module class Before: @Ts.prim_func(private=True) - def make(n: T.int64, out: T.Buffer((n,), "float32")): # noqa: F821 + def make(n: T.int64, out: T.Tensor((n,), "float32")): # noqa: F821 T.func_attr({"op_pattern": 0, "tirx.noalias": True}) for i in range(n): @@ -1761,7 +1761,7 @@ def make(n: T.int64, out: T.Buffer((n,), "float32")): # noqa: F821 out[vi] = T.float32(1) @Ts.prim_func(private=True) - def double(x: T.Buffer((n,), "float32"), n: T.int64, out: T.Buffer((n,), "float32")): # noqa: F821 + def double(x: T.Tensor((n,), "float32"), n: T.int64, out: T.Tensor((n,), "float32")): # noqa: F821 T.func_attr({"op_pattern": 0, "tirx.noalias": True}) for i in range(n): @@ -1828,7 +1828,7 @@ class Before: make_n = T.int64() @Ts.prim_func(private=True) - def make(n: make_n, out: T.Buffer((make_n,), "float32")): + def make(n: make_n, out: T.Tensor((make_n,), "float32")): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) for i in range(n): @@ -1840,7 +1840,7 @@ def make(n: make_n, out: T.Buffer((make_n,), "float32")): @Ts.prim_func(private=True) def double( - x: T.Buffer((double_n,), "float32"), n: double_n, out: T.Buffer((double_n,), "float32") + x: T.Tensor((double_n,), "float32"), n: double_n, out: T.Tensor((double_n,), "float32") ): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) @@ -2067,9 +2067,9 @@ def test_call_tir_inplace(): class Module: @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(10), T.int64(20)), "float32"), - B: T.Buffer((), "float32"), - Out: T.Buffer((T.int64(10), T.int64(20)), "float32"), + A: T.Tensor((T.int64(10), T.int64(20)), "float32"), + B: T.Tensor((), "float32"), + Out: T.Tensor((T.int64(10), T.int64(20)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): @@ -2080,7 +2080,7 @@ def add( Out[v_ax0, v_ax1] = A[v_ax0, v_ax1] + B[()] @Ts.prim_func(private=True) - def exp_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): + def exp_inplace(A: T.Tensor((T.int64(10), T.int64(20)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(10), T.int64(20)): with Ts.sblock("compute"): @@ -2090,7 +2090,7 @@ def exp_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): A[v_i0, v_i1] = T.exp(A[v_i0, v_i1]) @Ts.prim_func(private=True) - def squeeze_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): + def squeeze_inplace(A: T.Tensor((T.int64(10), T.int64(20)), "float32")): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): with Ts.sblock("T_squeeze"): @@ -2129,9 +2129,9 @@ def main( class Expected: @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(10), T.int64(20)), "float32"), - B: T.Buffer((), "float32"), - Out: T.Buffer((T.int64(10), T.int64(20)), "float32"), + A: T.Tensor((T.int64(10), T.int64(20)), "float32"), + B: T.Tensor((), "float32"), + Out: T.Tensor((T.int64(10), T.int64(20)), "float32"), ): T.func_attr({"tirx.noalias": True, "op_pattern": 0}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): @@ -2142,7 +2142,7 @@ def add( Out[v_ax0, v_ax1] = A[v_ax0, v_ax1] + B[()] @Ts.prim_func(private=True) - def exp_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): + def exp_inplace(A: T.Tensor((T.int64(10), T.int64(20)), "float32")): T.func_attr({"tirx.noalias": True, "op_pattern": 0}) for i0, i1 in T.grid(T.int64(10), T.int64(20)): with Ts.sblock("compute"): @@ -2152,7 +2152,7 @@ def exp_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): A[v_i0, v_i1] = T.exp(A[v_i0, v_i1]) @Ts.prim_func(private=True) - def squeeze_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): + def squeeze_inplace(A: T.Tensor((T.int64(10), T.int64(20)), "float32")): T.func_attr({"tirx.noalias": True, "op_pattern": 0}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): with Ts.sblock("T_squeeze"): @@ -2208,7 +2208,7 @@ def test_packed_params(): @I.ir_module class Before: @Ts.prim_func(private=True) - def cast(lv: T.Buffer((T.int64(16), T.int64(16)), "float16"), compute: T.Buffer((T.int64(16), T.int64(16)), "float32")): + def cast(lv: T.Tensor((T.int64(16), T.int64(16)), "float16"), compute: T.Tensor((T.int64(16), T.int64(16)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1 in T.grid(T.int64(16), T.int64(16)): @@ -2219,7 +2219,7 @@ def cast(lv: T.Buffer((T.int64(16), T.int64(16)), "float16"), compute: T.Buffer( compute[v_i0, v_i1] = T.Cast("float32", lv[v_i0, v_i1]) @Ts.prim_func(private=True) - def matmul(x: T.Buffer((T.int64(16), T.int64(16)), "float32"), lv2: T.Buffer((T.int64(16), T.int64(16)), "float32"), T_matmul: T.Buffer((T.int64(16), T.int64(16)), "float32")): + def matmul(x: T.Tensor((T.int64(16), T.int64(16)), "float32"), lv2: T.Tensor((T.int64(16), T.int64(16)), "float32"), T_matmul: T.Tensor((T.int64(16), T.int64(16)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1, k in T.grid(T.int64(16), T.int64(16), T.int64(16)): 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 a1ded6cfbc9e..8abe9fb7b5fc 100644 --- a/tests/python/relax/test_transform_fuse_ops_by_pattern.py +++ b/tests/python/relax/test_transform_fuse_ops_by_pattern.py @@ -703,8 +703,8 @@ def test_ignore_call_tir(): class Conv2dReLUCallTIR: @Ts.prim_func def relu( - data: T.Buffer((1, 64, 56, 56), "float32"), - out: T.Buffer((1, 64, 56, 56), "float32"), + data: T.Tensor((1, 64, 56, 56), "float32"), + out: T.Tensor((1, 64, 56, 56), "float32"), ): for ax0, ax1, ax2, ax3 in T.grid(1, 64, 56, 56): with Ts.sblock("root"): @@ -731,8 +731,8 @@ def main( class Conv2dReLUCallTIR_partitioned: @Ts.prim_func def relu( - data: T.Buffer((1, 64, 56, 56), "float32"), - out: T.Buffer((1, 64, 56, 56), "float32"), + data: T.Tensor((1, 64, 56, 56), "float32"), + out: T.Tensor((1, 64, 56, 56), "float32"), ): # with Ts.sblock("root"): for ax0, ax1, ax2, ax3 in T.grid(1, 64, 56, 56): diff --git a/tests/python/relax/test_transform_fuse_tir.py b/tests/python/relax/test_transform_fuse_tir.py index 6da47c544da2..85c51f886980 100644 --- a/tests/python/relax/test_transform_fuse_tir.py +++ b/tests/python/relax/test_transform_fuse_tir.py @@ -634,9 +634,9 @@ def func2(x: R.Tensor((20, 10), dtype="float32")) -> R.Tensor((20, 10), dtype="f @Ts.prim_func(private=True) def fused_add1_exp1_squeeze1( - x: T.Buffer((T.int64(20), T.int64(10)), "float32"), - p0: T.Buffer((), "float32"), - T_squeeze: T.Buffer((T.int64(20), T.int64(10)), "float32"), + x: T.Tensor((T.int64(20), T.int64(10)), "float32"), + p0: T.Tensor((), "float32"), + T_squeeze: T.Tensor((T.int64(20), T.int64(10)), "float32"), ): T.func_attr({"tirx.noalias": True}) T_add = Ts.sblock_alloc_buffer((T.int64(20), T.int64(10))) @@ -662,9 +662,9 @@ def fused_add1_exp1_squeeze1( @Ts.prim_func(private=True) def fused_add_exp_squeeze( - x: T.Buffer((T.int64(10), T.int64(20)), "float32"), - p0: T.Buffer((), "float32"), - T_squeeze: T.Buffer((T.int64(10), T.int64(20)), "float32"), + x: T.Tensor((T.int64(10), T.int64(20)), "float32"), + p0: T.Tensor((), "float32"), + T_squeeze: T.Tensor((T.int64(10), T.int64(20)), "float32"), ): T.func_attr({"tirx.noalias": True}) T_add = Ts.sblock_alloc_buffer((T.int64(10), T.int64(20))) @@ -761,7 +761,7 @@ def test_fuse_of_dynamic_kernel_with_var_params_and_static_args(): @I.ir_module class Before: @Ts.prim_func(private=True) - def dynamic_tir_kernel(A: T.Buffer([m, n], "float32"), B: T.Buffer([m, n], "float32")): + def dynamic_tir_kernel(A: T.Tensor([m, n], "float32"), B: T.Tensor([m, n], "float32")): for (*iters,) in T.grid(m, n): with Ts.sblock("compute"): i, j = Ts.axis.remap("SS", iters) @@ -789,8 +789,8 @@ def main(x: R.Tensor([16, 32], "float32")) -> R.Tensor([16, 32], dtype="float32" class Expected: @Ts.prim_func(private=True) def fused_function( - X: T.Buffer([T.int64(16), T.int64(32)], "float32"), - Z: T.Buffer([T.int64(16), T.int64(32)], "float32"), + X: T.Tensor([T.int64(16), T.int64(32)], "float32"), + Z: T.Tensor([T.int64(16), T.int64(32)], "float32"), ): T.func_attr({"tirx.noalias": True}) Y = Ts.sblock_alloc_buffer(X.shape, "float32") @@ -828,10 +828,10 @@ def test_fuse_of_dynamic_kernel_with_expression_params_and_static_args(): class Before: @Ts.prim_func(private=True) def dynamic_tir_kernel( - A: T.Buffer([m * n], "float32"), - B: T.Buffer([m], "float32"), - C: T.Buffer([n], "float32"), - D: T.Buffer([m * n], "float32"), + A: T.Tensor([m * n], "float32"), + B: T.Tensor([m], "float32"), + C: T.Tensor([n], "float32"), + D: T.Tensor([m * n], "float32"), ): for i, j in T.grid(m, n): with Ts.sblock("compute"): @@ -872,10 +872,10 @@ def main( class Expected: @Ts.prim_func(private=True) def fused_function( - X: T.Buffer(T.int64(512), "float32"), - B: T.Buffer(T.int64(16), "float32"), - C: T.Buffer(T.int64(32), "float32"), - Z: T.Buffer(T.int64(512), "float32"), + X: T.Tensor(T.int64(512), "float32"), + B: T.Tensor(T.int64(16), "float32"), + C: T.Tensor(T.int64(32), "float32"), + Z: T.Tensor(T.int64(512), "float32"), ): T.func_attr({"tirx.noalias": True}) Y = Ts.sblock_alloc_buffer((T.int64(512),)) @@ -976,10 +976,10 @@ def test_symbolic_var_in_call_tir_args(): class Before: @Ts.prim_func(private=True) def foo( - X: T.Buffer((T.int64(1), T.int64(1), T.int64(32), T.int64(128)), "float32"), - Y: T.Buffer((T.int64(2048), T.int64(128)), "float32"), + X: T.Tensor((T.int64(1), T.int64(1), T.int64(32), T.int64(128)), "float32"), + Y: T.Tensor((T.int64(2048), T.int64(128)), "float32"), m: T.int64, - rotary: T.Buffer((T.int64(1), T.int64(1), T.int64(32), T.int64(128)), "float32"), + rotary: T.Tensor((T.int64(1), T.int64(1), T.int64(32), T.int64(128)), "float32"), ): for i0, i1, i2, i3 in T.grid(T.int64(1), T.int64(1), T.int64(32), T.int64(128)): with Ts.sblock("rotary"): @@ -1022,10 +1022,10 @@ def main( class Expected: @Ts.prim_func(private=True) def fused( - X: T.Buffer((T.int64(1), T.int64(1), T.int64(32), T.int64(128)), "float32"), - Y: T.Buffer((T.int64(2048), T.int64(128)), "float32"), + X: T.Tensor((T.int64(1), T.int64(1), T.int64(32), T.int64(128)), "float32"), + Y: T.Tensor((T.int64(2048), T.int64(128)), "float32"), m: T.int64, - rotary: T.Buffer((T.int64(1), T.int64(1), T.int64(32), T.int64(128)), "float32"), + rotary: T.Tensor((T.int64(1), T.int64(1), T.int64(32), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) T_add = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(32), T.int64(128))) @@ -1064,11 +1064,11 @@ def test_same_buffer_multiple_read(): class Module: @Ts.prim_func(private=True) def concatenate( - rxplaceholder: T.Buffer((T.int64(1), T.int64(4), T.int64(64), T.int64(64)), "float32"), - rxplaceholder_1: T.Buffer( + rxplaceholder: T.Tensor((T.int64(1), T.int64(4), T.int64(64), T.int64(64)), "float32"), + rxplaceholder_1: T.Tensor( (T.int64(1), T.int64(4), T.int64(64), T.int64(64)), "float32" ), - T_concat: T.Buffer((T.int64(2), T.int64(4), T.int64(64), T.int64(64)), "float32"), + T_concat: T.Tensor((T.int64(2), T.int64(4), T.int64(64), T.int64(64)), "float32"), ): T.func_attr({"op_pattern": 2, "tirx.noalias": True}) for ax0, ax1, ax2, ax3 in T.grid(T.int64(2), T.int64(4), T.int64(64), T.int64(64)): @@ -1087,8 +1087,8 @@ def concatenate( @Ts.prim_func(private=True) def transpose2( - rxplaceholder: T.Buffer((T.int64(2), T.int64(4), T.int64(64), T.int64(64)), "float32"), - T_transpose: T.Buffer((T.int64(2), T.int64(64), T.int64(64), T.int64(4)), "float32"), + rxplaceholder: T.Tensor((T.int64(2), T.int64(4), T.int64(64), T.int64(64)), "float32"), + T_transpose: T.Tensor((T.int64(2), T.int64(64), T.int64(64), T.int64(4)), "float32"), ): T.func_attr({"op_pattern": 2, "tirx.noalias": True}) for ax0, ax1, ax2, ax3 in T.grid(T.int64(2), T.int64(64), T.int64(64), T.int64(4)): @@ -1133,8 +1133,8 @@ def main(inp_0: R.Tensor((1, 4, 64, 64), dtype="float32")) -> R.Tensor( class Expected: @Ts.prim_func(private=True) def fused_concatenate_transpose2( - inp_0: T.Buffer((T.int64(1), T.int64(4), T.int64(64), T.int64(64)), "float32"), - T_transpose_handle_intermediate: T.Buffer( + inp_0: T.Tensor((T.int64(1), T.int64(4), T.int64(64), T.int64(64)), "float32"), + T_transpose_handle_intermediate: T.Tensor( (T.int64(2), T.int64(64), T.int64(64), T.int64(4)), "float32" ), ): @@ -1218,10 +1218,10 @@ class Expected: @Ts.prim_func(private=True) def fused_transpose_matmul( - x: T.Buffer((T.int64(3), T.int64(4)), "float32"), - y: T.Buffer((fused_transpose_matmul_n - T.int64(1), T.int64(4))), + x: T.Tensor((T.int64(3), T.int64(4)), "float32"), + y: T.Tensor((fused_transpose_matmul_n - T.int64(1), T.int64(4))), n: fused_transpose_matmul_n, - var_T_matmul_intermediate: T.Buffer( + var_T_matmul_intermediate: T.Tensor( (fused_transpose_matmul_n - T.int64(1), T.int64(3)) ), ): @@ -1266,8 +1266,8 @@ def test_tuple_input_unused_field(): class Module: @Ts.prim_func(private=True) def reshape( - A: T.Buffer((T.int64(4), T.int64(8), T.int64(2048)), "float32"), - T_reshape: T.Buffer((T.int64(4), T.int64(8), T.int64(32), T.int64(64)), "float32"), + A: T.Tensor((T.int64(4), T.int64(8), T.int64(2048)), "float32"), + T_reshape: T.Tensor((T.int64(4), T.int64(8), T.int64(32), T.int64(64)), "float32"), ): T.func_attr({"op_pattern": 2, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -1329,8 +1329,8 @@ def main( class Expected: @Ts.prim_func(private=True) def fused_reshape( - lv_0: T.Buffer((T.int64(4), T.int64(8), T.int64(2048)), "float32"), - T_reshape_handle_intermediate: T.Buffer( + lv_0: T.Tensor((T.int64(4), T.int64(8), T.int64(2048)), "float32"), + T_reshape_handle_intermediate: T.Tensor( (T.int64(4), T.int64(8), T.int64(32), T.int64(64)), "float32" ), ): @@ -1385,8 +1385,8 @@ def test_unique_duplicated_buffer_allocation(): class Module: @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - Out: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + Out: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), ): for i, j in T.grid(T.int64(4096), T.int64(4096)): with Ts.sblock("add"): @@ -1395,8 +1395,8 @@ def add( @Ts.prim_func(private=True) def add1( - A: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - Out: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + Out: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), ): for i, j in T.grid(T.int64(4096), T.int64(4096)): with Ts.sblock("add"): @@ -1431,8 +1431,8 @@ def fused_func( class Expected: @Ts.prim_func(private=True) def fused_func( - input_embeds: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - Out_intermediate_1: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), + input_embeds: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + Out_intermediate_1: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) Out_intermediate = Ts.sblock_alloc_buffer((T.int64(4096), T.int64(4096)), "float16") @@ -1493,10 +1493,10 @@ def test_symbolic_var_in_buffer_shape(): class Before: @Ts.prim_func(private=True) def foo( - X: T.Buffer([T.int64(1), sequence_length_foo, T.int64(32), T.int64(128)], "float32"), - Y: T.Buffer((T.int64(2048), T.int64(128)), "float32"), + X: T.Tensor([T.int64(1), sequence_length_foo, T.int64(32), T.int64(128)], "float32"), + Y: T.Tensor((T.int64(2048), T.int64(128)), "float32"), m: T.int64, - rotary: T.Buffer( + rotary: T.Tensor( [T.int64(1), sequence_length_foo, T.int64(32), T.int64(128)], "float32" ), ): @@ -1545,10 +1545,10 @@ def main( class Expected: @Ts.prim_func(private=True) def fused( - X: T.Buffer([T.int64(1), sequence_length_fused, T.int64(32), T.int64(128)], "float32"), - Y: T.Buffer((T.int64(2048), T.int64(128)), "float32"), + X: T.Tensor([T.int64(1), sequence_length_fused, T.int64(32), T.int64(128)], "float32"), + Y: T.Tensor((T.int64(2048), T.int64(128)), "float32"), m: T.int64, - rotary: T.Buffer( + rotary: T.Tensor( [T.int64(1), sequence_length_fused, T.int64(32), T.int64(128)], "float32" ), ): @@ -1600,8 +1600,8 @@ def test_symbolic_var_called_with_static_shape(): class Before: @Ts.prim_func(private=True) def sum_1d( - X: T.Buffer([num_elements], "float32"), - Y: T.Buffer([T.int64(1)], "float32"), + X: T.Tensor([num_elements], "float32"), + Y: T.Tensor([T.int64(1)], "float32"), ): for i in range(num_elements): with Ts.sblock("sum"): @@ -1639,8 +1639,8 @@ def main( class Expected: @Ts.prim_func(private=True) def fused( - X: T.Buffer([T.int64(64)], "float32"), - Y: T.Buffer([T.int64(1)], "float32"), + X: T.Tensor([T.int64(64)], "float32"), + Y: T.Tensor([T.int64(1)], "float32"), ): T.func_attr({"tirx.noalias": True}) @@ -1673,8 +1673,8 @@ def test_symbolic_var_called_with_multiple_static_shapes(): class Before: @Ts.prim_func(private=True) def sum_1d( - X: T.Buffer([num_elements], "float32"), - Sum: T.Buffer([T.int64(1)], "float32"), + X: T.Tensor([num_elements], "float32"), + Sum: T.Tensor([T.int64(1)], "float32"), ): for i in range(num_elements): with Ts.sblock("sum"): @@ -1685,9 +1685,9 @@ def sum_1d( @Ts.prim_func(private=True) def sum_scalar( - X: T.Buffer([T.int64(1)], "float32"), - Y: T.Buffer([T.int64(1)], "float32"), - Sum: T.Buffer([T.int64(1)], "float32"), + X: T.Tensor([T.int64(1)], "float32"), + Y: T.Tensor([T.int64(1)], "float32"), + Sum: T.Tensor([T.int64(1)], "float32"), ): for i in range(T.int64(1)): with Ts.sblock("Out"): @@ -1735,9 +1735,9 @@ def main( class Expected: @Ts.prim_func(private=True) def fused( - X: T.Buffer([T.int64(64)], "float32"), - Y: T.Buffer([T.int64(16)], "float32"), - Out: T.Buffer([T.int64(1)], "float32"), + X: T.Tensor([T.int64(64)], "float32"), + Y: T.Tensor([T.int64(16)], "float32"), + Out: T.Tensor([T.int64(1)], "float32"), ): T.func_attr({"tirx.noalias": True}) @@ -1793,9 +1793,9 @@ def test_symbolic_var_called_with_static_argument(): class Before: @Ts.prim_func(private=True) def sum_1d( - X: T.Buffer([num_elements], "float32"), # noqa: F821 + X: T.Tensor([num_elements], "float32"), # noqa: F821 num_elements: T.int64, - Y: T.Buffer([T.int64(1)], "float32"), + Y: T.Tensor([T.int64(1)], "float32"), ): for i in range(num_elements): with Ts.sblock("sum"): @@ -1833,8 +1833,8 @@ def main( class Expected: @Ts.prim_func(private=True) def fused( - X: T.Buffer([T.int64(64)], "float32"), - Y: T.Buffer([T.int64(1)], "float32"), + X: T.Tensor([T.int64(64)], "float32"), + Y: T.Tensor([T.int64(1)], "float32"), ): T.func_attr({"tirx.noalias": True}) @@ -1863,8 +1863,8 @@ def test_gather(): class Before: @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - Out: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + Out: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), ): for i, j in T.grid(T.int64(4096), T.int64(4096)): with Ts.sblock("add"): @@ -1873,9 +1873,9 @@ def add( @Ts.prim_func(private=True) def take( - A: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - B: T.Buffer((T.int64(1),), "int32"), - T_take: T.Buffer((T.int64(1), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + B: T.Tensor((T.int64(1),), "int32"), + T_take: T.Tensor((T.int64(1), T.int64(4096)), "float16"), ): for ax0, ax1 in T.grid(T.int64(1), T.int64(4096)): with Ts.sblock("T_take"): @@ -1914,9 +1914,9 @@ def fused_func( class After: @Ts.prim_func(private=True) def fused_func( - input_ids: T.Buffer((T.int64(1),), "int32"), - input_embeds: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - T_take: T.Buffer((T.int64(1), T.int64(4096)), "float16"), + input_ids: T.Tensor((T.int64(1),), "int32"), + input_embeds: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + T_take: T.Tensor((T.int64(1), T.int64(4096)), "float16"), ): T.func_attr({"tirx.noalias": True}) Out_handle_intermediate = Ts.sblock_alloc_buffer( @@ -1956,7 +1956,7 @@ class Module: @Ts.prim_func(private=True) def add_inplace( - A: T.Buffer((T.int64(10), T.int64(20)), "float32"), B: T.Buffer((), "float32") + A: T.Tensor((T.int64(10), T.int64(20)), "float32"), B: T.Tensor((), "float32") ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): @@ -1967,7 +1967,7 @@ def add_inplace( A[v_ax0, v_ax1] = A[v_ax0, v_ax1] + B[()] @Ts.prim_func(private=True) - def exp_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): + def exp_inplace(A: T.Tensor((T.int64(10), T.int64(20)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(10), T.int64(20)): with Ts.sblock("compute"): @@ -1977,7 +1977,7 @@ def exp_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): A[v_i0, v_i1] = T.exp(A[v_i0, v_i1]) @Ts.prim_func(private=True) - def squeeze_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): + def squeeze_inplace(A: T.Tensor((T.int64(10), T.int64(20)), "float32")): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): with Ts.sblock("T_squeeze"): @@ -2035,7 +2035,7 @@ class Expected: @Ts.prim_func(private=True) def fused_add_exp_squeeze( - x: T.Buffer((T.int64(10), T.int64(20)), "float32"), p0: T.Buffer((), "float32") + x: T.Tensor((T.int64(10), T.int64(20)), "float32"), p0: T.Tensor((), "float32") ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): @@ -2077,9 +2077,9 @@ class Module: @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(10), T.int64(20)), "float32"), - B: T.Buffer((), "float32"), - Out: T.Buffer((T.int64(10), T.int64(20)), "float32"), + A: T.Tensor((T.int64(10), T.int64(20)), "float32"), + B: T.Tensor((), "float32"), + Out: T.Tensor((T.int64(10), T.int64(20)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): @@ -2088,7 +2088,7 @@ def add( Out[v_ax0, v_ax1] = A[v_ax0, v_ax1] + B[()] @Ts.prim_func(private=True) - def exp_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): + def exp_inplace(A: T.Tensor((T.int64(10), T.int64(20)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(10), T.int64(20)): with Ts.sblock("compute"): @@ -2096,7 +2096,7 @@ def exp_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): A[v_i0, v_i1] = T.exp(A[v_i0, v_i1]) @Ts.prim_func(private=True) - def squeeze_inplace(A: T.Buffer((T.int64(10), T.int64(20)), "float32")): + def squeeze_inplace(A: T.Tensor((T.int64(10), T.int64(20)), "float32")): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): with Ts.sblock("T_squeeze"): @@ -2146,9 +2146,9 @@ class Expected: @Ts.prim_func(private=True) def fused_add_exp_squeeze( - x: T.Buffer((T.int64(10), T.int64(20)), "float32"), - p0: T.Buffer((), "float32"), - p_output0: T.Buffer((T.int64(10), T.int64(20)), "float32"), + x: T.Tensor((T.int64(10), T.int64(20)), "float32"), + p0: T.Tensor((), "float32"), + p_output0: T.Tensor((T.int64(10), T.int64(20)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): @@ -2187,9 +2187,9 @@ class Module: # we will use it both in-place and normally (DPS) @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(10), T.int64(20)), "float32"), - B: T.Buffer((), "float32"), - Out: T.Buffer((T.int64(10), T.int64(20)), "float32"), + A: T.Tensor((T.int64(10), T.int64(20)), "float32"), + B: T.Tensor((), "float32"), + Out: T.Tensor((T.int64(10), T.int64(20)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): @@ -2238,9 +2238,9 @@ def main( class Expected: @Ts.prim_func(private=True) def fused_sums( - x: T.Buffer((T.int64(10), T.int64(20)), "float32"), - p0: T.Buffer((), "float32"), - p_output0: T.Buffer((T.int64(10), T.int64(20)), "float32"), + x: T.Tensor((T.int64(10), T.int64(20)), "float32"), + p0: T.Tensor((), "float32"), + p_output0: T.Tensor((T.int64(10), T.int64(20)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(10), T.int64(20)): @@ -2311,8 +2311,8 @@ def fused_func( @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - Out: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + Out: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), ): for i, j in T.grid(T.int64(4096), T.int64(4096)): with Ts.sblock("add"): @@ -2321,9 +2321,9 @@ def add( @Ts.prim_func(private=True) def take( - A: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), - B: T.Buffer((T.int64(1),), "int32"), - T_take: T.Buffer((T.int64(1), T.int64(4096)), "float16"), + A: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), + B: T.Tensor((T.int64(1),), "int32"), + T_take: T.Tensor((T.int64(1), T.int64(4096)), "float16"), ): for ax0, ax1 in T.grid(T.int64(1), T.int64(4096)): with Ts.sblock("T_take"): @@ -2337,7 +2337,7 @@ def test_block_name_numeric_suffix_deduplication(): @I.ir_module class Before: @Ts.prim_func(private=True) - def add1(x: T.Buffer((10,), "float32"), y: T.Buffer((10,), "float32")): + def add1(x: T.Tensor((10,), "float32"), y: T.Tensor((10,), "float32")): T.func_attr({"tirx.noalias": True}) for i in range(10): with Ts.sblock("compute1"): @@ -2345,7 +2345,7 @@ def add1(x: T.Buffer((10,), "float32"), y: T.Buffer((10,), "float32")): y[vi] = x[vi] + T.float32(1.0) @Ts.prim_func(private=True) - def mul1(x: T.Buffer((10,), "float32"), y: T.Buffer((10,), "float32")): + def mul1(x: T.Tensor((10,), "float32"), y: T.Tensor((10,), "float32")): T.func_attr({"tirx.noalias": True}) for i in range(10): with Ts.sblock("compute1"): @@ -2374,8 +2374,8 @@ def main(x: R.Tensor((10,), dtype="float32")) -> R.Tensor((10,), dtype="float32" class Expected: @Ts.prim_func(private=True) def fused_add_mul( - x: T.Buffer((T.int64(10),)), - y_intermediate_1: T.Buffer((T.int64(10),), elem_offset=T.int32(0)), + x: T.Tensor((T.int64(10),)), + y_intermediate_1: T.Tensor((T.int64(10),), elem_offset=T.int32(0)), ): T.func_attr({"tirx.noalias": True}) @@ -2411,7 +2411,7 @@ def test_primitive_scalar_parameter_preserves_identity(): @I.ir_module class Before: @Ts.prim_func(private=True) - def add_scalar(x: T.Buffer((4,), "int64"), p: T.int64, y: T.Buffer((1,), "int64")): + def add_scalar(x: T.Tensor((4,), "int64"), p: T.int64, y: T.Tensor((1,), "int64")): for i in range(1): with Ts.sblock("add"): vi = Ts.axis.spatial(1, i) @@ -2443,7 +2443,7 @@ def test_inplace_argument_after_primitive_scalar(): @I.ir_module class Before: @Ts.prim_func(private=True) - def add_scalar_inplace(p: T.int64, x: T.Buffer((4,), "int64")): + def add_scalar_inplace(p: T.int64, x: T.Tensor((4,), "int64")): for i in range(4): with Ts.sblock("add"): vi = Ts.axis.spatial(4, i) diff --git a/tests/python/relax/test_transform_fuse_transpose_matmul.py b/tests/python/relax/test_transform_fuse_transpose_matmul.py index 311431d0c410..f4a4a8f311e5 100644 --- a/tests/python/relax/test_transform_fuse_transpose_matmul.py +++ b/tests/python/relax/test_transform_fuse_transpose_matmul.py @@ -45,9 +45,9 @@ def main( class Expected: @Ts.prim_func(private=True) def NT_matmul( - x: T.Buffer((T.int64(128), T.int64(256)), "float32"), - w: T.Buffer((T.int64(128), T.int64(256)), "float32"), - NT_matmul: T.Buffer((T.int64(128), T.int64(128)), "float32"), + x: T.Tensor((T.int64(128), T.int64(256)), "float32"), + w: T.Tensor((T.int64(128), T.int64(256)), "float32"), + NT_matmul: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -98,9 +98,9 @@ def main( class Expected: @Ts.prim_func(private=True) def NT_matmul( - x: T.Buffer((T.int64(128), T.int64(256)), "float32"), - w: T.Buffer((T.int64(128), T.int64(256)), "float32"), - NT_matmul: T.Buffer((T.int64(128), T.int64(128)), "float32"), + x: T.Tensor((T.int64(128), T.int64(256)), "float32"), + w: T.Tensor((T.int64(128), T.int64(256)), "float32"), + NT_matmul: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): diff --git a/tests/python/relax/test_transform_gradient.py b/tests/python/relax/test_transform_gradient.py index c5f6f4a7def9..500a1f75ad28 100644 --- a/tests/python/relax/test_transform_gradient.py +++ b/tests/python/relax/test_transform_gradient.py @@ -1211,8 +1211,8 @@ def main(x0: R.Tensor((3, 3), "float32"), x1: R.Tensor((3, 3), "float32")): @Ts.prim_func def sum( - rxplaceholder: T.Buffer((T.int64(3), T.int64(3)), "float32"), - rxplaceholder_red: T.Buffer((), "float32"), + rxplaceholder: T.Tensor((T.int64(3), T.int64(3)), "float32"), + rxplaceholder_red: T.Tensor((), "float32"), ): T.func_attr({"tirx.noalias": True}) for k0, k1 in T.grid(T.int64(3), T.int64(3)): diff --git a/tests/python/relax/test_transform_gradient_te_register.py b/tests/python/relax/test_transform_gradient_te_register.py index 95fbd2c6a65c..90e4a94c0183 100644 --- a/tests/python/relax/test_transform_gradient_te_register.py +++ b/tests/python/relax/test_transform_gradient_te_register.py @@ -64,7 +64,7 @@ def get_expected_1(): @I.ir_module class Expected: @Ts.prim_func(private=True) - def f_mul(A: T.Buffer((T.int64(5), T.int64(5)), "float32"), B: T.Buffer((T.int64(5), T.int64(5)), "float32"), f_mul_1: T.Buffer((T.int64(5), T.int64(5)), "float32")): + def f_mul(A: T.Tensor((T.int64(5), T.int64(5)), "float32"), B: T.Tensor((T.int64(5), T.int64(5)), "float32"), f_mul_1: T.Tensor((T.int64(5), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1 in T.grid(T.int64(5), T.int64(5)): @@ -75,7 +75,7 @@ def f_mul(A: T.Buffer((T.int64(5), T.int64(5)), "float32"), B: T.Buffer((T.int64 f_mul_1[v_i0, v_i1] = A[v_i0, v_i1] * B[v_i0, v_i1] @Ts.prim_func(private=True) - def f_mul_grad(A: T.Buffer((T.int64(5), T.int64(5)), "float32"), B: T.Buffer((T.int64(5), T.int64(5)), "float32"), C: T.Buffer((T.int64(5), T.int64(5)), "float32"), f_mul_grad_1: T.Buffer((T.int64(5), T.int64(5)), "float32"), f_mul_grad_2: T.Buffer((T.int64(5), T.int64(5)), "float32")): + def f_mul_grad(A: T.Tensor((T.int64(5), T.int64(5)), "float32"), B: T.Tensor((T.int64(5), T.int64(5)), "float32"), C: T.Tensor((T.int64(5), T.int64(5)), "float32"), f_mul_grad_1: T.Tensor((T.int64(5), T.int64(5)), "float32"), f_mul_grad_2: T.Tensor((T.int64(5), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1 in T.grid(T.int64(5), T.int64(5)): @@ -151,7 +151,7 @@ def test_call_tir(register_te_grads): @I.ir_module class Before: @Ts.prim_func(private=True) - def f_mul(A: T.Buffer((T.int64(5), T.int64(5)), "float32"), B: T.Buffer((T.int64(5), T.int64(5)), "float32"), f_mul_1: T.Buffer((T.int64(5), T.int64(5)), "float32")): + def f_mul(A: T.Tensor((T.int64(5), T.int64(5)), "float32"), B: T.Tensor((T.int64(5), T.int64(5)), "float32"), f_mul_1: T.Tensor((T.int64(5), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1 in T.grid(T.int64(5), T.int64(5)): @@ -180,7 +180,7 @@ def get_expected_2(): @I.ir_module class Expected: @Ts.prim_func(private=True) - def f_mul(A: T.Buffer((T.int64(5), T.int64(5)), "float32"), f_mul2: T.Buffer((T.int64(5), T.int64(5)), "float32")): + def f_mul(A: T.Tensor((T.int64(5), T.int64(5)), "float32"), f_mul2: T.Tensor((T.int64(5), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1 in T.grid(T.int64(5), T.int64(5)): @@ -191,7 +191,7 @@ def f_mul(A: T.Buffer((T.int64(5), T.int64(5)), "float32"), f_mul2: T.Buffer((T. f_mul2[v_i0, v_i1] = A[v_i0, v_i1] * T.float32(2) @Ts.prim_func(private=True) - def f_mulk_grad(A: T.Buffer((T.int64(5), T.int64(5)), "float32"), B: T.Buffer((T.int64(5), T.int64(5)), "float32"), f_mulk_grad_1: T.Buffer((T.int64(5), T.int64(5)), "float32")): + def f_mulk_grad(A: T.Tensor((T.int64(5), T.int64(5)), "float32"), B: T.Tensor((T.int64(5), T.int64(5)), "float32"), f_mulk_grad_1: T.Tensor((T.int64(5), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1 in T.grid(T.int64(5), T.int64(5)): @@ -260,7 +260,7 @@ def test_call_tir_kwargs(register_te_grads): @I.ir_module class Before: @Ts.prim_func(private=True) - def f_mul(A: T.Buffer((T.int64(5), T.int64(5)), "float32"), f_mul2: T.Buffer((T.int64(5), T.int64(5)), "float32")): + def f_mul(A: T.Tensor((T.int64(5), T.int64(5)), "float32"), f_mul2: T.Tensor((T.int64(5), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1 in T.grid(T.int64(5), T.int64(5)): @@ -294,7 +294,7 @@ def get_expected_3(): @I.ir_module class Expected: @Ts.prim_func(private=True) - def f_mul(A: T.Buffer((n_f_mul, n_f_mul)), B: T.Buffer((n_f_mul, n_f_mul)), f_mul_1: T.Buffer((n_f_mul, n_f_mul))): + def f_mul(A: T.Tensor((n_f_mul, n_f_mul)), B: T.Tensor((n_f_mul, n_f_mul)), f_mul_1: T.Tensor((n_f_mul, n_f_mul))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -306,7 +306,7 @@ def f_mul(A: T.Buffer((n_f_mul, n_f_mul)), B: T.Buffer((n_f_mul, n_f_mul)), f_mu f_mul_1[v_i0, v_i1] = A[v_i0, v_i1] * B[v_i0, v_i1] @Ts.prim_func(private=True) - def f_mul_grad(A: T.Buffer((n_f_mul_grad, n_f_mul_grad)), B: T.Buffer((n_f_mul_grad, n_f_mul_grad)), C: T.Buffer((n_f_mul_grad, n_f_mul_grad)), f_mul_grad_1: T.Buffer((n_f_mul_grad, n_f_mul_grad)), f_mul_grad_2: T.Buffer((n_f_mul_grad, n_f_mul_grad))): + def f_mul_grad(A: T.Tensor((n_f_mul_grad, n_f_mul_grad)), B: T.Tensor((n_f_mul_grad, n_f_mul_grad)), C: T.Tensor((n_f_mul_grad, n_f_mul_grad)), f_mul_grad_1: T.Tensor((n_f_mul_grad, n_f_mul_grad)), f_mul_grad_2: T.Tensor((n_f_mul_grad, n_f_mul_grad))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): diff --git a/tests/python/relax/test_transform_lambda_lift.py b/tests/python/relax/test_transform_lambda_lift.py index 22ed281a2038..b2496f635751 100644 --- a/tests/python/relax/test_transform_lambda_lift.py +++ b/tests/python/relax/test_transform_lambda_lift.py @@ -332,9 +332,9 @@ def test_no_local_func(): class Before: @Ts.prim_func def sub( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: for i, j in T.grid(16, 16): with Ts.sblock("sub"): diff --git a/tests/python/relax/test_transform_lazy_transform_params.py b/tests/python/relax/test_transform_lazy_transform_params.py index 1feeef94d181..78c4708e3457 100644 --- a/tests/python/relax/test_transform_lazy_transform_params.py +++ b/tests/python/relax/test_transform_lazy_transform_params.py @@ -32,7 +32,7 @@ def test_lazy_transform_params(): class Before: @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((3, 16, 3, 3), "float32"), out: T.Buffer((16, 3, 3, 3), "float32") + w1: T.Tensor((3, 16, 3, 3), "float32"), out: T.Tensor((16, 3, 3, 3), "float32") ): for ax0, ax1, ax2, ax3 in T.grid(16, 3, 3, 3): with Ts.sblock("layout_transform"): @@ -69,7 +69,7 @@ def main_transform_params( class Expected: @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((3, 16, 3, 3), "float32"), out: T.Buffer((16, 3, 3, 3), "float32") + w1: T.Tensor((3, 16, 3, 3), "float32"), out: T.Tensor((16, 3, 3, 3), "float32") ): # with Ts.sblock("root"): for ax0, ax1, ax2, ax3 in T.grid(16, 3, 3, 3): @@ -113,7 +113,7 @@ def test_get_item_only(): class Before: @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((3, 16, 3, 3), "float32"), out: T.Buffer((16, 3, 3, 3), "float32") + w1: T.Tensor((3, 16, 3, 3), "float32"), out: T.Tensor((16, 3, 3, 3), "float32") ): for ax0, ax1, ax2, ax3 in T.grid(16, 3, 3, 3): with Ts.sblock("layout_transform"): @@ -151,7 +151,7 @@ def main_transform_params( class Expected: @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((3, 16, 3, 3), "float32"), out: T.Buffer((16, 3, 3, 3), "float32") + w1: T.Tensor((3, 16, 3, 3), "float32"), out: T.Tensor((16, 3, 3, 3), "float32") ): # with Ts.sblock("root"): for ax0, ax1, ax2, ax3 in T.grid(16, 3, 3, 3): @@ -196,7 +196,7 @@ def test_extra_get_item_params(): class Before: @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((3, 16, 3, 3), "float32"), out: T.Buffer((16, 3, 3, 3), "float32") + w1: T.Tensor((3, 16, 3, 3), "float32"), out: T.Tensor((16, 3, 3, 3), "float32") ): for ax0, ax1, ax2, ax3 in T.grid(16, 3, 3, 3): with Ts.sblock("layout_transform"): @@ -234,7 +234,7 @@ def main_transform_params( class Expected: @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((3, 16, 3, 3), "float32"), out: T.Buffer((16, 3, 3, 3), "float32") + w1: T.Tensor((3, 16, 3, 3), "float32"), out: T.Tensor((16, 3, 3, 3), "float32") ): # with Ts.sblock("root"): for ax0, ax1, ax2, ax3 in T.grid(16, 3, 3, 3): @@ -281,7 +281,7 @@ def test_extra_set_item_params(): class Before: @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((3, 16, 3, 3), "float32"), out: T.Buffer((16, 3, 3, 3), "float32") + w1: T.Tensor((3, 16, 3, 3), "float32"), out: T.Tensor((16, 3, 3, 3), "float32") ): for ax0, ax1, ax2, ax3 in T.grid(16, 3, 3, 3): with Ts.sblock("layout_transform"): @@ -319,7 +319,7 @@ def main_transform_params( class Expected: @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((3, 16, 3, 3), "float32"), out: T.Buffer((16, 3, 3, 3), "float32") + w1: T.Tensor((3, 16, 3, 3), "float32"), out: T.Tensor((16, 3, 3, 3), "float32") ): # with Ts.sblock("root"): for ax0, ax1, ax2, ax3 in T.grid(16, 3, 3, 3): @@ -431,9 +431,9 @@ def main_transform_params( @Ts.prim_func(private=True) def slice_buffer( - Input: T.Buffer((16, 16), "float32"), + Input: T.Tensor((16, 16), "float32"), slice_index: T.int64, - Output: T.Buffer(16, "float32"), + Output: T.Tensor(16, "float32"), ): for i in T.grid(16): with Ts.sblock("slice_buffer"): @@ -466,9 +466,9 @@ def main_transform_params(slice_shape_expr: R.Shape([slice_index])): @Ts.prim_func(private=True) def slice_buffer( - Input: T.Buffer((16, 16), "float32"), + Input: T.Tensor((16, 16), "float32"), slice_index: T.int64, - Output: T.Buffer(16, "float32"), + Output: T.Tensor(16, "float32"), ): for i in T.grid(16): with Ts.sblock("slice_buffer"): @@ -487,8 +487,8 @@ def test_param_shape_symbolic(): class Before: @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((ic_transform_layout_IOHW_to_OIHW, 16, 3, 3), "float32"), - out: T.Buffer((16, ic_transform_layout_IOHW_to_OIHW, 3, 3), "float32"), + w1: T.Tensor((ic_transform_layout_IOHW_to_OIHW, 16, 3, 3), "float32"), + out: T.Tensor((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"): @@ -530,8 +530,8 @@ def main_transform_params( class Expected: @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((ic_transform_layout_IOHW_to_OIHW, 16, 3, 3), "float32"), - out: T.Buffer((16, ic_transform_layout_IOHW_to_OIHW, 3, 3), "float32"), + w1: T.Tensor((ic_transform_layout_IOHW_to_OIHW, 16, 3, 3), "float32"), + out: T.Tensor((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"): @@ -573,7 +573,7 @@ def test_output_with_use_site(): @I.ir_module class Module: @Ts.prim_func - def copy(x: T.Buffer((), "float32"), y: T.Buffer((), "float32")): + def copy(x: T.Tensor((), "float32"), y: T.Tensor((), "float32")): with Ts.sblock("block"): Ts.reads(x[()]) Ts.writes(y[()]) @@ -595,7 +595,7 @@ def main_transform_params(params: R.Tuple(R.Tensor((), dtype="float32"))) -> R.T @I.ir_module class Expected: @Ts.prim_func - def copy(x: T.Buffer((), "float32"), y: T.Buffer((), "float32")): + def copy(x: T.Tensor((), "float32"), y: T.Tensor((), "float32")): with Ts.sblock("block"): Ts.reads(x[()]) Ts.writes(y[()]) diff --git a/tests/python/relax/test_transform_legalize_ops.py b/tests/python/relax/test_transform_legalize_ops.py index b603b0a3d346..ac3a92fd71eb 100644 --- a/tests/python/relax/test_transform_legalize_ops.py +++ b/tests/python/relax/test_transform_legalize_ops.py @@ -50,7 +50,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def add(rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), T_add: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def add(rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), T_add: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_add"): @@ -79,7 +79,7 @@ def mul2(x: R.Tensor((3, 3), "float32")): return gv @Ts.prim_func(private=True) - def identity(rxplaceholder: T.Buffer((T.int64(3), T.int64(3)), "float32"), T_id: T.Buffer((T.int64(3), T.int64(3)), "float32")): + def identity(rxplaceholder: T.Tensor((T.int64(3), T.int64(3)), "float32"), T_id: T.Tensor((T.int64(3), T.int64(3)), "float32")): for ax0, ax1 in T.grid(T.int64(3), T.int64(3)): with Ts.sblock("T_add"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) @@ -104,7 +104,7 @@ def mul2(x: R.Tensor((3, 3), dtype="float32")) -> R.Tensor((3, 3), dtype="float3 return gv @Ts.prim_func(private=True) - def identity(rxplaceholder: T.Buffer((T.int64(3), T.int64(3)), "float32"), T_id: T.Buffer((T.int64(3), T.int64(3)), "float32")): + def identity(rxplaceholder: T.Tensor((T.int64(3), T.int64(3)), "float32"), T_id: T.Tensor((T.int64(3), T.int64(3)), "float32")): for ax0, ax1 in T.grid(T.int64(3), T.int64(3)): with Ts.sblock("T_add"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) @@ -113,7 +113,7 @@ def identity(rxplaceholder: T.Buffer((T.int64(3), T.int64(3)), "float32"), T_id: T_id[v_ax0, v_ax1] = rxplaceholder[v_ax0, v_ax1] @Ts.prim_func(private=True) - def multiply(rxplaceholder: T.Buffer((T.int64(3), T.int64(3)), "float32"), T_multiply: T.Buffer((T.int64(3), T.int64(3)), "float32")): + def multiply(rxplaceholder: T.Tensor((T.int64(3), T.int64(3)), "float32"), T_multiply: T.Tensor((T.int64(3), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(3), T.int64(3)): with Ts.sblock("T_multiply"): @@ -195,8 +195,8 @@ def main(x: R.Tensor((3, 3), "bool")): class Expected0: @Ts.prim_func(private=True) def multiply( - rxplaceholder: T.Buffer((T.int64(3), T.int64(3)), "float16"), - T_multiply: T.Buffer((T.int64(3), T.int64(3)), "float16"), + rxplaceholder: T.Tensor((T.int64(3), T.int64(3)), "float16"), + T_multiply: T.Tensor((T.int64(3), T.int64(3)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -219,8 +219,8 @@ def main(x: R.Tensor((3, 3), dtype="float16")) -> R.Tensor((3, 3), dtype="float1 class Expected1: @Ts.prim_func(private=True) def multiply( - rxplaceholder: T.Buffer((T.int64(3), T.int64(3)), "uint8"), - T_multiply: T.Buffer((T.int64(3), T.int64(3)), "uint8"), + rxplaceholder: T.Tensor((T.int64(3), T.int64(3)), "uint8"), + T_multiply: T.Tensor((T.int64(3), T.int64(3)), "uint8"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -241,8 +241,8 @@ def main(x: R.Tensor((3, 3), dtype="uint8")) -> R.Tensor((3, 3), dtype="uint8"): class Expected2: @Ts.prim_func(private=True) def equal( - rxplaceholder: T.Buffer((T.int64(3), T.int64(3)), "bool"), - T_equal: T.Buffer((T.int64(3), T.int64(3)), "bool"), + rxplaceholder: T.Tensor((T.int64(3), T.int64(3)), "bool"), + T_equal: T.Tensor((T.int64(3), T.int64(3)), "bool"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -402,9 +402,9 @@ def func_cuda( @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(32), T.int64(32)), "float32"), - B: T.Buffer((T.int64(32), T.int64(32)), "float32"), - C: T.Buffer((T.int64(32), T.int64(32)), "float32"), + A: T.Tensor((T.int64(32), T.int64(32)), "float32"), + B: T.Tensor((T.int64(32), T.int64(32)), "float32"), + C: T.Tensor((T.int64(32), T.int64(32)), "float32"), ): T.func_attr({"tirx.noalias": True}) for (*iters,) in T.grid(T.int64(32), T.int64(32)): @@ -427,9 +427,9 @@ def func_llvm( @Ts.prim_func(private=True) def add_llvm( - A: T.Buffer((T.int64(32), T.int64(32)), "float32"), - B: T.Buffer((T.int64(32), T.int64(32)), "float32"), - C: T.Buffer((T.int64(32), T.int64(32)), "float32"), + A: T.Tensor((T.int64(32), T.int64(32)), "float32"), + B: T.Tensor((T.int64(32), T.int64(32)), "float32"), + C: T.Tensor((T.int64(32), T.int64(32)), "float32"), ): T.func_attr({"target": T.target("llvm"), "tirx.noalias": True}) for (*iters,) in T.grid(T.int64(32), T.int64(32)): diff --git a/tests/python/relax/test_transform_legalize_ops_binary.py b/tests/python/relax/test_transform_legalize_ops_binary.py index e1a424a2beb8..4be9a041b44a 100644 --- a/tests/python/relax/test_transform_legalize_ops_binary.py +++ b/tests/python/relax/test_transform_legalize_ops_binary.py @@ -44,7 +44,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def add(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_add: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def add(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_add: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_add"): @@ -75,7 +75,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def add(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_add: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def add(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_add: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_add"): @@ -106,7 +106,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def add(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_add: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def add(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_add: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_add"): @@ -151,7 +151,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def add(rxplaceholder: T.Buffer([T.int64(1), c_add, d_add], dtype='float32'), rxplaceholder_1: T.Buffer([a_add, b_add, c_add, T.int64(1)], dtype='float32'), T_add: T.Buffer([a_add, b_add, c_add, d_add], dtype='float32')): + def add(rxplaceholder: T.Tensor([T.int64(1), c_add, d_add], dtype='float32'), rxplaceholder_1: T.Tensor([a_add, b_add, c_add, T.int64(1)], dtype='float32'), T_add: T.Tensor([a_add, b_add, c_add, d_add], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_add, b_add, c_add, d_add): @@ -190,9 +190,9 @@ def main( @Ts.prim_func(private=True) def add( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -221,7 +221,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def divide(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_divide: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def divide(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_divide: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_divide"): @@ -252,7 +252,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def divide(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_divide: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def divide(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_divide: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_divide"): @@ -283,7 +283,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def divide(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_divide: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def divide(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_divide: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_divide"): @@ -328,7 +328,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def divide(rxplaceholder: T.Buffer([T.int64(1), c_divide, d_divide], dtype='float32'), rxplaceholder_1: T.Buffer([a_divide, b_divide, c_divide, T.int64(1)], dtype='float32'), T_divide: T.Buffer([a_divide, b_divide, c_divide, d_divide], dtype='float32')): + def divide(rxplaceholder: T.Tensor([T.int64(1), c_divide, d_divide], dtype='float32'), rxplaceholder_1: T.Tensor([a_divide, b_divide, c_divide, T.int64(1)], dtype='float32'), T_divide: T.Tensor([a_divide, b_divide, c_divide, d_divide], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_divide, b_divide, c_divide, d_divide): @@ -367,9 +367,9 @@ def main( @Ts.prim_func(private=True) def divide( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -398,7 +398,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def floor_divide(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_floor_divide: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def floor_divide(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_floor_divide: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_floor_divide"): @@ -429,7 +429,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def floor_divide(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_floor_divide: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def floor_divide(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_floor_divide: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_floor_divide"): @@ -460,7 +460,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def floor_divide(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_floor_divide: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def floor_divide(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_floor_divide: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_floor_divide"): @@ -505,7 +505,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def floor_divide(rxplaceholder: T.Buffer([T.int64(1), c_floor_divide, d_floor_divide], dtype='float32'), rxplaceholder_1: T.Buffer([a_floor_divide, b_floor_divide, c_floor_divide, T.int64(1)], dtype='float32'), T_floor_divide: T.Buffer([a_floor_divide, b_floor_divide, c_floor_divide, d_floor_divide], dtype='float32')): + def floor_divide(rxplaceholder: T.Tensor([T.int64(1), c_floor_divide, d_floor_divide], dtype='float32'), rxplaceholder_1: T.Tensor([a_floor_divide, b_floor_divide, c_floor_divide, T.int64(1)], dtype='float32'), T_floor_divide: T.Tensor([a_floor_divide, b_floor_divide, c_floor_divide, d_floor_divide], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_floor_divide, b_floor_divide, c_floor_divide, d_floor_divide): @@ -544,9 +544,9 @@ def main( @Ts.prim_func(private=True) def floor_divide( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -575,7 +575,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def multiply(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_multiply: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def multiply(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_multiply: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_multiply"): @@ -620,7 +620,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def multiply(rxplaceholder: T.Buffer([T.int64(1), c_multiply, d_multiply], dtype='float32'), rxplaceholder_1: T.Buffer([a_multiply, b_multiply, c_multiply, T.int64(1)], dtype='float32'), T_multiply: T.Buffer([a_multiply, b_multiply, c_multiply, d_multiply], dtype='float32')): + def multiply(rxplaceholder: T.Tensor([T.int64(1), c_multiply, d_multiply], dtype='float32'), rxplaceholder_1: T.Tensor([a_multiply, b_multiply, c_multiply, T.int64(1)], dtype='float32'), T_multiply: T.Tensor([a_multiply, b_multiply, c_multiply, d_multiply], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_multiply, b_multiply, c_multiply, d_multiply): @@ -659,9 +659,9 @@ def main( @Ts.prim_func(private=True) def multiply( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -685,7 +685,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def power(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_power: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def power(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_power: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1, ax2, ax3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): @@ -732,7 +732,7 @@ def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32") @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def power(rxplaceholder: T.Buffer((T.int64(1), c_power, d_power)), rxplaceholder_1: T.Buffer((a_power, b_power, c_power, T.int64(1))), T_power: T.Buffer((a_power, b_power, c_power, d_power))): + def power(rxplaceholder: T.Tensor((T.int64(1), c_power, d_power)), rxplaceholder_1: T.Tensor((a_power, b_power, c_power, T.int64(1))), T_power: T.Tensor((a_power, b_power, c_power, d_power))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -777,9 +777,9 @@ def main( @Ts.prim_func(private=True) def power( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -803,7 +803,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def atan2(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_atan2: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def atan2(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_atan2: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1, ax2, ax3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): @@ -850,7 +850,7 @@ def main(x: R.Tensor((1, c, d), "float32"), y: R.Tensor((a, b, c, 1), "float32") @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def atan2(rxplaceholder: T.Buffer((T.int64(1), c_atan2, d_atan2)), rxplaceholder_1: T.Buffer((a_atan2, b_atan2, c_atan2, T.int64(1))), T_atan2: T.Buffer((a_atan2, b_atan2, c_atan2, d_atan2))): + def atan2(rxplaceholder: T.Tensor((T.int64(1), c_atan2, d_atan2)), rxplaceholder_1: T.Tensor((a_atan2, b_atan2, c_atan2, T.int64(1))), T_atan2: T.Tensor((a_atan2, b_atan2, c_atan2, d_atan2))): T.func_attr({"tirx.noalias": True}) for ax0, ax1, ax2, ax3 in T.grid(a_atan2, b_atan2, c_atan2, d_atan2): @@ -894,9 +894,9 @@ def main( @Ts.prim_func(private=True) def atan2( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -925,7 +925,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def subtract(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_subtract: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def subtract(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_subtract: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_subtract"): @@ -970,7 +970,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def subtract(rxplaceholder: T.Buffer([T.int64(1), c_subtract, d_subtract], dtype='float32'), rxplaceholder_1: T.Buffer([a_subtract, b_subtract, c_subtract, T.int64(1)], dtype='float32'), T_subtract: T.Buffer([a_subtract, b_subtract, c_subtract, d_subtract], dtype='float32')): + def subtract(rxplaceholder: T.Tensor([T.int64(1), c_subtract, d_subtract], dtype='float32'), rxplaceholder_1: T.Tensor([a_subtract, b_subtract, c_subtract, T.int64(1)], dtype='float32'), T_subtract: T.Tensor([a_subtract, b_subtract, c_subtract, d_subtract], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_subtract, b_subtract, c_subtract, d_subtract): @@ -1009,9 +1009,9 @@ def main( @Ts.prim_func(private=True) def subtract( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -1043,7 +1043,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def equal(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_equal: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): + def equal(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_equal: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_equal"): @@ -1074,7 +1074,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "bool"): return gv @Ts.prim_func(private=True) - def equal(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_equal: T.Buffer((T.int64(2), T.int64(3)), "bool")): + def equal(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_equal: T.Tensor((T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_equal"): @@ -1105,7 +1105,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "bool"): return gv @Ts.prim_func(private=True) - def equal(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_equal: T.Buffer((T.int64(2), T.int64(3)), "bool")): + def equal(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_equal: T.Tensor((T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_equal"): @@ -1150,7 +1150,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def equal(rxplaceholder: T.Buffer([T.int64(1), c_equal, d_equal], dtype='float32'), rxplaceholder_1: T.Buffer([a_equal, b_equal, c_equal, T.int64(1)], dtype='float32'), T_equal: T.Buffer([a_equal, b_equal, c_equal, d_equal], dtype='bool')): + def equal(rxplaceholder: T.Tensor([T.int64(1), c_equal, d_equal], dtype='float32'), rxplaceholder_1: T.Tensor([a_equal, b_equal, c_equal, T.int64(1)], dtype='float32'), T_equal: T.Tensor([a_equal, b_equal, c_equal, d_equal], dtype='bool')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_equal, b_equal, c_equal, d_equal): @@ -1189,9 +1189,9 @@ def main( @Ts.prim_func(private=True) def equal( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "bool"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "bool"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -1220,7 +1220,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def greater(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_greater: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): + def greater(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_greater: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_greater"): @@ -1251,7 +1251,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "bool"): return gv @Ts.prim_func(private=True) - def greater(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_greater: T.Buffer((T.int64(2), T.int64(3)), "bool")): + def greater(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_greater: T.Tensor((T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_greater"): @@ -1282,7 +1282,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "bool"): return gv @Ts.prim_func(private=True) - def greater(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_greater: T.Buffer((T.int64(2), T.int64(3)), "bool")): + def greater(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_greater: T.Tensor((T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_greater"): @@ -1327,7 +1327,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def greater(rxplaceholder: T.Buffer([T.int64(1), c_greater, d_greater], dtype='float32'), rxplaceholder_1: T.Buffer([a_greater, b_greater, c_greater, T.int64(1)], dtype='float32'), T_greater: T.Buffer([a_greater, b_greater, c_greater, d_greater], dtype='bool')): + def greater(rxplaceholder: T.Tensor([T.int64(1), c_greater, d_greater], dtype='float32'), rxplaceholder_1: T.Tensor([a_greater, b_greater, c_greater, T.int64(1)], dtype='float32'), T_greater: T.Tensor([a_greater, b_greater, c_greater, d_greater], dtype='bool')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_greater, b_greater, c_greater, d_greater): @@ -1366,9 +1366,9 @@ def main( @Ts.prim_func(private=True) def greater( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "bool"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "bool"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -1397,7 +1397,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def greater_equal(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_greater_equal: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): + def greater_equal(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_greater_equal: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_greater_equal"): @@ -1442,7 +1442,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def greater_equal(rxplaceholder: T.Buffer([T.int64(1), c_greater_equal, d_greater_equal], dtype='float32'), rxplaceholder_1: T.Buffer([a_greater_equal, b_greater_equal, c_greater_equal, T.int64(1)], dtype='float32'), T_greater_equal: T.Buffer([a_greater_equal, b_greater_equal, c_greater_equal, d_greater_equal], dtype='bool')): + def greater_equal(rxplaceholder: T.Tensor([T.int64(1), c_greater_equal, d_greater_equal], dtype='float32'), rxplaceholder_1: T.Tensor([a_greater_equal, b_greater_equal, c_greater_equal, T.int64(1)], dtype='float32'), T_greater_equal: T.Tensor([a_greater_equal, b_greater_equal, c_greater_equal, d_greater_equal], dtype='bool')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_greater_equal, b_greater_equal, c_greater_equal, d_greater_equal): @@ -1481,9 +1481,9 @@ def main( @Ts.prim_func(private=True) def greater_equal( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "bool"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "bool"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -1512,7 +1512,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def less(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_less: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): + def less(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_less: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_less"): @@ -1557,7 +1557,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def less(rxplaceholder: T.Buffer([T.int64(1), c_less, d_less], dtype='float32'), rxplaceholder_1: T.Buffer([a_less, b_less, c_less, T.int64(1)], dtype='float32'), T_less: T.Buffer([a_less, b_less, c_less, d_less], dtype='bool')): + def less(rxplaceholder: T.Tensor([T.int64(1), c_less, d_less], dtype='float32'), rxplaceholder_1: T.Tensor([a_less, b_less, c_less, T.int64(1)], dtype='float32'), T_less: T.Tensor([a_less, b_less, c_less, d_less], dtype='bool')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_less, b_less, c_less, d_less): @@ -1596,9 +1596,9 @@ def main( @Ts.prim_func(private=True) def less( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "bool"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "bool"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -1627,7 +1627,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def less_equal(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_less_equal: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): + def less_equal(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_less_equal: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_less_equal"): @@ -1658,7 +1658,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "bool"): return gv @Ts.prim_func(private=True) - def less_equal(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_less_equal: T.Buffer((T.int64(2), T.int64(3)), "bool")): + def less_equal(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_less_equal: T.Tensor((T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_less_equal"): @@ -1689,7 +1689,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "bool"): return gv @Ts.prim_func(private=True) - def less_equal(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_less_equal: T.Buffer((T.int64(2), T.int64(3)), "bool")): + def less_equal(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_less_equal: T.Tensor((T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_less_equal"): @@ -1734,7 +1734,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def less_equal(rxplaceholder: T.Buffer([T.int64(1), c_less_equal, d_less_equal], dtype='float32'), rxplaceholder_1: T.Buffer([a_less_equal, b_less_equal, c_less_equal, T.int64(1)], dtype='float32'), T_less_equal: T.Buffer([a_less_equal, b_less_equal, c_less_equal, d_less_equal], dtype='bool')): + def less_equal(rxplaceholder: T.Tensor([T.int64(1), c_less_equal, d_less_equal], dtype='float32'), rxplaceholder_1: T.Tensor([a_less_equal, b_less_equal, c_less_equal, T.int64(1)], dtype='float32'), T_less_equal: T.Tensor([a_less_equal, b_less_equal, c_less_equal, d_less_equal], dtype='bool')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_less_equal, b_less_equal, c_less_equal, d_less_equal): @@ -1773,9 +1773,9 @@ def main( @Ts.prim_func(private=True) def less_equal( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "bool"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "bool"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -1804,7 +1804,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def not_equal(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_not_equal: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): + def not_equal(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_not_equal: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_not_equal"): @@ -1849,7 +1849,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def not_equal(rxplaceholder: T.Buffer([T.int64(1), c_not_equal, d_not_equal], dtype='float32'), rxplaceholder_1: T.Buffer([a_not_equal, b_not_equal, c_not_equal, T.int64(1)], dtype='float32'), T_not_equal: T.Buffer([a_not_equal, b_not_equal, c_not_equal, d_not_equal], dtype='bool')): + def not_equal(rxplaceholder: T.Tensor([T.int64(1), c_not_equal, d_not_equal], dtype='float32'), rxplaceholder_1: T.Tensor([a_not_equal, b_not_equal, c_not_equal, T.int64(1)], dtype='float32'), T_not_equal: T.Tensor([a_not_equal, b_not_equal, c_not_equal, d_not_equal], dtype='bool')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_not_equal, b_not_equal, c_not_equal, d_not_equal): @@ -1888,9 +1888,9 @@ def main( @Ts.prim_func(private=True) def not_equal( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "bool"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "bool"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -1919,7 +1919,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def maximum(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_maximum: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def maximum(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_maximum: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_maximum"): @@ -1950,7 +1950,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def maximum(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_maximum: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def maximum(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_maximum: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_maximum"): @@ -1981,7 +1981,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def maximum(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_maximum: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def maximum(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_maximum: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_maximum"): @@ -2026,7 +2026,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def maximum(rxplaceholder: T.Buffer([T.int64(1), c_maximum, d_maximum], dtype='float32'), rxplaceholder_1: T.Buffer([a_maximum, b_maximum, c_maximum, T.int64(1)], dtype='float32'), T_maximum: T.Buffer([a_maximum, b_maximum, c_maximum, d_maximum], dtype='float32')): + def maximum(rxplaceholder: T.Tensor([T.int64(1), c_maximum, d_maximum], dtype='float32'), rxplaceholder_1: T.Tensor([a_maximum, b_maximum, c_maximum, T.int64(1)], dtype='float32'), T_maximum: T.Tensor([a_maximum, b_maximum, c_maximum, d_maximum], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_maximum, b_maximum, c_maximum, d_maximum): @@ -2065,9 +2065,9 @@ def main( @Ts.prim_func(private=True) def maximum( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): @@ -2096,7 +2096,7 @@ def main(x: R.Tensor((1, 2, 3), "float32"), y: R.Tensor((4, 3, 2, 1), "float32") return gv @Ts.prim_func(private=True) - def minimum(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_minimum: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def minimum(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_minimum: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_minimum"): @@ -2127,7 +2127,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def minimum(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_minimum: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def minimum(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_minimum: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_minimum"): @@ -2158,7 +2158,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def minimum(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_minimum: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def minimum(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_minimum: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_minimum"): @@ -2203,7 +2203,7 @@ def main(x: R.Tensor((1, c_main, d_main), "float32"), y: R.Tensor((a_main, b_mai return gv @Ts.prim_func(private=True) - def minimum(rxplaceholder: T.Buffer([T.int64(1), c_minimum, d_minimum], dtype='float32'), rxplaceholder_1: T.Buffer([a_minimum, b_minimum, c_minimum, T.int64(1)], dtype='float32'), T_minimum: T.Buffer([a_minimum, b_minimum, c_minimum, d_minimum], dtype='float32')): + def minimum(rxplaceholder: T.Tensor([T.int64(1), c_minimum, d_minimum], dtype='float32'), rxplaceholder_1: T.Tensor([a_minimum, b_minimum, c_minimum, T.int64(1)], dtype='float32'), T_minimum: T.Tensor([a_minimum, b_minimum, c_minimum, d_minimum], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_minimum, b_minimum, c_minimum, d_minimum): @@ -2242,9 +2242,9 @@ def main( @Ts.prim_func(private=True) def minimum( - lhs: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + lhs: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), rhs: T.float32, - output: T.Buffer([T.int64(64), T.int64(32), T.int64(16)], "float32"), + output: T.Tensor([T.int64(64), T.int64(32), T.int64(16)], "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j, k in T.grid(*lhs.shape): diff --git a/tests/python/relax/test_transform_legalize_ops_ccl.py b/tests/python/relax/test_transform_legalize_ops_ccl.py index 18ddd88dcb1a..418c93a21fe1 100644 --- a/tests/python/relax/test_transform_legalize_ops_ccl.py +++ b/tests/python/relax/test_transform_legalize_ops_ccl.py @@ -110,7 +110,7 @@ def main(x: R.Tensor((10, 10), "float32")) -> R.Tensor((10,5), "float32"): @I.ir_module class Expected: @Ts.prim_func(private=True) - def reshape(A: T.Buffer((T.int64(10), T.int64(10)), "float32"), T_reshape: T.Buffer((T.int64(10), T.int64(2), T.int64(5)), "float32")): + def reshape(A: T.Tensor((T.int64(10), T.int64(10)), "float32"), T_reshape: T.Tensor((T.int64(10), T.int64(2), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1, ax2 in T.grid(T.int64(10), T.int64(2), T.int64(5)): @@ -121,7 +121,7 @@ def reshape(A: T.Buffer((T.int64(10), T.int64(10)), "float32"), T_reshape: T.Buf T_reshape[v_ax0, v_ax1, v_ax2] = A[((v_ax1 * T.int64(5) + v_ax2) // T.int64(10) + v_ax0) % T.int64(10), (v_ax1 * T.int64(5) + v_ax2) % T.int64(10)] @Ts.prim_func(private=True) - def transpose(A: T.Buffer((T.int64(10), T.int64(2), T.int64(5)), "float32"), T_transpose: T.Buffer((T.int64(2), T.int64(10), T.int64(5)), "float32")): + def transpose(A: T.Tensor((T.int64(10), T.int64(2), T.int64(5)), "float32"), T_transpose: T.Tensor((T.int64(2), T.int64(10), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1, ax2 in T.grid(T.int64(2), T.int64(10), T.int64(5)): 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 b31bff855834..c5e32bc9bd1a 100644 --- a/tests/python/relax/test_transform_legalize_ops_create_datatype.py +++ b/tests/python/relax/test_transform_legalize_ops_create_datatype.py @@ -43,7 +43,7 @@ def main(v: R.Tensor((), "int32")) -> R.Tensor((2, 3), "int32"): return gv @Ts.prim_func(private=True) - def full(rxplaceholder: T.Buffer((), "int32"), T_full: T.Buffer((T.int64(2), T.int64(3)), "int32")): + def full(rxplaceholder: T.Tensor((), "int32"), T_full: T.Tensor((T.int64(2), T.int64(3)), "int32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -74,7 +74,7 @@ def main() -> R.Tensor((2, 3), "int32"): return gv @Ts.prim_func(private=True) - def full(T_full: T.Buffer((T.int64(2), T.int64(3)), "int32")): + def full(T_full: T.Tensor((T.int64(2), T.int64(3)), "int32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -105,7 +105,7 @@ def main(v: R.Tensor((), "int32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def full(rxplaceholder: T.Buffer((), "int32"), T_full: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def full(rxplaceholder: T.Tensor((), "int32"), T_full: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -144,7 +144,7 @@ def main(dumb_param: R.Tensor((m_main, n_main)), v: R.Tensor((), "int32")) -> R. return gv @Ts.prim_func(private=True) - def full(rxplaceholder: T.Buffer((), "int32"), T_full: T.Buffer([m_full, n_full], dtype='int32')): + def full(rxplaceholder: T.Tensor((), "int32"), T_full: T.Tensor([m_full, n_full], dtype='int32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(m_full, n_full): @@ -176,7 +176,7 @@ def main(x: R.Tensor((2, 3), "int32"), v: R.Tensor((), "float32")) -> R.Tensor(( return gv @Ts.prim_func(private=True) - def full(rxplaceholder: T.Buffer((), "float32"), T_full: T.Buffer((T.int64(2), T.int64(3)), "int32")): + def full(rxplaceholder: T.Tensor((), "float32"), T_full: T.Tensor((T.int64(2), T.int64(3)), "int32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -207,7 +207,7 @@ def main(x: R.Tensor((2, 3), "int32")) -> R.Tensor((2, 3), "int32"): return gv @Ts.prim_func(private=True) - def full(T_full: T.Buffer((T.int64(2), T.int64(3)), "int32")): + def full(T_full: T.Tensor((T.int64(2), T.int64(3)), "int32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -238,7 +238,7 @@ def main(x: R.Tensor((2, 3), "int32"), v: R.Tensor((), "float32")) -> R.Tensor(( return gv @Ts.prim_func(private=True) - def full(rxplaceholder: T.Buffer((), "float32"), T_full: T.Buffer((T.int64(2), T.int64(3)), "float64")): + def full(rxplaceholder: T.Tensor((), "float32"), T_full: T.Tensor((T.int64(2), T.int64(3)), "float64")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -277,7 +277,7 @@ def main(x: R.Tensor((m_main, n_main), "int32"), v: R.Tensor((), "float32")) -> return gv @Ts.prim_func(private=True) - def full(rxplaceholder: T.Buffer((), "float32"), T_full: T.Buffer([m_full, n_full], dtype='int32')): + def full(rxplaceholder: T.Tensor((), "float32"), T_full: T.Tensor([m_full, n_full], dtype='int32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(m_full, n_full): @@ -309,7 +309,7 @@ def main() -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def ones(T_full: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def ones(T_full: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -348,7 +348,7 @@ def main(dumb_param: R.Tensor((m_main, n_main))) -> R.Tensor((m_main, n_main), " return gv @Ts.prim_func(private=True) - def ones(T_full: T.Buffer([m_ones, n_ones], dtype='float32')): + def ones(T_full: T.Tensor([m_ones, n_ones], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(m_ones, n_ones): @@ -380,7 +380,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "int32"): return gv @Ts.prim_func(private=True) - def ones(T_full: T.Buffer((T.int64(2), T.int64(3)), "int32")): + def ones(T_full: T.Tensor((T.int64(2), T.int64(3)), "int32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -419,7 +419,7 @@ def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), return gv @Ts.prim_func(private=True) - def ones(T_full: T.Buffer([m_ones, n_ones], dtype='float32')): + def ones(T_full: T.Tensor([m_ones, n_ones], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(m_ones, n_ones): @@ -451,7 +451,7 @@ def main() -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def zeros(T_full: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def zeros(T_full: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -490,7 +490,7 @@ def main(dumb_param: R.Tensor((m_main, n_main))) -> R.Tensor((m_main, n_main), " return gv @Ts.prim_func(private=True) - def zeros(T_full: T.Buffer([m_zeros, n_zeros], dtype='float32')): + def zeros(T_full: T.Tensor([m_zeros, n_zeros], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(m_zeros, n_zeros): @@ -522,7 +522,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "int32"): return gv @Ts.prim_func(private=True) - def zeros(T_full: T.Buffer((T.int64(2), T.int64(3)), "int32")): + def zeros(T_full: T.Tensor((T.int64(2), T.int64(3)), "int32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -561,7 +561,7 @@ def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), return gv @Ts.prim_func(private=True) - def zeros(T_full: T.Buffer([m_zeros, n_zeros], dtype='float32')): + def zeros(T_full: T.Tensor([m_zeros, n_zeros], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(m_zeros, n_zeros): @@ -621,7 +621,7 @@ def main(x: R.Tensor([n], "float32")): arange_n = T.int64() @Ts.prim_func(private=True) - def arange(n: arange_n, T_arange: T.Buffer((arange_n // T.int64(2),), 'int64')): + def arange(n: arange_n, T_arange: T.Tensor((arange_n // T.int64(2),), 'int64')): T.func_attr({"tirx.noalias": True}) for ax0 in range(n // T.int64(2)): @@ -653,7 +653,7 @@ def main(x: R.Tensor((2, 3, 4), "float32")) -> R.Tensor((3,), "int64"): return gv_1 @Ts.prim_func(private=True) - def shape_to_tensor(shape_to_tensor: T.Buffer((T.int64(3),), "int64")): + def shape_to_tensor(shape_to_tensor: T.Tensor((T.int64(3),), "int64")): T.func_attr({"tirx.noalias": True}) for i in range(T.int64(3)): with Ts.sblock("shape_to_tensor"): @@ -690,7 +690,7 @@ def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((2,), "int64"): return gv_1 @Ts.prim_func(private=True) - def shape_to_tensor(m: T.int64, n: T.int64, shape_to_tensor: T.Buffer((T.int64(2),), "int64")): + def shape_to_tensor(m: T.int64, n: T.int64, shape_to_tensor: T.Tensor((T.int64(2),), "int64")): T.func_attr({"tirx.noalias": True}) for i in range(T.int64(2)): with Ts.sblock("shape_to_tensor"): @@ -725,7 +725,7 @@ def main(x: R.Tensor((m, 3), "float32")) -> R.Tensor((2,), "int64"): return gv_1 @Ts.prim_func(private=True) - def shape_to_tensor(m: T.int64, shape_to_tensor: T.Buffer((T.int64(2),), "int64")): + def shape_to_tensor(m: T.int64, shape_to_tensor: T.Tensor((T.int64(2),), "int64")): T.func_attr({"tirx.noalias": True}) for i in range(T.int64(2)): with Ts.sblock("shape_to_tensor"): @@ -775,7 +775,7 @@ def main(x: R.Tensor((2, 3, 4), "float32")) -> R.Tensor((2, 3, 4), "float32"): return gv @Ts.prim_func(private=True) - def tril(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32"), trilu: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32")): + def tril(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32"), trilu: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(T.int64(2), T.int64(3), T.int64(4)): with Ts.sblock("trilu"): @@ -817,7 +817,7 @@ def main(x: R.Tensor((m_main, n_main, k_main), "int8")) -> R.Tensor((m_main, n_m return gv @Ts.prim_func(private=True) - def tril(rxplaceholder: T.Buffer([m_tril, n_tril, k_tril], dtype='int8'), trilu: T.Buffer([m_tril, n_tril, k_tril], dtype='int8')): + def tril(rxplaceholder: T.Tensor([m_tril, n_tril, k_tril], dtype='int8'), trilu: T.Tensor([m_tril, n_tril, k_tril], dtype='int8')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(m_tril, n_tril, k_tril): @@ -849,7 +849,7 @@ def main(x: R.Tensor((2, 3, 4), "float32")) -> R.Tensor((2, 3, 4), "float32"): return gv @Ts.prim_func(private=True) - def triu(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32"), trilu: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32")): + def triu(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32"), trilu: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(T.int64(2), T.int64(3), T.int64(4)): with Ts.sblock("trilu"): @@ -891,7 +891,7 @@ def main(x: R.Tensor((m_main, n_main, k_main), "int8")) -> R.Tensor((m_main, n_m return gv @Ts.prim_func(private=True) - def triu(rxplaceholder: T.Buffer([m_triu, n_triu, k_triu], dtype='int8'), trilu: T.Buffer([m_triu, n_triu, k_triu], dtype='int8')): + def triu(rxplaceholder: T.Tensor([m_triu, n_triu, k_triu], dtype='int8'), trilu: T.Tensor([m_triu, n_triu, k_triu], dtype='int8')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(m_triu, n_triu, k_triu): @@ -926,7 +926,7 @@ def main(x: R.Tensor((2, 3, 4), "float32")) -> R.Tensor((2, 3, 4), "int32"): return gv @Ts.prim_func(private=True) - def cast(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32"), compute: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "int32")): + def cast(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32"), compute: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "int32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(T.int64(2), T.int64(3), T.int64(4)): with Ts.sblock("compute"): @@ -986,7 +986,7 @@ def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), return gv @Ts.prim_func(private=True) - def cast(rxplaceholder: T.Buffer([m_cast, n_cast], dtype='float32'), compute: T.Buffer([m_cast, n_cast], dtype='int32')): + def cast(rxplaceholder: T.Tensor([m_cast, n_cast], dtype='float32'), compute: T.Tensor([m_cast, n_cast], dtype='int32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(m_cast, n_cast): diff --git a/tests/python/relax/test_transform_legalize_ops_distributed.py b/tests/python/relax/test_transform_legalize_ops_distributed.py index 6727bb6ebd80..885f1b58ad66 100644 --- a/tests/python/relax/test_transform_legalize_ops_distributed.py +++ b/tests/python/relax/test_transform_legalize_ops_distributed.py @@ -40,7 +40,7 @@ def main(x: R.Tensor((10, 10), "float32")) -> R.Tensor((10, 5), "float32"): @I.ir_module class Expected: @Ts.prim_func(private=True) - def strided_slice(A: T.Buffer((T.int64(10), T.int64(10)), "float32"), worker_id: T.int64, redistribute_replica_to_shard: T.Buffer((T.int64(10), T.int64(5)), "float32")): + def strided_slice(A: T.Tensor((T.int64(10), T.int64(10)), "float32"), worker_id: T.int64, redistribute_replica_to_shard: T.Tensor((T.int64(10), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1 in T.grid(T.int64(10), T.int64(5)): diff --git a/tests/python/relax/test_transform_legalize_ops_grad.py b/tests/python/relax/test_transform_legalize_ops_grad.py index c61c29220733..8b6db70db806 100644 --- a/tests/python/relax/test_transform_legalize_ops_grad.py +++ b/tests/python/relax/test_transform_legalize_ops_grad.py @@ -37,7 +37,7 @@ def main(output_grad: R.Tensor((), "float32"), predictions: R.Tensor((2, 3, 4, 5 @I.ir_module class Expected: @Ts.prim_func(private=True) - def nll_loss_backward(rxplaceholder: T.Buffer((), "float32"), rxplaceholder_1: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_2: T.Buffer((T.int64(2), T.int64(4), T.int64(5)), "int64"), rxplaceholder_3: T.Buffer((T.int64(4),), "float32"), pred_grad: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32")): + def nll_loss_backward(rxplaceholder: T.Tensor((), "float32"), rxplaceholder_1: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_2: T.Tensor((T.int64(2), T.int64(4), T.int64(5)), "int64"), rxplaceholder_3: T.Tensor((T.int64(4),), "float32"), pred_grad: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): all_weights = Ts.sblock_alloc_buffer((T.int64(2), T.int64(4), T.int64(5))) @@ -100,7 +100,7 @@ def main(output_grad: R.Tensor((), "float32"), predictions: R.Tensor((2, 3, 4, 5 @I.ir_module class Expected: @Ts.prim_func(private=True) - def te_nll_loss_backward_no_weight(rxplaceholder: T.Buffer((), "float32"), rxplaceholder_1: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_2: T.Buffer((T.int64(2), T.int64(4), T.int64(5)), "int64"), pred_grad: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32")): + def te_nll_loss_backward_no_weight(rxplaceholder: T.Tensor((), "float32"), rxplaceholder_1: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_2: T.Tensor((T.int64(2), T.int64(4), T.int64(5)), "int64"), pred_grad: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): T_full = Ts.sblock_alloc_buffer((T.int64(3),)) @@ -176,7 +176,7 @@ def main(output_grad: R.Tensor((), dtype="float32"), predictions: R.Tensor((4,), return gv @Ts.prim_func(private=True) - def nll_loss_backward(rxplaceholder: T.Buffer((), "float32"), rxplaceholder_1: T.Buffer((T.int64(4),), "float32"), rxplaceholder_2: T.Buffer((), "int64"), rxplaceholder_3: T.Buffer((T.int64(4),), "float32"), pred_grad: T.Buffer((T.int64(4),), "float32")): + def nll_loss_backward(rxplaceholder: T.Tensor((), "float32"), rxplaceholder_1: T.Tensor((T.int64(4),), "float32"), rxplaceholder_2: T.Tensor((), "int64"), rxplaceholder_3: T.Tensor((T.int64(4),), "float32"), pred_grad: T.Tensor((T.int64(4),), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): all_weights = Ts.sblock_alloc_buffer(()) @@ -221,7 +221,7 @@ def main(output_grad: R.Tensor((3, 2, 6, 5), "float32"), data: R.Tensor((3, 2, 1 @I.ir_module class Expected: @Ts.prim_func(private=True) - def max_pool2d_backward(A: T.Buffer((T.int64(3), T.int64(2), T.int64(6), T.int64(5)), "float32"), B: T.Buffer((T.int64(3), T.int64(2), T.int64(10), T.int64(10)), "float32"), T_pool_grad: T.Buffer((T.int64(3), T.int64(2), T.int64(10), T.int64(10)), "float32")): + def max_pool2d_backward(A: T.Tensor((T.int64(3), T.int64(2), T.int64(6), T.int64(5)), "float32"), B: T.Tensor((T.int64(3), T.int64(2), T.int64(10), T.int64(10)), "float32"), T_pool_grad: T.Tensor((T.int64(3), T.int64(2), T.int64(10), T.int64(10)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): pad_temp = Ts.sblock_alloc_buffer((T.int64(3), T.int64(2), T.int64(15), T.int64(13))) @@ -277,7 +277,7 @@ def main(output_grad: R.Tensor((3, 2, 6, 5), "float32"), data: R.Tensor((3, 2, 1 @I.ir_module class Expected: @Ts.prim_func(private=True) - def avg_pool2d_backward(output_grad: T.Buffer((T.int64(3), T.int64(2), T.int64(6), T.int64(5)), "float32"), data: T.Buffer((T.int64(3), T.int64(2), T.int64(10), T.int64(10)), "float32"), T_pool_grad: T.Buffer((T.int64(3), T.int64(2), T.int64(10), T.int64(10)), "float32")): + def avg_pool2d_backward(output_grad: T.Tensor((T.int64(3), T.int64(2), T.int64(6), T.int64(5)), "float32"), data: T.Tensor((T.int64(3), T.int64(2), T.int64(10), T.int64(10)), "float32"), T_pool_grad: T.Tensor((T.int64(3), T.int64(2), T.int64(10), T.int64(10)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1, ax2, ax3, wh, ww in T.grid(T.int64(3), T.int64(2), T.int64(10), T.int64(10), T.int64(3), T.int64(3)): @@ -312,7 +312,7 @@ def main(output_grad: R.Tensor((3, 2, 5), "float32"), x: R.Tensor((3, 4, 5), "fl @I.ir_module class Expected: @Ts.prim_func(private=True) - def take_backward(rxplaceholder: T.Buffer((T.int64(3), T.int64(2), T.int64(5)), offset_factor=1), rxplaceholder_1: T.Buffer((T.int64(3), T.int64(4), T.int64(5)), offset_factor=1), rxplaceholder_2: T.Buffer((T.int64(2),), 'int32', offset_factor=1), out_buf: T.Buffer((T.int64(3), T.int64(4), T.int64(5)), "float32")): + def take_backward(rxplaceholder: T.Tensor((T.int64(3), T.int64(2), T.int64(5)), offset_factor=1), rxplaceholder_1: T.Tensor((T.int64(3), T.int64(4), T.int64(5)), offset_factor=1), rxplaceholder_2: T.Tensor((T.int64(2),), 'int32', offset_factor=1), out_buf: T.Tensor((T.int64(3), T.int64(4), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) with Ts.sblock("take_backward"): @@ -356,7 +356,7 @@ def main(output_grad: R.Tensor((m, i), "float32"), x: R.Tensor((m, n), "float32" @I.ir_module class Expected: @Ts.prim_func(private=True) - def take_backward(rxplaceholder: T.Buffer((m_take_backward, i_take_backward), offset_factor=1), rxplaceholder_1: T.Buffer((m_take_backward, n_take_backward), offset_factor=1), rxplaceholder_2: T.Buffer((i_take_backward,), 'int32', offset_factor=1), out_buf: T.Buffer((m_take_backward, n_take_backward))): + def take_backward(rxplaceholder: T.Tensor((m_take_backward, i_take_backward), offset_factor=1), rxplaceholder_1: T.Tensor((m_take_backward, n_take_backward), offset_factor=1), rxplaceholder_2: T.Tensor((i_take_backward,), 'int32', offset_factor=1), out_buf: T.Tensor((m_take_backward, n_take_backward))): T.func_attr({"tirx.noalias": True}) with Ts.sblock("take_backward"): diff --git a/tests/python/relax/test_transform_legalize_ops_image.py b/tests/python/relax/test_transform_legalize_ops_image.py index 142f2dfbdead..35dbc8e6cdca 100644 --- a/tests/python/relax/test_transform_legalize_ops_image.py +++ b/tests/python/relax/test_transform_legalize_ops_image.py @@ -42,7 +42,7 @@ def main(x: R.Tensor((2, 8, 8, 3), "float32")) -> R.Tensor((2, 16, 16, 3), "floa return gv @Ts.prim_func(private=True) - def resize2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(8), T.int64(8), T.int64(3)), "float32"), resize: T.Buffer((T.int64(2), T.int64(16), T.int64(16), T.int64(3)), "float32")): + def resize2d(rxplaceholder: T.Tensor((T.int64(2), T.int64(8), T.int64(8), T.int64(3)), "float32"), resize: T.Tensor((T.int64(2), T.int64(16), T.int64(16), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(16), T.int64(16), T.int64(3)): with Ts.sblock("resize"): @@ -93,7 +93,7 @@ def main(dumb_param: R.Tensor((oh_main, ow_main)), x: R.Tensor((n_main, c_main, return gv @Ts.prim_func(private=True) - def resize2d(rxplaceholder: T.Buffer([n_resize2d, c_resize2d, h_resize2d, w_resize2d, T.int64(16)], dtype='float32'), resize: T.Buffer([n_resize2d, c_resize2d, oh_resize2d, ow_resize2d, T.int64(16)], dtype='float32')): + def resize2d(rxplaceholder: T.Tensor([n_resize2d, c_resize2d, h_resize2d, w_resize2d, T.int64(16)], dtype='float32'), resize: T.Tensor([n_resize2d, c_resize2d, oh_resize2d, ow_resize2d, T.int64(16)], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3, i4 in T.grid(n_resize2d, c_resize2d, oh_resize2d, ow_resize2d, T.int64(16)): @@ -125,7 +125,7 @@ def main(theta: R.Tensor((2, 2, 3), "float32")) -> R.Tensor((2, 2, 16, 16), "flo return gv @Ts.prim_func(private=True) - def affine_grid(theta: T.Buffer((T.int64(2), T.int64(2), T.int64(3))), compute: T.Buffer((T.int64(2), T.int64(2), T.int64(16), T.int64(16)))): + def affine_grid(theta: T.Tensor((T.int64(2), T.int64(2), T.int64(3))), compute: T.Tensor((T.int64(2), T.int64(2), T.int64(16), T.int64(16)))): T.func_attr({"tirx.noalias": True}) with Ts.sblock("root"): 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 04ec65f0075f..593ea050eff5 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 @@ -45,7 +45,7 @@ def main(x: R.Tensor((2, 3, 4), "float32"), indices: R.Tensor((4,), "int64")) -> return gv @Ts.prim_func(private=True) - def take(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32"), rxplaceholder_1: T.Buffer(T.int64(4), "int64"), T_take: T.Buffer((T.int64(2), T.int64(4), T.int64(4)), "float32")): + def take(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32"), rxplaceholder_1: T.Tensor(T.int64(4), "int64"), T_take: T.Tensor((T.int64(2), T.int64(4), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(T.int64(2), T.int64(4), T.int64(4)): with Ts.sblock("T_take"): @@ -76,7 +76,7 @@ def main(x: R.Tensor((2, 3, 4), "float32"), index: T.int64) -> R.Tensor((2, 4), return gv @Ts.prim_func(private=True) - def take(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32"), index: T.int64, T_take: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def take(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32"), index: T.int64, T_take: T.Tensor((T.int64(2), 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"): @@ -107,7 +107,7 @@ def main(x: R.Tensor((2, 3, 4), "float32")) -> R.Tensor((2, 4), "float32"): return gv @Ts.prim_func(private=True) - def take(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32"), T_take: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def take(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32"), T_take: T.Tensor((T.int64(2), 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"): @@ -149,7 +149,7 @@ def main(x: R.Tensor((m_main, n_main), "float32"), indices: R.Tensor((i_main,), return gv @Ts.prim_func(private=True) - def take(rxplaceholder: T.Buffer([m_take, n_take], dtype='float32'), rxplaceholder_1: T.Buffer([i_take], dtype='int64'), T_take: T.Buffer([m_take, i_take], dtype='float32')): + def take(rxplaceholder: T.Tensor([m_take, n_take], dtype='float32'), rxplaceholder_1: T.Tensor([i_take], dtype='int64'), T_take: T.Tensor([m_take, i_take], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(m_take, i_take): @@ -186,7 +186,7 @@ def main(x: R.Tensor((2, n_main, 4), "float32")) -> R.Tensor((2, 4), "float32"): return gv @Ts.prim_func(private=True) - def take(rxplaceholder: T.Buffer((T.int64(2), n_take, T.int64(4)), 'float32'), T_take: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def take(rxplaceholder: T.Tensor((T.int64(2), n_take, T.int64(4)), 'float32'), T_take: T.Tensor((T.int64(2), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i2 in T.grid(T.int64(2), T.int64(4)): @@ -218,7 +218,7 @@ def main(x: R.Tensor((8, 9, 10, 10), dtype="float32")) -> R.Tensor((4, 9, 10, 3) return gv @Ts.prim_func(private=True) - def strided_slice(rxplaceholder: T.Buffer((T.int64(8), T.int64(9), T.int64(10), T.int64(10)), "float32"), T_strided_slice_with_axes: T.Buffer((T.int64(4), T.int64(9), T.int64(10), T.int64(3)), "float32")): + def strided_slice(rxplaceholder: T.Tensor((T.int64(8), T.int64(9), T.int64(10), T.int64(10)), "float32"), T_strided_slice_with_axes: T.Tensor((T.int64(4), T.int64(9), T.int64(10), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(9), T.int64(10), T.int64(3)): with Ts.sblock("T_strided_slice_with_axes"): @@ -249,7 +249,7 @@ def main(x: R.Tensor((8, 9, 10, 10), dtype="float32")): return gv @Ts.prim_func(private=True) - def strided_slice(rxplaceholder: T.Buffer((T.int64(8), T.int64(9), T.int64(10), T.int64(10)), "float32"), T_strided_slice_with_axes: T.Buffer((T.int64(7), T.int64(9), T.int64(10), T.int64(2)), "float32")): + def strided_slice(rxplaceholder: T.Tensor((T.int64(8), T.int64(9), T.int64(10), T.int64(10)), "float32"), T_strided_slice_with_axes: T.Tensor((T.int64(7), T.int64(9), T.int64(10), T.int64(2)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1, ax2, ax3 in T.grid(T.int64(7), T.int64(9), T.int64(10), T.int64(2)): @@ -281,7 +281,7 @@ def main(x: R.Tensor((8, 9, 10), dtype="float32")) -> R.Tensor((8, 9, 3), dtype= return gv @Ts.prim_func(private=True) - def strided_slice(rxplaceholder: T.Buffer((T.int64(8), T.int64(9), T.int64(10)), "float32"), T_strided_slice_with_axes: T.Buffer((T.int64(8), T.int64(9), T.int64(3)), "float32")): + def strided_slice(rxplaceholder: T.Tensor((T.int64(8), T.int64(9), T.int64(10)), "float32"), T_strided_slice_with_axes: T.Tensor((T.int64(8), T.int64(9), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for ax0, ax1, ax2 in T.grid(T.int64(8), T.int64(9), T.int64(3)): with Ts.sblock("T_strided_slice_with_axes"): @@ -315,7 +315,7 @@ def main(x: R.Tensor((m, n), "float32")) -> R.Tensor((2, n), "float32"): @I.ir_module class Expected: @Ts.prim_func(private=True) - def strided_slice(A: T.Buffer((m_strided_slice, n_strided_slice)), T_dynamic_strided_slice_with_axes: T.Buffer((T.int64(3), n_strided_slice))): + def strided_slice(A: T.Tensor((m_strided_slice, n_strided_slice)), T_dynamic_strided_slice_with_axes: T.Tensor((T.int64(3), n_strided_slice))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -359,7 +359,7 @@ def main(x: R.Tensor((10, n_main), dtype="float32")) -> R.Tensor((3, n_main), dt return gv @Ts.prim_func(private=True) - def strided_slice(rxplaceholder: T.Buffer([T.int64(10), n_strided_slice], dtype='float32'), T_strided_slice_with_axes: T.Buffer([T.int64(3), n_strided_slice], dtype='float32')): + def strided_slice(rxplaceholder: T.Tensor([T.int64(10), n_strided_slice], dtype='float32'), T_strided_slice_with_axes: T.Tensor([T.int64(3), n_strided_slice], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(3), n_strided_slice): @@ -396,7 +396,7 @@ def main(x: R.Tensor((10, n_main), dtype="float32")) -> R.Tensor((3, n_main), dt return gv @Ts.prim_func(private=True) - def strided_slice(rxplaceholder: T.Buffer([T.int64(10), n_strided_slice], dtype='float32'), T_strided_slice_with_axes: T.Buffer([T.int64(3), n_strided_slice], dtype='float32')): + def strided_slice(rxplaceholder: T.Tensor([T.int64(10), n_strided_slice], dtype='float32'), T_strided_slice_with_axes: T.Tensor([T.int64(3), n_strided_slice], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(3), n_strided_slice): @@ -429,7 +429,7 @@ def main(x: R.Tensor((10, n_main), dtype="float32")) -> R.Tensor((3, n_main), dt return gv @Ts.prim_func(private=True) - def strided_slice(rxplaceholder: T.Buffer([T.int64(10), n_strided_slice], dtype='float32'), T_strided_slice_with_axes: T.Buffer([T.int64(3), n_strided_slice], dtype='float32')): + def strided_slice(rxplaceholder: T.Tensor([T.int64(10), n_strided_slice], dtype='float32'), T_strided_slice_with_axes: T.Tensor([T.int64(3), n_strided_slice], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(3), n_strided_slice): @@ -461,13 +461,13 @@ def main(x: R.Tensor((8, 9, 10, 10), "float32"), begin: R.Tensor((4,),"int64"), class Expected: @Ts.prim_func(private=True) def dynamic_strided_slice( - rxplaceholder: T.Buffer( + rxplaceholder: T.Tensor( (T.int64(8), T.int64(9), T.int64(10), T.int64(10)), "float32" ), - rxplaceholder_1: T.Buffer((T.int64(4),), "int64"), - rxplaceholder_2: T.Buffer((T.int64(4),), "int64"), - rxplaceholder_3: T.Buffer((T.int64(4),), "int64"), - T_strided_slice_dynamic: T.Buffer((s_dynamic_strided_slice, s_1_dynamic_strided_slice, s_2_dynamic_strided_slice, s_3_dynamic_strided_slice)), + rxplaceholder_1: T.Tensor((T.int64(4),), "int64"), + rxplaceholder_2: T.Tensor((T.int64(4),), "int64"), + rxplaceholder_3: T.Tensor((T.int64(4),), "int64"), + T_strided_slice_dynamic: T.Tensor((s_dynamic_strided_slice, s_1_dynamic_strided_slice, s_2_dynamic_strided_slice, s_3_dynamic_strided_slice)), ): T.func_attr({"tirx.noalias": True}) @@ -555,13 +555,13 @@ def dynamic_strided_slice( @Ts.prim_func(private=True) def shape_func( - rxplaceholder: T.Buffer( + rxplaceholder: T.Tensor( (T.int64(8), T.int64(9), T.int64(10), T.int64(10)), "float32" ), - rxplaceholder_1: T.Buffer((T.int64(4),), "int64"), - rxplaceholder_2: T.Buffer((T.int64(4),), "int64"), - rxplaceholder_3: T.Buffer((T.int64(4),), "int64"), - T_shape_func_strided_slice_dynamic: T.Buffer((T.int64(4),), "int64"), + rxplaceholder_1: T.Tensor((T.int64(4),), "int64"), + rxplaceholder_2: T.Tensor((T.int64(4),), "int64"), + rxplaceholder_3: T.Tensor((T.int64(4),), "int64"), + T_shape_func_strided_slice_dynamic: T.Tensor((T.int64(4),), "int64"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -805,11 +805,11 @@ def main(x: R.Tensor((10, n), "float32"), begin:R.Tensor((2,), "int64"), end:R.T class Expected: @Ts.prim_func(private=True) def dynamic_strided_slice( - rxplaceholder_3: T.Buffer((T.int64(10), n_dynamic_strided_slice)), - rxplaceholder: T.Buffer((T.int64(2),), "int64"), - rxplaceholder_1: T.Buffer((T.int64(2),), "int64"), - rxplaceholder_2: T.Buffer((T.int64(2),), "int64"), - T_strided_slice_dynamic: T.Buffer((s_dynamic_strided_slice, s_1_dynamic_strided_slice)), + rxplaceholder_3: T.Tensor((T.int64(10), n_dynamic_strided_slice)), + rxplaceholder: T.Tensor((T.int64(2),), "int64"), + rxplaceholder_1: T.Tensor((T.int64(2),), "int64"), + rxplaceholder_2: T.Tensor((T.int64(2),), "int64"), + T_strided_slice_dynamic: T.Tensor((s_dynamic_strided_slice, s_1_dynamic_strided_slice)), ): T.func_attr({"tirx.noalias": True}) @@ -863,11 +863,11 @@ def dynamic_strided_slice( @Ts.prim_func(private=True) def shape_func( - rxplaceholder_3: T.Buffer((T.int64(10), n_shape_func)), - rxplaceholder: T.Buffer((T.int64(2),), "int64"), - rxplaceholder_1: T.Buffer((T.int64(2),), "int64"), - rxplaceholder_2: T.Buffer((T.int64(2),), "int64"), - T_shape_func_strided_slice_dynamic: T.Buffer((T.int64(2),), "int64"), + rxplaceholder_3: T.Tensor((T.int64(10), n_shape_func)), + rxplaceholder: T.Tensor((T.int64(2),), "int64"), + rxplaceholder_1: T.Tensor((T.int64(2),), "int64"), + rxplaceholder_2: T.Tensor((T.int64(2),), "int64"), + T_shape_func_strided_slice_dynamic: T.Tensor((T.int64(2),), "int64"), ): T.func_attr({"tirx.noalias": True}) @@ -1027,7 +1027,7 @@ def main(x: R.Tensor((4,), "float32"), y: R.Tensor((2, 3, 4, 5), "float32")) -> return gv @Ts.prim_func(private=True) - def matmul(rxplaceholder: T.Buffer(T.int64(4), "float32"), rxplaceholder_1: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), matmul: T.Buffer((T.int64(2), T.int64(3), T.int64(5)), "float32")): + def matmul(rxplaceholder: T.Tensor(T.int64(4), "float32"), rxplaceholder_1: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), matmul: T.Tensor((T.int64(2), T.int64(3), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(3), T.int64(5), T.int64(4)): with Ts.sblock("matmul"): @@ -1060,7 +1060,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32"), y: R.Tensor((5,), "float32")) -> return gv @Ts.prim_func(private=True) - def matmul(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_1: T.Buffer(T.int64(5), "float32"), matmul: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32")): + def matmul(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_1: T.Tensor(T.int64(5), "float32"), matmul: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(3), T.int64(4), T.int64(5)): with Ts.sblock("matmul"): @@ -1093,7 +1093,7 @@ def main(x: R.Tensor((4,), "float32"), y: R.Tensor((4,), "float32")) -> R.Tensor return gv @Ts.prim_func(private=True) - def matmul(rxplaceholder: T.Buffer(T.int64(4), "float32"), rxplaceholder_1: T.Buffer(T.int64(4), "float32"), matmul: T.Buffer((), "float32")): + def matmul(rxplaceholder: T.Tensor(T.int64(4), "float32"), rxplaceholder_1: T.Tensor(T.int64(4), "float32"), matmul: T.Tensor((), "float32")): T.func_attr({"tirx.noalias": True}) for i0 in T.serial(T.int64(4)): with Ts.sblock("matmul"): @@ -1126,7 +1126,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float16"), y: R.Tensor((6, 2, 3, 5, 7), "flo return gv @Ts.prim_func(private=True) - def matmul(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float16"), rxplaceholder_1: T.Buffer((T.int64(6), T.int64(2), T.int64(3), T.int64(5), T.int64(7)), "float16"), matmul: T.Buffer((T.int64(6), T.int64(2), T.int64(3), T.int64(4), T.int64(7)), "float32")): + def matmul(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float16"), rxplaceholder_1: T.Tensor((T.int64(6), T.int64(2), T.int64(3), T.int64(5), T.int64(7)), "float16"), matmul: T.Tensor((T.int64(6), T.int64(2), T.int64(3), T.int64(4), T.int64(7)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3, i4, i5 in T.grid(T.int64(6), T.int64(2), T.int64(3), T.int64(4), T.int64(7), T.int64(5)): with Ts.sblock("matmul"): @@ -1179,7 +1179,7 @@ def main(x: R.Tensor((b_main, 1, m_main, k_main), "float32"), y: R.Tensor((a_mai return gv @Ts.prim_func(private=True) - def matmul(rxplaceholder: T.Buffer([b_matmul, T.int64(1), m_matmul, k_matmul], dtype='float32'), rxplaceholder_1: T.Buffer([a_matmul, T.int64(1), c_matmul, k_matmul, n_matmul], dtype='float32'), matmul: T.Buffer([a_matmul, b_matmul, c_matmul, m_matmul, n_matmul], dtype='float32')): + def matmul(rxplaceholder: T.Tensor([b_matmul, T.int64(1), m_matmul, k_matmul], dtype='float32'), rxplaceholder_1: T.Tensor([a_matmul, T.int64(1), c_matmul, k_matmul, n_matmul], dtype='float32'), matmul: T.Tensor([a_matmul, b_matmul, c_matmul, m_matmul, n_matmul], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3, i4, i5 in T.grid(a_matmul, b_matmul, c_matmul, m_matmul, n_matmul, k_matmul): @@ -1208,7 +1208,7 @@ def main(x: R.Tensor((1, 1, 4, 5), "float32"), y: R.Tensor((1, 1, 5, 7), "float3 @I.ir_module class Expected: @Ts.prim_func(private=True) - 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")): + def matmul(A: T.Tensor((T.int64(1), T.int64(1), T.int64(4), T.int64(5)), "float32"), B: T.Tensor((T.int64(1), T.int64(1), T.int64(5), T.int64(7)), "float32"), matmul_1: T.Tensor((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_index in T.grid(T.int64(1), T.int64(1), T.int64(4), T.int64(7), T.int64(5)): @@ -1268,9 +1268,9 @@ def main( @Ts.prim_func(private=True) def einsum( - rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), - rxplaceholder_1: T.Buffer((T.int64(3), T.int64(4)), "float32"), - T_einsum: T.Buffer((T.int64(2), T.int64(4)), "float32"), + rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), + rxplaceholder_1: T.Tensor((T.int64(3), T.int64(4)), "float32"), + T_einsum: T.Tensor((T.int64(2), T.int64(4)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1, j in T.grid(T.int64(2), T.int64(4), T.int64(3)): @@ -1323,9 +1323,9 @@ def main( @Ts.prim_func(private=True) def einsum( - rxplaceholder: T.Buffer((a_einsum, b_einsum)), - rxplaceholder_1: T.Buffer((b_einsum, c_einsum)), - T_einsum: T.Buffer((a_einsum, c_einsum)), + rxplaceholder: T.Tensor((a_einsum, b_einsum)), + rxplaceholder_1: T.Tensor((b_einsum, c_einsum)), + T_einsum: T.Tensor((a_einsum, c_einsum)), ): T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/relax/test_transform_legalize_ops_manipulate.py b/tests/python/relax/test_transform_legalize_ops_manipulate.py index 85ad10d49047..fc04589c968d 100644 --- a/tests/python/relax/test_transform_legalize_ops_manipulate.py +++ b/tests/python/relax/test_transform_legalize_ops_manipulate.py @@ -46,7 +46,7 @@ def main(x: R.Tensor((2, 1, 3), "float32")) -> R.Tensor((4, 2, 5, 3), "float32") return gv @Ts.prim_func(private=True) - def broadcast_to(rxplaceholder: T.Buffer((T.int64(2), T.int64(1), T.int64(3)), "float32"), T_broadcast_to: T.Buffer((T.int64(4), T.int64(2), T.int64(5), T.int64(3)), "float32")): + def broadcast_to(rxplaceholder: T.Tensor((T.int64(2), T.int64(1), T.int64(3)), "float32"), T_broadcast_to: T.Tensor((T.int64(4), T.int64(2), T.int64(5), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(2), T.int64(5), T.int64(3)): with Ts.sblock("T_broadcast_to"): @@ -91,7 +91,7 @@ def main(dumb_param: R.Tensor((a_main, c_main)), x: R.Tensor((b_main, 1, d_main) return gv @Ts.prim_func(private=True) - def broadcast_to(rxplaceholder: T.Buffer([b_broadcast_to, T.int64(1), d_broadcast_to], dtype='float32'), T_broadcast_to: T.Buffer([a_broadcast_to, b_broadcast_to, c_broadcast_to, d_broadcast_to], dtype='float32')): + def broadcast_to(rxplaceholder: T.Tensor([b_broadcast_to, T.int64(1), d_broadcast_to], dtype='float32'), T_broadcast_to: T.Tensor([a_broadcast_to, b_broadcast_to, c_broadcast_to, d_broadcast_to], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_broadcast_to, b_broadcast_to, c_broadcast_to, d_broadcast_to): @@ -123,7 +123,7 @@ def main(x1: R.Tensor((1, 2, 3), "float32"), x2: R.Tensor((1, 3, 3), "float32"), return gv @Ts.prim_func(private=True) - def concatenate(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(1), T.int64(3), T.int64(3)), "float32"), rxplaceholder_2: T.Buffer((T.int64(1), T.int64(4), T.int64(3)), "float32"), T_concat: T.Buffer((T.int64(1), T.int64(9), T.int64(3)), "float32")): + def concatenate(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(1), T.int64(3), T.int64(3)), "float32"), rxplaceholder_2: T.Tensor((T.int64(1), T.int64(4), T.int64(3)), "float32"), T_concat: T.Tensor((T.int64(1), T.int64(9), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(T.int64(1), T.int64(9), T.int64(3)): with Ts.sblock("T_concat"): @@ -156,7 +156,7 @@ def main(t: R.Tuple(R.Tensor((3, 4), "float32"), R.Tensor((3, 5), "float32"))) - return gv2 @Ts.prim_func(private=True) - def concatenate(rxplaceholder: T.Buffer((T.int64(3), T.int64(4)), "float32"), rxplaceholder_1: T.Buffer((T.int64(3), T.int64(5)), "float32"), T_concat: T.Buffer((T.int64(3), T.int64(9)), "float32")): + def concatenate(rxplaceholder: T.Tensor((T.int64(3), T.int64(4)), "float32"), rxplaceholder_1: T.Tensor((T.int64(3), T.int64(5)), "float32"), T_concat: T.Tensor((T.int64(3), T.int64(9)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(3), T.int64(9)): with Ts.sblock("T_concat"): @@ -204,7 +204,7 @@ def main(t: R.Tuple(R.Tensor((a_main, b0_main), "float32"), R.Tensor((a_main, b1 return gv3 @Ts.prim_func(private=True) - def concatenate(rxplaceholder: T.Buffer([a_concatenate, b0_concatenate], dtype='float32'), rxplaceholder_1: T.Buffer([a_concatenate, b1_concatenate], dtype='float32'), rxplaceholder_2: T.Buffer([a_concatenate, b2_concatenate], dtype='float32'), T_concat: T.Buffer([a_concatenate, b0_concatenate + b1_concatenate + b2_concatenate], dtype='float32')): + def concatenate(rxplaceholder: T.Tensor([a_concatenate, b0_concatenate], dtype='float32'), rxplaceholder_1: T.Tensor([a_concatenate, b1_concatenate], dtype='float32'), rxplaceholder_2: T.Tensor([a_concatenate, b2_concatenate], dtype='float32'), T_concat: T.Tensor([a_concatenate, b0_concatenate + b1_concatenate + b2_concatenate], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(a_concatenate, b0_concatenate + b1_concatenate + b2_concatenate): @@ -236,7 +236,7 @@ def main(x: R.Tensor((2, 3, 4), "float32")) -> R.Tensor((2, 1, 1, 1, 3, 1, 4, 1) return gv @Ts.prim_func(private=True) - def expand_dims(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32"), expand_dims: T.Buffer((T.int64(2), T.int64(1), T.int64(1), T.int64(1), T.int64(3), T.int64(1), T.int64(4), T.int64(1)), "float32")): + def expand_dims(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32"), expand_dims: T.Tensor((T.int64(2), T.int64(1), T.int64(1), T.int64(1), T.int64(3), T.int64(1), T.int64(4), T.int64(1)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3, i4, i5, i6, i7 in T.grid(T.int64(2), T.int64(1), T.int64(1), T.int64(1), T.int64(3), T.int64(1), T.int64(4), T.int64(1)): with Ts.sblock("expand_dims"): @@ -278,7 +278,7 @@ def main(x: R.Tensor((a_main, b_main, c_main), "float32")) -> R.Tensor((a_main, return gv @Ts.prim_func(private=True) - def expand_dims(rxplaceholder: T.Buffer([a_expand_dims, b_expand_dims, c_expand_dims], dtype='float32'), expand_dims: T.Buffer([a_expand_dims, T.int64(1), b_expand_dims, T.int64(1), c_expand_dims, T.int64(1)], dtype='float32')): + def expand_dims(rxplaceholder: T.Tensor([a_expand_dims, b_expand_dims, c_expand_dims], dtype='float32'), expand_dims: T.Tensor([a_expand_dims, T.int64(1), b_expand_dims, T.int64(1), c_expand_dims, T.int64(1)], dtype='float32')): T.func_attr({"tirx.noalias": True}) 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)): @@ -310,7 +310,7 @@ def main(x: R.Tensor((2, 3, 4), "float32")) -> R.Tensor((24,), "float32"): return gv @Ts.prim_func(private=True) - def reshape(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32"), T_reshape: T.Buffer(T.int64(24), "float32")): + def reshape(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32"), T_reshape: T.Tensor(T.int64(24), "float32")): T.func_attr({"tirx.noalias": True}) for i0 in T.serial(T.int64(24)): with Ts.sblock("T_reshape"): @@ -341,7 +341,7 @@ def main(x: R.Tensor((), "float32")) -> R.Tensor((1,), "float32"): return gv @Ts.prim_func(private=True) - def reshape(rxplaceholder: T.Buffer((), "float32"), T_reshape: T.Buffer(T.int64(1), "float32")): + def reshape(rxplaceholder: T.Tensor((), "float32"), T_reshape: T.Tensor(T.int64(1), "float32")): T.func_attr({"tirx.noalias": True}) for i0 in T.serial(T.int64(1)): with Ts.sblock("T_reshape"): @@ -383,7 +383,7 @@ def main(x: R.Tensor((a_main, b_main, c_main), "float32")) -> R.Tensor((a_main * return gv @Ts.prim_func(private=True) - def reshape(rxplaceholder: T.Buffer([a_reshape, b_reshape, c_reshape], dtype='float32'), T_reshape: T.Buffer([a_reshape * b_reshape * c_reshape], dtype='float32')): + def reshape(rxplaceholder: T.Tensor([a_reshape, b_reshape, c_reshape], dtype='float32'), T_reshape: T.Tensor([a_reshape * b_reshape * c_reshape], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0 in T.serial(a_reshape * b_reshape * c_reshape): @@ -415,7 +415,7 @@ def main(x: R.Tensor((1, 2, 3, 4), "float32")) -> R.Tensor((2, 4, 3, 1), "float3 return gv @Ts.prim_func(private=True) - def transpose(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3), T.int64(4)), "float32"), T_transpose: T.Buffer((T.int64(2), T.int64(4), T.int64(3), T.int64(1)), "float32")): + def transpose(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3), T.int64(4)), "float32"), T_transpose: T.Tensor((T.int64(2), T.int64(4), T.int64(3), T.int64(1)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(4), T.int64(3), T.int64(1)): with Ts.sblock("T_transpose"): @@ -460,7 +460,7 @@ def main(x: R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) -> R.Te return gv @Ts.prim_func(private=True) - def transpose(rxplaceholder: T.Buffer([a_transpose, b_transpose, c_transpose, d_transpose], dtype='float32'), T_transpose: T.Buffer([b_transpose, d_transpose, c_transpose, a_transpose], dtype='float32')): + def transpose(rxplaceholder: T.Tensor([a_transpose, b_transpose, c_transpose, d_transpose], dtype='float32'), T_transpose: T.Tensor([b_transpose, d_transpose, c_transpose, a_transpose], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(b_transpose, d_transpose, c_transpose, a_transpose): @@ -492,7 +492,7 @@ def main(x: R.Tensor((1, 2, 3, 4), "float32")) -> R.Tensor((8, 3), "float32"): return gv @Ts.prim_func(private=True) - def reshape(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3), T.int64(4)), "float32"), T_reshape: T.Buffer((T.int64(8), T.int64(3)), "float32")): + def reshape(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3), T.int64(4)), "float32"), T_reshape: T.Tensor((T.int64(8), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(8), T.int64(3)): with Ts.sblock("T_reshape"): @@ -520,8 +520,8 @@ def main(x: R.Tensor((1, 2, 3, 4), "float32")) -> R.Tensor((8, 3), "float32"): class Expected2: @Ts.prim_func(private=True) def reshape( - rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3), T.int64(4)), "float32"), - T_reshape: T.Buffer((T.int64(8), T.int64(3)), "float32"), + rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3), T.int64(4)), "float32"), + T_reshape: T.Tensor((T.int64(8), T.int64(3)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -580,7 +580,7 @@ def main(x: R.Tensor((a_main, b_main), "float32")) -> R.Tensor((a_main // 2, b_m return gv @Ts.prim_func(private=True) - def reshape(rxplaceholder: T.Buffer([a_reshape, b_reshape], dtype='float32'), T_reshape: T.Buffer([a_reshape // T.int64(2), b_reshape * T.int64(2)], dtype='float32')): + def reshape(rxplaceholder: T.Tensor([a_reshape, b_reshape], dtype='float32'), T_reshape: T.Tensor([a_reshape // T.int64(2), b_reshape * T.int64(2)], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(a_reshape // T.int64(2), b_reshape * T.int64(2)): @@ -626,8 +626,8 @@ def main(x: R.Tensor((a_main, b_main), "float32")) -> R.Tensor( @Ts.prim_func(private=True) def reshape( - rxplaceholder: T.Buffer([a_reshape, b_reshape], dtype="float32"), - T_reshape: T.Buffer([a_reshape // T.int64(2), b_reshape * T.int64(2)], dtype="float32"), + rxplaceholder: T.Tensor([a_reshape, b_reshape], dtype="float32"), + T_reshape: T.Tensor([a_reshape // T.int64(2), b_reshape * T.int64(2)], dtype="float32"), ): T.func_attr({"tirx.noalias": True}) @@ -669,8 +669,8 @@ def main(x: R.Tensor((10, b), "float32")) -> R.Tensor((5, b * 2), "float32"): class Expected3: @Ts.prim_func(private=True) def reshape( - rxplaceholder: T.Buffer((T.int64(10), b_reshape)), - T_reshape: T.Buffer((T.int64(5), b_reshape * T.int64(2))), + rxplaceholder: T.Tensor((T.int64(10), b_reshape)), + T_reshape: T.Tensor((T.int64(5), b_reshape * T.int64(2))), ): T.func_attr({"tirx.noalias": True}) @@ -743,8 +743,8 @@ def main( @Ts.prim_func(private=True) def reshape( - rxplaceholder: T.Buffer(T.int64(16), "float32"), - T_reshape: T.Buffer([M_reshape, N_reshape], 'float32'), + rxplaceholder: T.Tensor(T.int64(16), "float32"), + T_reshape: T.Tensor([M_reshape, N_reshape], 'float32'), ): T.func_attr({"tirx.noalias": True}) @@ -776,7 +776,7 @@ def main(x: R.Tensor((2, 10, 4), "float32")) -> R.Tuple([R.Tensor((2, 3, 4), "fl return gv @Ts.prim_func(private=True) - def split(rxplaceholder: T.Buffer((T.int64(2), T.int64(10), T.int64(4)), "float32"), T_split: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32"), T_split_1: T.Buffer((T.int64(2), T.int64(4), T.int64(4)), "float32"), T_split_2: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32")): + def split(rxplaceholder: T.Tensor((T.int64(2), T.int64(10), T.int64(4)), "float32"), T_split: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32"), T_split_1: T.Tensor((T.int64(2), T.int64(4), T.int64(4)), "float32"), T_split_2: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(T.int64(2), T.int64(3), T.int64(4)): with Ts.sblock("T_split"): @@ -819,7 +819,7 @@ def main(x: R.Tensor((2, 10, 4), "float32")) -> R.Tuple([R.Tensor((2, 4, 4), "fl return gv @Ts.prim_func(private=True) - def split(rxplaceholder: T.Buffer((T.int64(2), T.int64(10), T.int64(4)), "float32"), T_split_sections: T.Buffer((T.int64(2), T.int64(4), T.int64(4)), "float32"), T_split_sections_1: T.Buffer((T.int64(2), T.int64(4), T.int64(4)), "float32"), T_split_sections_2: T.Buffer((T.int64(2), T.int64(2), T.int64(4)), "float32")): + def split(rxplaceholder: T.Tensor((T.int64(2), T.int64(10), T.int64(4)), "float32"), T_split_sections: T.Tensor((T.int64(2), T.int64(4), T.int64(4)), "float32"), T_split_sections_1: T.Tensor((T.int64(2), T.int64(4), T.int64(4)), "float32"), T_split_sections_2: T.Tensor((T.int64(2), T.int64(2), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(T.int64(2), T.int64(4), T.int64(4)): with Ts.sblock("T_split_sections"): @@ -863,7 +863,7 @@ def main(x: R.Tensor((2, 10, 4), "float32")) -> R.Tuple([R.Tensor((2, 5, 4), "fl return gv @Ts.prim_func(private=True) - def split(rxplaceholder: T.Buffer((T.int64(2), T.int64(10), T.int64(4)), "float32"), T_split_sections: T.Buffer((T.int64(2), T.int64(5), T.int64(4)), "float32"), T_split_sections_1: T.Buffer((T.int64(2), T.int64(5), T.int64(4)), "float32")): + def split(rxplaceholder: T.Tensor((T.int64(2), T.int64(10), T.int64(4)), "float32"), T_split_sections: T.Tensor((T.int64(2), T.int64(5), T.int64(4)), "float32"), T_split_sections_1: T.Tensor((T.int64(2), T.int64(5), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(T.int64(2), T.int64(5), T.int64(4)): with Ts.sblock("T_split_sections"): @@ -909,7 +909,7 @@ def main(dumb_param: R.Tensor((n,)), x: R.Tensor((m_main, n * 3), "float32")) -> split_n = T.int64() @Ts.prim_func(private=True) - def split(rxplaceholder: T.Buffer([m_split, split_n * T.int64(3)], dtype='float32'), n: split_n, T_split_sections: T.Buffer([m_split, (split_n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3)], dtype='float32'), T_split_sections_1: T.Buffer([m_split, (split_n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3) * T.int64(2) - (split_n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3)], dtype='float32'), T_split_sections_2: T.Buffer([m_split, split_n * T.int64(3) - (split_n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3) * T.int64(2)], dtype='float32')): + def split(rxplaceholder: T.Tensor([m_split, split_n * T.int64(3)], dtype='float32'), n: split_n, T_split_sections: T.Tensor([m_split, (split_n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3)], dtype='float32'), T_split_sections_1: T.Tensor([m_split, (split_n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3) * T.int64(2) - (split_n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3)], dtype='float32'), T_split_sections_2: T.Tensor([m_split, split_n * T.int64(3) - (split_n * T.int64(3) + T.int64(3) - T.int64(1)) // T.int64(3) * T.int64(2)], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(m_split, n): @@ -953,7 +953,7 @@ def main(x: R.Tensor((2, 1, 3, 1, 1, 4), "float32")) -> R.Tensor((2, 3, 1, 4), " return gv @Ts.prim_func(private=True) - def squeeze(rxplaceholder: T.Buffer((T.int64(2), T.int64(1), T.int64(3), T.int64(1), T.int64(1), T.int64(4)), "float32"), T_squeeze: T.Buffer((T.int64(2), T.int64(3), T.int64(1), T.int64(4)), "float32")): + def squeeze(rxplaceholder: T.Tensor((T.int64(2), T.int64(1), T.int64(3), T.int64(1), T.int64(1), T.int64(4)), "float32"), T_squeeze: T.Tensor((T.int64(2), T.int64(3), T.int64(1), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(3), T.int64(1), T.int64(4)): with Ts.sblock("T_squeeze"): @@ -984,7 +984,7 @@ def main(x: R.Tensor((2, 1, 3, 1, 1, 4), "float32")) : return gv @Ts.prim_func(private=True) - def squeeze(rxplaceholder: T.Buffer((T.int64(2), T.int64(1), T.int64(3), T.int64(1), T.int64(1), T.int64(4)), "float32"), T_squeeze: T.Buffer((T.int64(2), T.int64(3), T.int64(4)), "float32")): + def squeeze(rxplaceholder: T.Tensor((T.int64(2), T.int64(1), T.int64(3), T.int64(1), T.int64(1), T.int64(4)), "float32"), T_squeeze: T.Tensor((T.int64(2), T.int64(3), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(T.int64(2), T.int64(3), T.int64(4)): with Ts.sblock("T_squeeze"): @@ -1023,7 +1023,7 @@ def main(x: R.Tensor((a_main, 1, b_main, 1), "float32")) -> R.Tensor((a_main, b_ return gv @Ts.prim_func(private=True) - def squeeze(rxplaceholder: T.Buffer([a_squeeze, T.int64(1), b_squeeze, T.int64(1)], dtype='float32'), T_squeeze: T.Buffer([a_squeeze, b_squeeze, T.int64(1)], dtype='float32')): + def squeeze(rxplaceholder: T.Tensor([a_squeeze, T.int64(1), b_squeeze, T.int64(1)], dtype='float32'), T_squeeze: T.Tensor([a_squeeze, b_squeeze, T.int64(1)], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(a_squeeze, b_squeeze, T.int64(1)): @@ -1055,7 +1055,7 @@ def main(x: R.Tensor((2, 3), "float32"), y: R.Tensor((1, 3), "float32")) -> R.Te return gv @Ts.prim_func(private=True) - def collapse_sum(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), rxplaceholder_red: T.Buffer((T.int64(1), T.int64(3)), "float32")): + def collapse_sum(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), rxplaceholder_red: T.Tensor((T.int64(1), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(T.int64(1), T.int64(3), T.int64(2)): with Ts.sblock("rxplaceholder_red"): @@ -1091,7 +1091,7 @@ def main( return gv @Ts.prim_func(private=True) - def collapse_sum(rxplaceholder: T.Buffer((T.int64(3), T.int64(2), T.int64(3)), "float32"), rxplaceholder_red: T.Buffer((T.int64(2), T.int64(1)), "float32")): + def collapse_sum(rxplaceholder: T.Tensor((T.int64(3), T.int64(2), T.int64(3)), "float32"), rxplaceholder_red: T.Tensor((T.int64(2), T.int64(1)), "float32")): T.func_attr({"tirx.noalias": True}) for ax0, ax1, k0, k2 in T.grid(T.int64(2), T.int64(1), T.int64(3), T.int64(3)): with Ts.sblock("rxplaceholder_red"): @@ -1124,7 +1124,7 @@ def main(x: R.Tensor((3, 2, 3), dtype="float32")) -> R.Tensor((6, 2, 3), dtype=" return gv @Ts.prim_func(private=True) - def repeat(rxplaceholder: T.Buffer((T.int64(3), T.int64(2), T.int64(3)), "float32"), T_repeat: T.Buffer((T.int64(6), T.int64(2), T.int64(3)), "float32")): + def repeat(rxplaceholder: T.Tensor((T.int64(3), T.int64(2), T.int64(3)), "float32"), T_repeat: T.Tensor((T.int64(6), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1, ax2 in T.grid(T.int64(6), T.int64(2), T.int64(3)): @@ -1159,8 +1159,8 @@ def main( @Ts.prim_func(private=True) def repeat( - rxplaceholder: T.Buffer((T.int64(3), T.int64(2), T.int64(3)), "float32"), - T_repeat: T.Buffer((T.int64(36),), "float32"), + rxplaceholder: T.Tensor((T.int64(3), T.int64(2), T.int64(3)), "float32"), + T_repeat: T.Tensor((T.int64(36),), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1216,7 +1216,7 @@ def main(x: R.Tensor((a, b, c), "float32")): @I.ir_module class Expected: @Ts.prim_func(private=True) - def repeat(rxplaceholder: T.Buffer((a_repeat, b_repeat, c_repeat)), T_repeat: T.Buffer((T.int64(2) * a_repeat, b_repeat, c_repeat))): + def repeat(rxplaceholder: T.Tensor((a_repeat, b_repeat, c_repeat)), T_repeat: T.Tensor((T.int64(2) * a_repeat, b_repeat, c_repeat))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1249,7 +1249,7 @@ def main(x: R.Tensor((3, 2, 3), "float32")): @I.ir_module class Expected: @Ts.prim_func(private=True) - def tile(rxplaceholder: T.Buffer((T.int64(3), T.int64(2), T.int64(3)), "float32"), T_tile: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(9)), "float32")): + def tile(rxplaceholder: T.Tensor((T.int64(3), T.int64(2), T.int64(3)), "float32"), T_tile: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(9)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for ax0, ax1, ax2, ax3 in T.grid(T.int64(2), T.int64(3), T.int64(4), T.int64(9)): @@ -1292,7 +1292,7 @@ def main(x: R.Tensor((a, b, c), "float32")): @I.ir_module class Expected: @Ts.prim_func(private=True) - def tile(rxplaceholder: T.Buffer((a_tile, b_tile, c_tile)), T_tile: T.Buffer((T.int64(2), a_tile, b_tile * T.int64(2), c_tile * T.int64(3)))): + def tile(rxplaceholder: T.Tensor((a_tile, b_tile, c_tile)), T_tile: T.Tensor((T.int64(2), a_tile, b_tile * T.int64(2), c_tile * T.int64(3)))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1331,8 +1331,8 @@ def main(x: R.Tensor((2, 3), dtype="float32")) -> R.Tensor((2, 3), dtype="float3 @Ts.prim_func(private=True) def flip( - rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), - T_reverse_sequence: T.Buffer((T.int64(2), T.int64(3)), "float32"), + rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), + T_reverse_sequence: T.Tensor((T.int64(2), T.int64(3)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(2), T.int64(3)): @@ -1378,7 +1378,7 @@ def main( return gv @Ts.prim_func(private=True) - def flip(rxplaceholder: T.Buffer((a_flip, b_flip)), T_reverse_sequence: T.Buffer((a_flip, b_flip))): + def flip(rxplaceholder: T.Tensor((a_flip, b_flip)), T_reverse_sequence: T.Tensor((a_flip, b_flip))): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(a_flip, b_flip): @@ -1422,9 +1422,9 @@ def main( @Ts.prim_func(private=True) def reverse_sequence( - rxplaceholder: T.Buffer((T.int64(4), T.int64(2), T.int64(3)), "float32"), - seq_lengths: T.Buffer((T.int64(2),), "int64"), - T_reverse_sequence: T.Buffer((T.int64(4), T.int64(2), T.int64(3)), "float32"), + rxplaceholder: T.Tensor((T.int64(4), T.int64(2), T.int64(3)), "float32"), + seq_lengths: T.Tensor((T.int64(2),), "int64"), + T_reverse_sequence: T.Tensor((T.int64(4), T.int64(2), T.int64(3)), "float32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1, ax2 in T.grid(T.int64(4), T.int64(2), T.int64(3)): @@ -1464,10 +1464,10 @@ def main(x: R.Tensor((4,4), "float32"), indices: R.Tensor((2,2), "int64"), updat class Expected: @Ts.prim_func(private=True) def scatter_elements( - rxplaceholder: T.Buffer((T.int64(4), T.int64(4)), offset_factor=1), - rxplaceholder_1: T.Buffer((T.int64(2), T.int64(2)), 'int64', offset_factor=1), - rxplaceholder_2: T.Buffer((T.int64(2), T.int64(2)), offset_factor=1), - out_buf: T.Buffer((T.int64(4), T.int64(4)), "float32"), + rxplaceholder: T.Tensor((T.int64(4), T.int64(4)), offset_factor=1), + rxplaceholder_1: T.Tensor((T.int64(2), T.int64(2)), 'int64', offset_factor=1), + rxplaceholder_2: T.Tensor((T.int64(2), T.int64(2)), offset_factor=1), + out_buf: T.Tensor((T.int64(4), T.int64(4)), "float32"), ): T.func_attr({"tirx.noalias": True}) @@ -1567,10 +1567,10 @@ def main(x: R.Tensor((a, b), "float32"), indices:R.Tensor((m, n), "int64"), upda class Expected: @Ts.prim_func(private=True) def scatter_elements( - rxplaceholder: T.Buffer((a_scatter_elements, b_scatter_elements), offset_factor=1), - rxplaceholder_1: T.Buffer((m_scatter_elements, n_scatter_elements), 'int64', offset_factor=1), - rxplaceholder_2: T.Buffer((m_scatter_elements, n_scatter_elements), offset_factor=1), - out_buf: T.Buffer((a_scatter_elements, b_scatter_elements)), + rxplaceholder: T.Tensor((a_scatter_elements, b_scatter_elements), offset_factor=1), + rxplaceholder_1: T.Tensor((m_scatter_elements, n_scatter_elements), 'int64', offset_factor=1), + rxplaceholder_2: T.Tensor((m_scatter_elements, n_scatter_elements), offset_factor=1), + out_buf: T.Tensor((a_scatter_elements, b_scatter_elements)), ): T.func_attr({"tirx.noalias": True}) @@ -1677,7 +1677,7 @@ def main(x: R.Tensor((10, 21, 30), "float32")): @I.ir_module class Expected: @Ts.prim_func(private=True) - def te_layout_transform(A: T.Buffer((T.int64(10), T.int64(21), T.int64(30)), "float32"), te_layout_transform_1: T.Buffer((T.int64(10), T.int64(30), T.int64(7), T.int64(3)), "float32")): + def te_layout_transform(A: T.Tensor((T.int64(10), T.int64(21), T.int64(30)), "float32"), te_layout_transform_1: T.Tensor((T.int64(10), T.int64(30), T.int64(7), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1, i2 in T.grid(T.int64(10), T.int64(21), T.int64(30)): @@ -1715,7 +1715,7 @@ def main(x: R.Tensor((10, 20, 30), "float32")): @I.ir_module class Expected: @Ts.prim_func(private=True) - def te_layout_transform_with_pad(A: T.Buffer((T.int64(10), T.int64(20), T.int64(30)), "float32"), te_layout_transform_with_pad_1: T.Buffer((T.int64(10), T.int64(30), T.int64(7), T.int64(3)), "float32")): + def te_layout_transform_with_pad(A: T.Tensor((T.int64(10), T.int64(20), T.int64(30)), "float32"), te_layout_transform_with_pad_1: T.Tensor((T.int64(10), T.int64(30), T.int64(7), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for axis0, axis1, axis2, axis3 in T.grid(T.int64(10), T.int64(30), T.int64(7), T.int64(3)): @@ -1764,7 +1764,7 @@ def main(x: R.Tensor((a, b, c), "float32")): @I.ir_module class Expected: @Ts.prim_func(private=True) - def te_layout_transform_with_pad(A: T.Buffer((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.Buffer((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)))): + def te_layout_transform_with_pad(A: T.Tensor((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.Tensor((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)))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1840,8 +1840,8 @@ def main( @Ts.prim_func(private=True) def te_layout_transform( - A: T.Buffer((T.int64(16),), "float32"), - te_layout_transform: T.Buffer((T.int64(4), T.int64(4)), "float32"), + A: T.Tensor((T.int64(16),), "float32"), + te_layout_transform: T.Tensor((T.int64(4), T.int64(4)), "float32"), ): T.func_attr({"tirx.noalias": True}) for i in range(T.int64(16)): @@ -1881,7 +1881,7 @@ def main( return gv @Ts.prim_func(private=True) - def scatter_nd(data: T.Buffer((T.int64(8),), offset_factor=1), indices: T.Buffer((T.int64(4), T.int64(1)), 'int64'), updates: T.Buffer((T.int64(4),), offset_factor=1), out_buf: T.Buffer((T.int64(8),))): + def scatter_nd(data: T.Tensor((T.int64(8),), offset_factor=1), indices: T.Tensor((T.int64(4), T.int64(1)), 'int64'), updates: T.Tensor((T.int64(4),), offset_factor=1), out_buf: T.Tensor((T.int64(8),))): T.func_attr({"tirx.noalias": True}) with Ts.sblock("root"): diff --git a/tests/python/relax/test_transform_legalize_ops_nn.py b/tests/python/relax/test_transform_legalize_ops_nn.py index d00ae6912adc..4a8db5da9565 100644 --- a/tests/python/relax/test_transform_legalize_ops_nn.py +++ b/tests/python/relax/test_transform_legalize_ops_nn.py @@ -46,7 +46,7 @@ def main(x: R.Tensor((2, 128, 28), dtype="float32"), w: R.Tensor((64, 16, 3), dt return gv @Ts.prim_func(private=True) - def conv1d(A: T.Buffer((T.int64(2), T.int64(128), T.int64(28)), "float32"), B: T.Buffer((T.int64(64), T.int64(16), T.int64(3)), "float32"), group_conv1d_ncw: T.Buffer((T.int64(2), T.int64(64), T.int64(13)), "float32")): + def conv1d(A: T.Tensor((T.int64(2), T.int64(128), T.int64(28)), "float32"), B: T.Tensor((T.int64(64), T.int64(16), T.int64(3)), "float32"), group_conv1d_ncw: T.Tensor((T.int64(2), T.int64(64), T.int64(13)), "float32")): T.func_attr({"tirx.noalias": True}) pad_temp = Ts.sblock_alloc_buffer((T.int64(2), T.int64(128), T.int64(30))) for i0, i1, i2 in T.grid(T.int64(2), T.int64(128), T.int64(30)): @@ -86,7 +86,7 @@ def main(x: R.Tensor((2, 3, 28), dtype="float32"), w: R.Tensor((4, 3, 3), dtype= return gv @Ts.prim_func(private=True) - def conv1d(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(28)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(3)), "float32"), conv1d_ncw: T.Buffer((T.int64(2), T.int64(4), T.int64(26)), "float16")): + def conv1d(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(28)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(3)), "float32"), conv1d_ncw: T.Tensor((T.int64(2), T.int64(4), T.int64(26)), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): pad_temp = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3), T.int64(28))) @@ -127,7 +127,7 @@ def main(x: R.Tensor((2, 28, 128), dtype="float32"), w: R.Tensor((64, 128, 3), d return gv @Ts.prim_func(private=True) - def conv1d(rxplaceholder: T.Buffer((T.int64(2), T.int64(28), T.int64(128)), "float32"), rxplaceholder_1: T.Buffer((T.int64(64), T.int64(128), T.int64(3)), "float32"), conv1d_nwc: T.Buffer((T.int64(2), T.int64(26), T.int64(64)), "float32")): + def conv1d(rxplaceholder: T.Tensor((T.int64(2), T.int64(28), T.int64(128)), "float32"), rxplaceholder_1: T.Tensor((T.int64(64), T.int64(128), T.int64(3)), "float32"), conv1d_nwc: T.Tensor((T.int64(2), T.int64(26), T.int64(64)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): pad_temp = Ts.sblock_alloc_buffer((T.int64(2), T.int64(28), T.int64(128))) @@ -185,7 +185,7 @@ def main(x: R.Tensor((n_main, c_main, w_main), dtype="float32"), kernel: R.Tenso return gv @Ts.prim_func(private=True) - def conv1d(rxplaceholder: T.Buffer((n_conv1d, c_conv1d, w_conv1d)), rxplaceholder_1: T.Buffer((f_conv1d, c_conv1d, kw_conv1d)), conv1d_ncw: T.Buffer((n_conv1d, f_conv1d, w_conv1d + T.int64(1) - kw_conv1d))): + def conv1d(rxplaceholder: T.Tensor((n_conv1d, c_conv1d, w_conv1d)), rxplaceholder_1: T.Tensor((f_conv1d, c_conv1d, kw_conv1d)), conv1d_ncw: T.Tensor((n_conv1d, f_conv1d, w_conv1d + T.int64(1) - kw_conv1d))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -222,7 +222,7 @@ def main(x: R.Tensor((2, 128, 28), "float32"), w: R.Tensor((128, 16, 3), "float3 @I.ir_module class Expected: @Ts.prim_func(private=True) - def conv1d_transpose(x: T.Buffer((T.int64(2), T.int64(128), T.int64(28)), "float32"), w: T.Buffer((T.int64(128), T.int64(16), T.int64(3)), "float32"), compute: T.Buffer((T.int64(2), T.int64(128), T.int64(56)), "float32")): + def conv1d_transpose(x: T.Tensor((T.int64(2), T.int64(128), T.int64(28)), "float32"), w: T.Tensor((T.int64(128), T.int64(16), T.int64(3)), "float32"), compute: T.Tensor((T.int64(2), T.int64(128), T.int64(56)), "float32")): T.func_attr({"tirx.noalias": True}) data_dilate = Ts.sblock_alloc_buffer((T.int64(2), T.int64(128), T.int64(55))) data_pad = Ts.sblock_alloc_buffer((T.int64(2), T.int64(128), T.int64(58))) @@ -274,7 +274,7 @@ def main(x: R.Tensor((2, 128, 28, 28), "float32"), w: R.Tensor((64, 16, 3, 3), " return gv @Ts.prim_func(private=True) - def conv2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(128), T.int64(28), T.int64(28)), "float32"), rxplaceholder_1: T.Buffer((T.int64(64), T.int64(16), T.int64(3), T.int64(3)), "float32"), group_conv2d_nchw: T.Buffer((T.int64(2), T.int64(64), T.int64(13), T.int64(13)), "float32")): + def conv2d(rxplaceholder: T.Tensor((T.int64(2), T.int64(128), T.int64(28), T.int64(28)), "float32"), rxplaceholder_1: T.Tensor((T.int64(64), T.int64(16), T.int64(3), T.int64(3)), "float32"), group_conv2d_nchw: T.Tensor((T.int64(2), T.int64(64), T.int64(13), T.int64(13)), "float32")): T.func_attr({"tirx.noalias": True}) pad_temp = Ts.sblock_alloc_buffer([T.int64(2), T.int64(128), T.int64(30), T.int64(30)], dtype="float32") for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(128), T.int64(30), T.int64(30)): @@ -314,7 +314,7 @@ def main(x: R.Tensor((2, 3, 28, 28), "float32"), w: R.Tensor((4, 3, 3, 3), "floa return gv @Ts.prim_func(private=True) - def conv2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(28), T.int64(28)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(3), T.int64(3)), "float32"), conv2d_nchw: T.Buffer((T.int64(2), T.int64(4), T.int64(26), T.int64(26)), "float16")): + def conv2d(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(28), T.int64(28)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(3), T.int64(3)), "float32"), conv2d_nchw: T.Tensor((T.int64(2), T.int64(4), T.int64(26), T.int64(26)), "float16")): T.func_attr({"tirx.noalias": True}) pad_temp = Ts.sblock_alloc_buffer([T.int64(2), T.int64(3), T.int64(28), T.int64(28)], dtype="float32") for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(3), T.int64(28), T.int64(28)): @@ -354,7 +354,7 @@ def main(x: R.Tensor((2, 28, 28, 128), "float32"), w: R.Tensor((64, 128, 3, 3), return gv @Ts.prim_func(private=True) - def conv2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(28), T.int64(28), T.int64(128)), "float32"), rxplaceholder_1: T.Buffer((T.int64(64), T.int64(128), T.int64(3), T.int64(3)), "float32"), conv2d_nhwc: T.Buffer((T.int64(2), T.int64(26), T.int64(26), T.int64(64)), "float32")): + def conv2d(rxplaceholder: T.Tensor((T.int64(2), T.int64(28), T.int64(28), T.int64(128)), "float32"), rxplaceholder_1: T.Tensor((T.int64(64), T.int64(128), T.int64(3), T.int64(3)), "float32"), conv2d_nhwc: T.Tensor((T.int64(2), T.int64(26), T.int64(26), T.int64(64)), "float32")): T.func_attr({"tirx.noalias": True}) pad_temp = Ts.sblock_alloc_buffer([T.int64(2), T.int64(28), T.int64(28), T.int64(128)], dtype="float32") for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(28), T.int64(28), T.int64(128)): @@ -417,7 +417,7 @@ def main(x: R.Tensor((n_main, c_main, h_main, w_main), "float32"), kernel: R.Ten return gv @Ts.prim_func(private=True) - def conv2d(rxplaceholder: T.Buffer([n_conv2d, c_conv2d, h_conv2d, w_conv2d], dtype='float32'), rxplaceholder_1: T.Buffer([f_conv2d, c_conv2d, kh_conv2d, kw_conv2d], dtype='float32'), conv2d_nchw: T.Buffer([n_conv2d, f_conv2d, h_conv2d + T.int64(1) - kh_conv2d, w_conv2d + T.int64(1) - kw_conv2d], dtype='float32')): + def conv2d(rxplaceholder: T.Tensor([n_conv2d, c_conv2d, h_conv2d, w_conv2d], dtype='float32'), rxplaceholder_1: T.Tensor([f_conv2d, c_conv2d, kh_conv2d, kw_conv2d], dtype='float32'), conv2d_nchw: T.Tensor([n_conv2d, f_conv2d, h_conv2d + T.int64(1) - kh_conv2d, w_conv2d + T.int64(1) - kw_conv2d], dtype='float32')): T.func_attr({"tirx.noalias": True}) pad_temp = Ts.sblock_alloc_buffer([n_conv2d, c_conv2d, h_conv2d, w_conv2d], dtype="float32") @@ -472,7 +472,7 @@ def main(x: R.Tensor((n_main, c_main, 28, 28), dtype="float32"), w: R.Tensor((f_ return gv @Ts.prim_func(private=True) - def conv2d(x: T.Buffer((n_conv2d, c_conv2d, T.int64(28), T.int64(28))), w: T.Buffer((f_conv2d, c_div_8_conv2d, T.int64(3), T.int64(3))), group_conv2d_nchw: T.Buffer((n_conv2d, f_conv2d, T.int64(26), T.int64(26)))): + def conv2d(x: T.Tensor((n_conv2d, c_conv2d, T.int64(28), T.int64(28))), w: T.Tensor((f_conv2d, c_div_8_conv2d, T.int64(3), T.int64(3))), group_conv2d_nchw: T.Tensor((n_conv2d, f_conv2d, T.int64(26), T.int64(26)))): T.func_attr({"tirx.noalias": True}) pad_temp = Ts.sblock_alloc_buffer((n_conv2d, c_conv2d, T.int64(28), T.int64(28))) @@ -513,7 +513,7 @@ def main(x: R.Tensor((2, 128, 28, 28), dtype="float32"), w: R.Tensor((128, 16, 3 return gv @Ts.prim_func(private=True) - def conv2d_transpose(rxplaceholder: T.Buffer((T.int64(2), T.int64(128), T.int64(28), T.int64(28)), "float32"), rxplaceholder_1: T.Buffer((T.int64(128), T.int64(16), T.int64(3), T.int64(3)), "float32"), compute: T.Buffer((T.int64(2), T.int64(128), T.int64(56), T.int64(84)), "float32")): + def conv2d_transpose(rxplaceholder: T.Tensor((T.int64(2), T.int64(128), T.int64(28), T.int64(28)), "float32"), rxplaceholder_1: T.Tensor((T.int64(128), T.int64(16), T.int64(3), T.int64(3)), "float32"), compute: T.Tensor((T.int64(2), T.int64(128), T.int64(56), T.int64(84)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): data_dilate = Ts.sblock_alloc_buffer((T.int64(2), T.int64(128), T.int64(55), T.int64(82))) @@ -568,7 +568,7 @@ def main(x: R.Tensor((2, 3, 4, 4, 4), dtype="float32"), w: R.Tensor((3, 4, 3, 3, return gv @Ts.prim_func(private=True) - def conv3d_transpose(x: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(4), T.int64(4)), "float32"), w: T.Buffer((T.int64(3), T.int64(4), T.int64(3), T.int64(3), T.int64(3)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4), T.int64(6), T.int64(6), T.int64(6)), "float32")): + def conv3d_transpose(x: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(4), T.int64(4)), "float32"), w: T.Tensor((T.int64(3), T.int64(4), T.int64(3), T.int64(3), T.int64(3)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4), T.int64(6), T.int64(6), T.int64(6)), "float32")): T.func_attr({"tirx.noalias": True}) data_dilate = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(4), T.int64(4))) data_pad = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3), T.int64(8), T.int64(8), T.int64(8))) @@ -622,7 +622,7 @@ def main(x: R.Tensor((2, 3, 4, 4, 4), dtype="float32"), w: R.Tensor((3, 4, 3, 3, return gv @Ts.prim_func(private=True) - def conv3d_transpose(x: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(4), T.int64(4)), "float32"), w: T.Buffer((T.int64(3), T.int64(4), T.int64(3), T.int64(3), T.int64(3)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4), T.int64(6), T.int64(6), T.int64(6)), "float16")): + def conv3d_transpose(x: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(4), T.int64(4)), "float32"), w: T.Tensor((T.int64(3), T.int64(4), T.int64(3), T.int64(3), T.int64(3)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4), T.int64(6), T.int64(6), T.int64(6)), "float16")): T.func_attr({"tirx.noalias": True}) data_dilate = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(4), T.int64(4))) data_pad = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3), T.int64(8), T.int64(8), T.int64(8))) @@ -676,7 +676,7 @@ def main(x: R.Tensor((2, 3, 28, 28), dtype="float32"), w: R.Tensor((3, 4, 3, 3), return gv @Ts.prim_func(private=True) - def conv2d_transpose(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(28), T.int64(28)), "float32"), rxplaceholder_1: T.Buffer((T.int64(3), T.int64(4), T.int64(3), T.int64(3)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4), T.int64(30), T.int64(30)), "float16")): + def conv2d_transpose(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(28), T.int64(28)), "float32"), rxplaceholder_1: T.Tensor((T.int64(3), T.int64(4), T.int64(3), T.int64(3)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4), T.int64(30), T.int64(30)), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): data_dilate = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3), T.int64(28), T.int64(28))) @@ -754,7 +754,7 @@ def main(x: R.Tensor((n_main, c_main, h_main, w_main), dtype="float32"), kernel: return gv @Ts.prim_func(private=True) - def conv2d_transpose(rxplaceholder: T.Buffer((n_conv2d_transpose, c_conv2d_transpose, h_conv2d_transpose, w_conv2d_transpose)), rxplaceholder_1: T.Buffer((f_conv2d_transpose, c_conv2d_transpose, kh_conv2d_transpose, kw_conv2d_transpose)), compute: T.Buffer((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)))): + def conv2d_transpose(rxplaceholder: T.Tensor((n_conv2d_transpose, c_conv2d_transpose, h_conv2d_transpose, w_conv2d_transpose)), rxplaceholder_1: T.Tensor((f_conv2d_transpose, c_conv2d_transpose, kh_conv2d_transpose, kw_conv2d_transpose)), compute: T.Tensor((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)))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -805,7 +805,7 @@ def main(x: R.Tensor((1, 1, 3, 3), "float32"), w: R.Tensor((1, 1, 2, 2), "float3 @I.ir_module class Expected: @Ts.prim_func(private=True) - def conv2d_transpose(x: T.Buffer((T.int64(1), T.int64(1), T.int64(3), T.int64(3)), "float32"), w: T.Buffer((T.int64(1), T.int64(1), T.int64(2), T.int64(2)), "float32"), compute: T.Buffer((T.int64(1), T.int64(1), T.int64(5), T.int64(5)), "float32")): + def conv2d_transpose(x: T.Tensor((T.int64(1), T.int64(1), T.int64(3), T.int64(3)), "float32"), w: T.Tensor((T.int64(1), T.int64(1), T.int64(2), T.int64(2)), "float32"), compute: T.Tensor((T.int64(1), T.int64(1), T.int64(5), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) data_dilate = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(3), T.int64(3))) data_pad = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(7), T.int64(7))) @@ -862,7 +862,7 @@ def main(x: R.Tensor((4, 112, 112, 6), "float32")) -> R.Tensor((4, 56, 56, 6), " return gv @Ts.prim_func(private=True) - def max_pool2d(rxplaceholder: T.Buffer((T.int64(4), T.int64(112), T.int64(112), T.int64(6)), "float32"), pool_max: T.Buffer((T.int64(4), T.int64(56), T.int64(56), T.int64(6)), "float32")): + def max_pool2d(rxplaceholder: T.Tensor((T.int64(4), T.int64(112), T.int64(112), T.int64(6)), "float32"), pool_max: T.Tensor((T.int64(4), T.int64(56), T.int64(56), T.int64(6)), "float32")): T.func_attr({"tirx.noalias": True}) pad_temp = Ts.sblock_alloc_buffer([T.int64(4), T.int64(114), T.int64(114), T.int64(6)], dtype="float32") for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(114), T.int64(114), T.int64(6)): @@ -903,7 +903,7 @@ def main(x: R.Tensor((4, 4, 112, 112, 16), "float32")) -> R.Tensor((4, 4, 110, 1 return gv @Ts.prim_func(private=True) - def max_pool2d(rxplaceholder: T.Buffer((T.int64(4), T.int64(4), T.int64(112), T.int64(112), T.int64(16)), "float32"), pool_max: T.Buffer((T.int64(4), T.int64(4), T.int64(110), T.int64(110), T.int64(16)), "float32")): + def max_pool2d(rxplaceholder: T.Tensor((T.int64(4), T.int64(4), T.int64(112), T.int64(112), T.int64(16)), "float32"), pool_max: T.Tensor((T.int64(4), T.int64(4), T.int64(110), T.int64(110), T.int64(16)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3, i4, i5, i6 in T.grid(T.int64(4), T.int64(4), T.int64(110), T.int64(110), T.int64(16), T.int64(3), T.int64(3)): with Ts.sblock("pool_max"): @@ -937,7 +937,7 @@ def main(x: R.Tensor((4, 6, 112, 112), dtype="float32")) -> R.Tensor((4, 6, 38, return gv @Ts.prim_func(private=True) - def max_pool2d(rxplaceholder: T.Buffer((T.int64(4), T.int64(6), T.int64(112), T.int64(112)), "float32"), pool_max: T.Buffer((T.int64(4), T.int64(6), T.int64(38), T.int64(38)), "float32")): + def max_pool2d(rxplaceholder: T.Tensor((T.int64(4), T.int64(6), T.int64(112), T.int64(112)), "float32"), pool_max: T.Tensor((T.int64(4), T.int64(6), T.int64(38), T.int64(38)), "float32")): T.func_attr({"tirx.noalias": True}) pad_temp = Ts.sblock_alloc_buffer([T.int64(4), T.int64(6), T.int64(116), T.int64(116)], dtype="float32") for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(6), T.int64(116), T.int64(116)): @@ -996,7 +996,7 @@ def main(x: R.Tensor((4, 112, 112, 6), "float32")) -> R.Tensor((4, 56, 56, 6), " @I.ir_module class Expected: @Ts.prim_func(private=True) - def avg_pool2d(rxplaceholder: T.Buffer((T.int64(4), T.int64(112), T.int64(112), T.int64(6)), "float32"), pool_avg: T.Buffer((T.int64(4), T.int64(56), T.int64(56), T.int64(6)), "float32")): + def avg_pool2d(rxplaceholder: T.Tensor((T.int64(4), T.int64(112), T.int64(112), T.int64(6)), "float32"), pool_avg: T.Tensor((T.int64(4), T.int64(56), T.int64(56), T.int64(6)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): pad_temp = Ts.sblock_alloc_buffer((T.int64(4), T.int64(114), T.int64(114), T.int64(6))) @@ -1045,7 +1045,7 @@ def main(x: R.Tensor((4, 4, 112, 112, 16), "float32")) -> R.Tensor((4, 4, 110, 1 @I.ir_module class Expected: @Ts.prim_func(private=True) - def avg_pool2d(rxplaceholder: T.Buffer((T.int64(4), T.int64(4), T.int64(112), T.int64(112), T.int64(16)), "float32"), pool_avg: T.Buffer((T.int64(4), T.int64(4), T.int64(110), T.int64(110), T.int64(16)), "float32")): + def avg_pool2d(rxplaceholder: T.Tensor((T.int64(4), T.int64(4), T.int64(112), T.int64(112), T.int64(16)), "float32"), pool_avg: T.Tensor((T.int64(4), T.int64(4), T.int64(110), T.int64(110), T.int64(16)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): pool_sum = Ts.sblock_alloc_buffer((T.int64(4), T.int64(4), T.int64(110), T.int64(110), T.int64(16))) @@ -1086,7 +1086,7 @@ def main(x: R.Tensor((4, 6, 112, 112), "float32")) -> R.Tensor((4, 6, 38, 38), " @I.ir_module class Expected: @Ts.prim_func(private=True) - def avg_pool2d(rxplaceholder: T.Buffer((T.int64(4), T.int64(6), T.int64(112), T.int64(112)), "float32"), pool_avg: T.Buffer((T.int64(4), T.int64(6), T.int64(38), T.int64(38)), "float32")): + def avg_pool2d(rxplaceholder: T.Tensor((T.int64(4), T.int64(6), T.int64(112), T.int64(112)), "float32"), pool_avg: T.Tensor((T.int64(4), T.int64(6), T.int64(38), T.int64(38)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): pad_temp = Ts.sblock_alloc_buffer((T.int64(4), T.int64(6), T.int64(116), T.int64(116))) @@ -1164,7 +1164,7 @@ def main(x: R.Tensor((2, 4, 7, 7, 16), "float32")) -> R.Tensor((2, 4, 1, 1, 16), return gv @Ts.prim_func(private=True) - def adaptive_avg_pool2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(4), T.int64(7), T.int64(7), T.int64(16)), "float32"), adaptive_pool_avg: T.Buffer((T.int64(2), T.int64(4), T.int64(1), T.int64(1), T.int64(16)), "float32")): + def adaptive_avg_pool2d(rxplaceholder: T.Tensor((T.int64(2), T.int64(4), T.int64(7), T.int64(7), T.int64(16)), "float32"), adaptive_pool_avg: T.Tensor((T.int64(2), T.int64(4), T.int64(1), T.int64(1), T.int64(16)), "float32")): T.func_attr({"tirx.noalias": True}) adaptive_pool_sum = Ts.sblock_alloc_buffer([T.int64(2), T.int64(4), T.int64(1), T.int64(1), T.int64(16)], dtype="float32") for i0, i1, i2, i3, i4, i5, i6 in T.grid(T.int64(2), T.int64(4), T.int64(1), T.int64(1), T.int64(16), T.int64(7), T.int64(7)): @@ -1205,7 +1205,7 @@ def main(x: R.Tensor((2, 16, 7, 7), "float32")) -> R.Tensor((2, 16, 7, 7), "floa return gv @Ts.prim_func(private=True) - def adaptive_avg_pool2d(rxplaceholder: T.Buffer((T.int64(2), T.int64(16), T.int64(7), T.int64(7)), "float32"), adaptive_pool_avg: T.Buffer((T.int64(2), T.int64(16), T.int64(7), T.int64(7)), "float32")): + def adaptive_avg_pool2d(rxplaceholder: T.Tensor((T.int64(2), T.int64(16), T.int64(7), T.int64(7)), "float32"), adaptive_pool_avg: T.Tensor((T.int64(2), T.int64(16), T.int64(7), T.int64(7)), "float32")): T.func_attr({"tirx.noalias": True}) adaptive_pool_sum = Ts.sblock_alloc_buffer([T.int64(2), T.int64(16), T.int64(7), T.int64(7)], dtype="float32") for i0, i1, i2, i3, i4, i5 in T.grid(T.int64(2), T.int64(16), T.int64(7), T.int64(7), T.int64(1), T.int64(1)): @@ -1268,7 +1268,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def relu(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), compute: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def relu(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), compute: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("compute"): @@ -1307,7 +1307,7 @@ def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), return gv @Ts.prim_func(private=True) - def relu(rxplaceholder: T.Buffer([m_relu, n_relu], dtype='float32'), compute: T.Buffer([m_relu, n_relu], dtype='float32')): + def relu(rxplaceholder: T.Tensor([m_relu, n_relu], dtype='float32'), compute: T.Tensor([m_relu, n_relu], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(m_relu, n_relu): @@ -1339,7 +1339,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def leaky_relu(x: T.Buffer((T.int64(2), T.int64(3)), "float32"), compute: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def leaky_relu(x: T.Tensor((T.int64(2), T.int64(3)), "float32"), compute: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("compute"): @@ -1378,7 +1378,7 @@ def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), return gv @Ts.prim_func(private=True) - def leaky_relu(x: T.Buffer((m_leaky_relu, n_leaky_relu)), compute: T.Buffer((m_leaky_relu, n_leaky_relu))): + def leaky_relu(x: T.Tensor((m_leaky_relu, n_leaky_relu)), compute: T.Tensor((m_leaky_relu, n_leaky_relu))): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(m_leaky_relu, n_leaky_relu): @@ -1410,7 +1410,7 @@ def main(x: R.Tensor((2, 3), dtype="float32"), y: R.Tensor((1,), dtype="float32" return gv @Ts.prim_func(private=True) - def prelu(x: T.Buffer((T.int64(2), T.int64(3)), "float32"), y: T.Buffer((T.int64(1),), "float32"), compute: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def prelu(x: T.Tensor((T.int64(2), T.int64(3)), "float32"), y: T.Tensor((T.int64(1),), "float32"), compute: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): slope_broadcasted = Ts.sblock_alloc_buffer((T.int64(3),)) @@ -1454,7 +1454,7 @@ def main(x: R.Tensor((m_main, 7), dtype="float32"), y: R.Tensor((1,), dtype="flo return gv @Ts.prim_func(private=True) - def prelu(x: T.Buffer((m_prelu, T.int64(7))), y: T.Buffer((T.int64(1),), "float32"), compute: T.Buffer((m_prelu, T.int64(7)))): + def prelu(x: T.Tensor((m_prelu, T.int64(7))), y: T.Tensor((T.int64(1),), "float32"), compute: T.Tensor((m_prelu, T.int64(7)))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1494,7 +1494,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def gelu(x: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_multiply: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def gelu(x: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_multiply: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) T_multiply_1 = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3))) compute = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3))) @@ -1561,7 +1561,7 @@ def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), return gv @Ts.prim_func(private=True) - def gelu(x: T.Buffer((m_gelu, n_gelu)), T_multiply: T.Buffer((m_gelu, n_gelu))): + def gelu(x: T.Tensor((m_gelu, n_gelu)), T_multiply: T.Tensor((m_gelu, n_gelu))): T.func_attr({"tirx.noalias": True}) T_multiply_1 = Ts.sblock_alloc_buffer((m_gelu, n_gelu)) @@ -1621,7 +1621,7 @@ def main(x: R.Tensor((2, 3), dtype="float32")) -> R.Tensor((2, 3), dtype="float3 return gv @Ts.prim_func(private=True) - def gelu_tanh(A: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_multiply: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def gelu_tanh(A: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_multiply: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) T_multiply_1 = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3))) T_multiply_2 = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3))) @@ -1715,7 +1715,7 @@ def main(x: R.Tensor((m_main, n_main), dtype="float32")) -> R.Tensor((m_main, n_ return gv @Ts.prim_func(private=True) - def gelu_tanh(A: T.Buffer((m_gelu_tanh, n_gelu_tanh)), T_multiply: T.Buffer((m_gelu_tanh, n_gelu_tanh))): + def gelu_tanh(A: T.Tensor((m_gelu_tanh, n_gelu_tanh)), T_multiply: T.Tensor((m_gelu_tanh, n_gelu_tanh))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1803,7 +1803,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tensor((2, 3), "float32"): return gv @Ts.prim_func(private=True) - def silu(rxplaceholder: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_multiply: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def silu(rxplaceholder: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_multiply: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) compute = Ts.sblock_alloc_buffer([T.int64(2), T.int64(3)], dtype="float32") for i0, i1 in T.grid(T.int64(2), T.int64(3)): @@ -1849,7 +1849,7 @@ def main(x: R.Tensor((m_main, n_main), "float32")) -> R.Tensor((m_main, n_main), return gv @Ts.prim_func(private=True) - def silu(rxplaceholder: T.Buffer([m_silu, n_silu], dtype='float32'), T_multiply: T.Buffer([m_silu, n_silu], dtype='float32')): + def silu(rxplaceholder: T.Tensor([m_silu, n_silu], dtype='float32'), T_multiply: T.Tensor([m_silu, n_silu], dtype='float32')): T.func_attr({"tirx.noalias": True}) compute = Ts.sblock_alloc_buffer([m_silu, n_silu], dtype="float32") @@ -1888,7 +1888,7 @@ def main(x: R.Tensor((2, 3, 16, 32), "float32")) -> R.Tensor((2, 3, 16, 32), "fl return gv @Ts.prim_func(private=True) - def softmax(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(16), T.int64(32)), "float32"), T_softmax_norm: T.Buffer((T.int64(2), T.int64(3), T.int64(16), T.int64(32)), "float32")): + def softmax(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(16), T.int64(32)), "float32"), T_softmax_norm: T.Tensor((T.int64(2), T.int64(3), T.int64(16), T.int64(32)), "float32")): T.func_attr({"tirx.noalias": True}) T_softmax_maxelem = Ts.sblock_alloc_buffer([T.int64(2), T.int64(3), T.int64(32)], dtype="float32") T_softmax_exp = Ts.sblock_alloc_buffer([T.int64(2), T.int64(3), T.int64(16), T.int64(32)], dtype="float32") @@ -1956,7 +1956,7 @@ def main(x: R.Tensor((a_main, b_main, c_main), "float32")) -> R.Tensor((a_main, return gv @Ts.prim_func(private=True) - def softmax(rxplaceholder: T.Buffer([a_softmax, b_softmax, c_softmax], dtype='float32'), T_softmax_norm: T.Buffer([a_softmax, b_softmax, c_softmax], dtype='float32')): + def softmax(rxplaceholder: T.Tensor([a_softmax, b_softmax, c_softmax], dtype='float32'), T_softmax_norm: T.Tensor([a_softmax, b_softmax, c_softmax], dtype='float32')): T.func_attr({"tirx.noalias": True}) T_softmax_maxelem = Ts.sblock_alloc_buffer([a_softmax, b_softmax], dtype="float32") @@ -2014,7 +2014,7 @@ def main(x: R.Tensor((2, 3, 16, 32), dtype="float32")) -> R.Tensor((2, 3, 16, 32 return gv @Ts.prim_func(private=True) - def log_softmax(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(16), T.int64(32)), "float32"), compute: T.Buffer((T.int64(2), T.int64(3), T.int64(16), T.int64(32)), "float32"),): + def log_softmax(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(16), T.int64(32)), "float32"), compute: T.Tensor((T.int64(2), T.int64(3), T.int64(16), T.int64(32)), "float32"),): T.func_attr({"tirx.noalias": True}) T_softmax_maxelem = Ts.sblock_alloc_buffer([T.int64(2), T.int64(3), T.int64(32)], dtype="float32") compute_1 = Ts.sblock_alloc_buffer([T.int64(2), T.int64(3), T.int64(32)], dtype="float32") @@ -2076,7 +2076,7 @@ def main(x: R.Tensor((a_main, b_main, c_main), dtype="float32")) -> R.Tensor((a_ return gv @Ts.prim_func(private=True) - def log_softmax(rxplaceholder: T.Buffer([a_log_softmax, b_log_softmax, c_log_softmax], dtype='float32'), compute: T.Buffer([a_log_softmax, b_log_softmax, c_log_softmax], dtype='float32')): + def log_softmax(rxplaceholder: T.Tensor([a_log_softmax, b_log_softmax, c_log_softmax], dtype='float32'), compute: T.Tensor([a_log_softmax, b_log_softmax, c_log_softmax], dtype='float32')): T.func_attr({"tirx.noalias": True}) T_softmax_maxelem = Ts.sblock_alloc_buffer([a_log_softmax, b_log_softmax], dtype="float32") @@ -2127,7 +2127,7 @@ def main(x: R.Tensor((3,), dtype="float32"), y: R.Tensor((3,), dtype="float32")) return gv @Ts.prim_func(private=True) - def cross_entropy_with_logits(x: T.Buffer((T.int64(3),), "float32"), y: T.Buffer((T.int64(3),), "float32"), T_multiply: T.Buffer((), "float32")): + def cross_entropy_with_logits(x: T.Tensor((T.int64(3),), "float32"), y: T.Tensor((T.int64(3),), "float32"), T_multiply: T.Tensor((), "float32")): T.func_attr({"tirx.noalias": True}) T_multiply_1 = Ts.sblock_alloc_buffer((T.int64(3),)) T_multiply_red = Ts.sblock_alloc_buffer(()) @@ -2173,7 +2173,7 @@ def main(x: R.Tensor((2, 3), dtype="float32"), y: R.Tensor((2, 3), dtype="float3 return gv @Ts.prim_func(private=True) - def cross_entropy_with_logits(x: T.Buffer((T.int64(2), T.int64(3)), "float32"), y: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_divide: T.Buffer((), "float32")): + def cross_entropy_with_logits(x: T.Tensor((T.int64(2), T.int64(3)), "float32"), y: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_divide: T.Tensor((), "float32")): T.func_attr({"tirx.noalias": True}) T_multiply = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3))) T_multiply_red = Ts.sblock_alloc_buffer(()) @@ -2233,7 +2233,7 @@ def main(x: R.Tensor((n_main, m_main), dtype="float32"), y: R.Tensor((n_main, m_ return gv @Ts.prim_func(private=True) - def cross_entropy_with_logits(x: T.Buffer((n_cross_entropy_with_logits, m_cross_entropy_with_logits)), y: T.Buffer((n_cross_entropy_with_logits, m_cross_entropy_with_logits)), T_divide: T.Buffer((), "float32")): + def cross_entropy_with_logits(x: T.Tensor((n_cross_entropy_with_logits, m_cross_entropy_with_logits)), y: T.Tensor((n_cross_entropy_with_logits, m_cross_entropy_with_logits)), T_divide: T.Tensor((), "float32")): T.func_attr({"tirx.noalias": True}) T_multiply = Ts.sblock_alloc_buffer((n_cross_entropy_with_logits, m_cross_entropy_with_logits)) @@ -2281,7 +2281,7 @@ def main(x: R.Tensor((2, 3, 28, 28), "float32"), gamma: R.Tensor((3,), "float32" @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def batch_norm(x: T.Buffer((T.int64(2), T.int64(3), T.int64(28), T.int64(28))), gamma: T.Buffer((T.int64(3),)), beta: T.Buffer((T.int64(3),)), moving_mean: T.Buffer((T.int64(3),)), moving_var: T.Buffer((T.int64(3),)), T_add: T.Buffer((T.int64(2), T.int64(3), T.int64(28), T.int64(28))), T_add_1: T.Buffer((T.int64(3),)), T_add_2: T.Buffer((T.int64(3),))): + def batch_norm(x: T.Tensor((T.int64(2), T.int64(3), T.int64(28), T.int64(28))), gamma: T.Tensor((T.int64(3),)), beta: T.Tensor((T.int64(3),)), moving_mean: T.Tensor((T.int64(3),)), moving_var: T.Tensor((T.int64(3),)), T_add: T.Tensor((T.int64(2), T.int64(3), T.int64(28), T.int64(28))), T_add_1: T.Tensor((T.int64(3),)), T_add_2: T.Tensor((T.int64(3),))): T.func_attr({"tirx.noalias": True}) with Ts.sblock("root"): @@ -2577,7 +2577,7 @@ def main(x: R.Tensor((n, h, w, c), "float32"), gamma: R.Tensor((c,), "float32"), @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def batch_norm(x: T.Buffer((n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)), gamma: T.Buffer((c_batch_norm,)), beta: T.Buffer((c_batch_norm,)), moving_mean: T.Buffer((c_batch_norm,)), moving_var: T.Buffer((c_batch_norm,)), T_add: T.Buffer((n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)), T_add_1: T.Buffer((T.max(c_batch_norm, h_batch_norm),)), T_add_2: T.Buffer((T.max(c_batch_norm, h_batch_norm),))): + def batch_norm(x: T.Tensor((n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)), gamma: T.Tensor((c_batch_norm,)), beta: T.Tensor((c_batch_norm,)), moving_mean: T.Tensor((c_batch_norm,)), moving_var: T.Tensor((c_batch_norm,)), T_add: T.Tensor((n_batch_norm, h_batch_norm, w_batch_norm, c_batch_norm)), T_add_1: T.Tensor((T.max(c_batch_norm, h_batch_norm),)), T_add_2: T.Tensor((T.max(c_batch_norm, h_batch_norm),))): T.func_attr({"tirx.noalias": True}) with Ts.sblock("root"): @@ -2863,7 +2863,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32"), gamma: R.Tensor((4, 5), "float32" return gv @Ts.prim_func(private=True) - def layer_norm(x: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), gamma: T.Buffer((T.int64(4), T.int64(5)), "float32"), beta: T.Buffer((T.int64(4), T.int64(5)), "float32"), T_layer_norm: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32")): + def layer_norm(x: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), gamma: T.Tensor((T.int64(4), T.int64(5)), "float32"), beta: T.Tensor((T.int64(4), T.int64(5)), "float32"), T_layer_norm: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): x_sum = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3))) @@ -2918,7 +2918,7 @@ def forward(x: R.Tensor((3,), dtype="float32"), layer_norm_weight: R.Tensor((3,) @I.ir_module class LayerNorm_1D_Expected: @Ts.prim_func(private=True) - def layer_norm(x: T.Buffer((T.int64(3),), "float32"), layer_norm_weight: T.Buffer((T.int64(3),), "float32"), layer_norm_bias: T.Buffer((T.int64(3),), "float32"), T_layer_norm: T.Buffer((T.int64(3),), "float32")): + def layer_norm(x: T.Tensor((T.int64(3),), "float32"), layer_norm_weight: T.Tensor((T.int64(3),), "float32"), layer_norm_bias: T.Tensor((T.int64(3),), "float32"), T_layer_norm: T.Tensor((T.int64(3),), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): x_sum = Ts.sblock_alloc_buffer(()) @@ -2979,10 +2979,10 @@ def main(x: R.Tensor((2, 3, 4, 5), "float16"), gamma: R.Tensor((4, 5), "float16" class Expected: @Ts.prim_func(private=True) def layer_norm( - x: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float16"), - gamma: T.Buffer((T.int64(4), T.int64(5)), "float16"), - beta: T.Buffer((T.int64(4), T.int64(5)), "float16"), - T_layer_norm: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float16"), + x: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float16"), + gamma: T.Tensor((T.int64(4), T.int64(5)), "float16"), + beta: T.Tensor((T.int64(4), T.int64(5)), "float16"), + T_layer_norm: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -3055,7 +3055,7 @@ def main(x: R.Tensor((n_main, s_main, f_main), "float32"), gamma: R.Tensor((s_ma return gv @Ts.prim_func(private=True) - def layer_norm(x: T.Buffer((n_layer_norm, s_layer_norm, f_layer_norm)), gamma: T.Buffer((s_layer_norm, f_layer_norm)), beta: T.Buffer((s_layer_norm, f_layer_norm)), T_layer_norm: T.Buffer((n_layer_norm, s_layer_norm, f_layer_norm))): + def layer_norm(x: T.Tensor((n_layer_norm, s_layer_norm, f_layer_norm)), gamma: T.Tensor((s_layer_norm, f_layer_norm)), beta: T.Tensor((s_layer_norm, f_layer_norm)), T_layer_norm: T.Tensor((n_layer_norm, s_layer_norm, f_layer_norm))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -3107,7 +3107,7 @@ def main(x: R.Tensor((2, 4, 4, 5), "float32"), gamma: R.Tensor((4,), "float32"), @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def group_norm(rxplaceholder: T.Buffer((T.int64(2), T.int64(4), T.int64(4), T.int64(5)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4),), "float32"), rxplaceholder_2: T.Buffer((T.int64(4),), "float32"), T_reshape: T.Buffer((T.int64(2), T.int64(4), T.int64(4), T.int64(5)), "float32")): + def group_norm(rxplaceholder: T.Tensor((T.int64(2), T.int64(4), T.int64(4), T.int64(5)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4),), "float32"), rxplaceholder_2: T.Tensor((T.int64(4),), "float32"), T_reshape: T.Tensor((T.int64(2), T.int64(4), T.int64(4), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) T_reshape_1 = Ts.sblock_alloc_buffer((T.int64(2), T.int64(2), T.int64(2), T.int64(4), T.int64(5))) rxplaceholder_red_temp_v0 = Ts.sblock_alloc_buffer((T.int64(2), T.int64(2))) @@ -3184,7 +3184,7 @@ def main(x: R.Tensor((2, 4, 4, 5), dtype="float16"), gamma: R.Tensor((4,), dtype return gv @Ts.prim_func(private=True) - def group_norm(rxplaceholder: T.Buffer((T.int64(2), T.int64(4), T.int64(4), T.int64(5)), "float16"), rxplaceholder_1: T.Buffer((T.int64(4),), "float16"), rxplaceholder_2: T.Buffer((T.int64(4),), "float16"), T_reshape: T.Buffer((T.int64(2), T.int64(4), T.int64(4), T.int64(5)), "float16")): + def group_norm(rxplaceholder: T.Tensor((T.int64(2), T.int64(4), T.int64(4), T.int64(5)), "float16"), rxplaceholder_1: T.Tensor((T.int64(4),), "float16"), rxplaceholder_2: T.Tensor((T.int64(4),), "float16"), T_reshape: T.Tensor((T.int64(2), T.int64(4), T.int64(4), T.int64(5)), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): T_reshape_1 = Ts.sblock_alloc_buffer((T.int64(2), T.int64(2), T.int64(2), T.int64(4), T.int64(5)), "float16") @@ -3275,7 +3275,7 @@ class Expected: group_norm_c = T.int64() @Ts.prim_func(private=True) - def group_norm(rxplaceholder: T.Buffer((n_group_norm, T.int64(4) * group_norm_c, h_group_norm, w_group_norm)), rxplaceholder_1: T.Buffer((T.int64(4) * group_norm_c,)), rxplaceholder_2: T.Buffer((T.int64(4) * group_norm_c,)), c: group_norm_c, T_reshape: T.Buffer((n_group_norm, T.int64(4) * group_norm_c, h_group_norm, w_group_norm))): + def group_norm(rxplaceholder: T.Tensor((n_group_norm, T.int64(4) * group_norm_c, h_group_norm, w_group_norm)), rxplaceholder_1: T.Tensor((T.int64(4) * group_norm_c,)), rxplaceholder_2: T.Tensor((T.int64(4) * group_norm_c,)), c: group_norm_c, T_reshape: T.Tensor((n_group_norm, T.int64(4) * group_norm_c, h_group_norm, w_group_norm))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -3349,7 +3349,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32"), weight: R.Tensor((4, 5), "float32 @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def rms_norm(A: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), B: T.Buffer((T.int64(4), T.int64(5)), "float32"), T_cast: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32")): + def rms_norm(A: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), B: T.Tensor((T.int64(4), T.int64(5)), "float32"), T_cast: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): T_cast_1 = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5))) @@ -3425,7 +3425,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float16"), weight: R.Tensor((4, 5), "float16 @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def rms_norm(A: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float16"), B: T.Buffer((T.int64(4), T.int64(5)), "float16"), T_cast: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float16")): + def rms_norm(A: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float16"), B: T.Tensor((T.int64(4), T.int64(5)), "float16"), T_cast: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): T_cast_1 = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5))) @@ -3512,7 +3512,7 @@ def main(x: R.Tensor((n, s, f), "float32"), weight: R.Tensor((s, f), "float32")) @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def rms_norm(A: T.Buffer((n_rms_norm, s_rms_norm, f_rms_norm)), B: T.Buffer((s_rms_norm, f_rms_norm)), T_cast: T.Buffer((n_rms_norm, s_rms_norm, f_rms_norm))): + def rms_norm(A: T.Tensor((n_rms_norm, s_rms_norm, f_rms_norm)), B: T.Tensor((s_rms_norm, f_rms_norm)), T_cast: T.Tensor((n_rms_norm, s_rms_norm, f_rms_norm))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -3589,7 +3589,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32"), weight: R.Tensor((4, 5), "float32 @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def rms_norm(A: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), B: T.Buffer((T.int64(4), T.int64(5)), "float32"), T_cast: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32")): + def rms_norm(A: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), B: T.Tensor((T.int64(4), T.int64(5)), "float32"), T_cast: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): T_cast_1 = Ts.sblock_alloc_buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5))) @@ -3665,7 +3665,7 @@ def main(q: R.Tensor((4, 16, 32, 8), "float32"), k: R.Tensor((4, 8, 32, 8), "flo @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def attention_bias(q: T.Buffer((T.int64(4), T.int64(16), T.int64(32), T.int64(8)), "float32"), k: T.Buffer((T.int64(4), T.int64(8), T.int64(32), T.int64(8)), "float32"), v: T.Buffer((T.int64(4), T.int64(8), T.int64(32), T.int64(16)), "float32"), bias: T.Buffer((T.int64(4), T.int64(32), T.int64(16), T.int64(8)), "float32"), T_transpose: T.Buffer((T.int64(4), T.int64(16), T.int64(32), T.int64(16)), "float32")): + def attention_bias(q: T.Tensor((T.int64(4), T.int64(16), T.int64(32), T.int64(8)), "float32"), k: T.Tensor((T.int64(4), T.int64(8), T.int64(32), T.int64(8)), "float32"), v: T.Tensor((T.int64(4), T.int64(8), T.int64(32), T.int64(16)), "float32"), bias: T.Tensor((T.int64(4), T.int64(32), T.int64(16), T.int64(8)), "float32"), T_transpose: T.Tensor((T.int64(4), T.int64(16), T.int64(32), T.int64(16)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): T_transpose_1 = Ts.sblock_alloc_buffer((T.int64(4), T.int64(32), T.int64(16), T.int64(8))) @@ -3929,10 +3929,10 @@ def main( @Ts.prim_func(private=True) def nll_loss( - predictions: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), - targets: T.Buffer((T.int64(2), T.int64(4), T.int64(5)), "int64"), - weights: T.Buffer(T.int64(3), "float32"), - output: T.Buffer((), "float32"), + predictions: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), + targets: T.Tensor((T.int64(2), T.int64(4), T.int64(5)), "int64"), + weights: T.Tensor(T.int64(3), "float32"), + output: T.Tensor((), "float32"), ): # function attr dict T.func_attr({"tirx.noalias": True}) @@ -3998,7 +3998,7 @@ def main(predictions: R.Tensor((2, 3, 4, 5), dtype="float32"), targets: R.Tensor return gv @Ts.prim_func(private=True) - def nll_loss_without_weight(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_1: T.Buffer((T.int64(2), T.int64(4), T.int64(5)), "int64"), T_divide: T.Buffer((), "float32"),): + def nll_loss_without_weight(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_1: T.Tensor((T.int64(2), T.int64(4), T.int64(5)), "int64"), T_divide: T.Tensor((), "float32"),): # function attr dict T.func_attr({"tirx.noalias": True}) # body @@ -4075,7 +4075,7 @@ def main(predictions: R.Tensor((C_main,), dtype="float32"), targets: R.Tensor(() return gv @Ts.prim_func(private=True) - def nll_loss(rxplaceholder_1: T.Buffer((C_nll_loss,)), rxplaceholder: T.Buffer((), "int64"), rxplaceholder_2: T.Buffer((C_nll_loss,)), T_divide: T.Buffer((), "float32")): + def nll_loss(rxplaceholder_1: T.Tensor((C_nll_loss,)), rxplaceholder: T.Tensor((), "int64"), rxplaceholder_2: T.Tensor((C_nll_loss,)), T_divide: T.Tensor((), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -4134,7 +4134,7 @@ def main(predictions: R.Tensor((N_main, C_main, d1_main, d2_main), dtype="float3 return gv @Ts.prim_func(private=True) - def nll_loss(rxplaceholder: T.Buffer([N_nll_loss, C_nll_loss, d1_nll_loss, d2_nll_loss], dtype='float32'), rxplaceholder_1: T.Buffer([N_nll_loss, d1_nll_loss, d2_nll_loss], dtype='int64'), rxplaceholder_2: T.Buffer([C_nll_loss], dtype='float32'), T_divide: T.Buffer((), "float32"),): + def nll_loss(rxplaceholder: T.Tensor([N_nll_loss, C_nll_loss, d1_nll_loss, d2_nll_loss], dtype='float32'), rxplaceholder_1: T.Tensor([N_nll_loss, d1_nll_loss, d2_nll_loss], dtype='int64'), rxplaceholder_2: T.Tensor([C_nll_loss], dtype='float32'), T_divide: T.Tensor((), "float32"),): # function attr dict T.func_attr({"tirx.noalias": True}) @@ -4201,8 +4201,8 @@ def main( @Ts.prim_func(private=True) def pad( - A: T.Buffer((T.int64(2), T.int64(128), T.int64(28)), "float32"), - PadInput: T.Buffer((T.int64(2), T.int64(130), T.int64(30)), "float32"), + A: T.Tensor((T.int64(2), T.int64(128), T.int64(28)), "float32"), + PadInput: T.Tensor((T.int64(2), T.int64(130), T.int64(30)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -4241,7 +4241,7 @@ def main(x: R.Tensor((2, 3, 4, 5), dtype="float32")) -> R.Tensor((2, 60), dtype= return gv @Ts.prim_func(private=True) - def reshape(x: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), T_reshape: T.Buffer((T.int64(2), T.int64(60)), "float32")): + def reshape(x: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), T_reshape: T.Tensor((T.int64(2), T.int64(60)), "float32")): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(2), T.int64(60)): with Ts.sblock("T_reshape"): @@ -4279,7 +4279,7 @@ def main(x: R.Tensor((2, 3), "float32")) -> R.Tuple(R.Tensor((2, 3), "float32"), @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def dropout(x: T.Buffer((T.int64(2), T.int64(3)), "float32"), compute: T.Buffer((T.int64(2), T.int64(3)), "float32"), T_full_like: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def dropout(x: T.Tensor((T.int64(2), T.int64(3)), "float32"), compute: T.Tensor((T.int64(2), T.int64(3)), "float32"), T_full_like: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("compute"): diff --git a/tests/python/relax/test_transform_legalize_ops_qdq.py b/tests/python/relax/test_transform_legalize_ops_qdq.py index 86d766afc8c0..293b98d55f4d 100644 --- a/tests/python/relax/test_transform_legalize_ops_qdq.py +++ b/tests/python/relax/test_transform_legalize_ops_qdq.py @@ -39,10 +39,10 @@ def main( class Expected: @Ts.prim_func(private=True) def quantize( - A: T.Buffer((T.int64(2), T.int64(4)), "float32"), - B: T.Buffer((T.int64(2),), "float32"), - C: T.Buffer((T.int64(2),), "int8"), - quantized: T.Buffer((T.int64(2), T.int64(4)), "int8"), + A: T.Tensor((T.int64(2), T.int64(4)), "float32"), + B: T.Tensor((T.int64(2),), "float32"), + C: T.Tensor((T.int64(2),), "int8"), + quantized: T.Tensor((T.int64(2), T.int64(4)), "int8"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -93,10 +93,10 @@ def main( class Expected: @Ts.prim_func(private=True) def quantize( - A: T.Buffer((T.int64(2), T.int64(4)), "float16"), - B: T.Buffer((T.int64(2),), "float16"), - C: T.Buffer((T.int64(2),), "int8"), - quantized: T.Buffer((T.int64(2), T.int64(4)), "uint8"), + A: T.Tensor((T.int64(2), T.int64(4)), "float16"), + B: T.Tensor((T.int64(2),), "float16"), + C: T.Tensor((T.int64(2),), "int8"), + quantized: T.Tensor((T.int64(2), T.int64(4)), "uint8"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -152,10 +152,10 @@ def main( class Expected: @Ts.prim_func(private=True) def quantize( - A: T.Buffer((T.int64(4), n_quantize)), - B: T.Buffer((n_quantize,)), - C: T.Buffer((n_quantize,), "int8"), - quantized: T.Buffer((T.int64(4), n_quantize), "int8"), + A: T.Tensor((T.int64(4), n_quantize)), + B: T.Tensor((n_quantize,)), + C: T.Tensor((n_quantize,), "int8"), + quantized: T.Tensor((T.int64(4), n_quantize), "int8"), ): T.func_attr({"tirx.noalias": True}) @@ -205,8 +205,8 @@ def main(data: R.Tensor((2, 4), "float32")) -> R.Tensor((2, 4), "int8"): class Expected: @Ts.prim_func(private=True) def quantize( - A: T.Buffer((T.int64(2), T.int64(4)), "float32"), - quantized: T.Buffer((T.int64(2), T.int64(4)), "int8"), + A: T.Tensor((T.int64(2), T.int64(4)), "float32"), + quantized: T.Tensor((T.int64(2), T.int64(4)), "int8"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -253,10 +253,10 @@ def main(data: R.Tensor((2, 4), "float32")) -> R.Tensor((2, 4), "int8"): class Expected: @Ts.prim_func(private=True) def quantize( - A: T.Buffer((T.int64(2), T.int64(4)), "float32"), - B: T.Buffer((T.int64(2),), "float32"), - C: T.Buffer((T.int64(2),), "int8"), - quantized: T.Buffer((T.int64(2), T.int64(4)), "int8"), + A: T.Tensor((T.int64(2), T.int64(4)), "float32"), + B: T.Tensor((T.int64(2),), "float32"), + C: T.Tensor((T.int64(2),), "int8"), + quantized: T.Tensor((T.int64(2), T.int64(4)), "int8"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -304,8 +304,8 @@ def main(data: R.Tensor((2, 4), "float16")) -> R.Tensor((2, 4), "int8"): class Expected: @Ts.prim_func(private=True) def quantize( - A: T.Buffer((T.int64(2), T.int64(4)), "float16"), - quantized: T.Buffer((T.int64(2), T.int64(4)), "int8"), + A: T.Tensor((T.int64(2), T.int64(4)), "float16"), + quantized: T.Tensor((T.int64(2), T.int64(4)), "int8"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -350,10 +350,10 @@ def main( class Expected: @Ts.prim_func(private=True) def dequantize( - A: T.Buffer((T.int64(2), T.int64(4)), "int8"), - B: T.Buffer((T.int64(2),), "float32"), - C: T.Buffer((T.int64(2),), "int8"), - dequantized: T.Buffer((T.int64(2), T.int64(4)), "float32"), + A: T.Tensor((T.int64(2), T.int64(4)), "int8"), + B: T.Tensor((T.int64(2),), "float32"), + C: T.Tensor((T.int64(2),), "int8"), + dequantized: T.Tensor((T.int64(2), T.int64(4)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -396,8 +396,8 @@ def main(data: R.Tensor((2, 4), "int8")) -> R.Tensor((2, 4), "float32"): class Expected: @Ts.prim_func(private=True) def dequantize( - A: T.Buffer((T.int64(2), T.int64(4)), "int8"), - dequantized: T.Buffer((T.int64(2), T.int64(4)), "float32"), + A: T.Tensor((T.int64(2), T.int64(4)), "int8"), + dequantized: T.Tensor((T.int64(2), T.int64(4)), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -441,10 +441,10 @@ def main( class Expected: @Ts.prim_func(private=True) def dequantize( - A: T.Buffer((T.int64(2), n_dequantize), "int8"), - B: T.Buffer((n_dequantize,)), - C: T.Buffer((n_dequantize,), "int8"), - dequantized: T.Buffer((T.int64(2), n_dequantize)), + A: T.Tensor((T.int64(2), n_dequantize), "int8"), + B: T.Tensor((n_dequantize,)), + C: T.Tensor((n_dequantize,), "int8"), + dequantized: T.Tensor((T.int64(2), n_dequantize)), ): T.func_attr({"tirx.noalias": True}) @@ -492,10 +492,10 @@ def main( class Expected: @Ts.prim_func(private=True) def dequantize( - A: T.Buffer((T.int64(2), T.int64(4)), "int8"), - B: T.Buffer((T.int64(2),), "float16"), - C: T.Buffer((T.int64(2),), "int8"), - dequantized: T.Buffer((T.int64(2), T.int64(4)), "float16"), + A: T.Tensor((T.int64(2), T.int64(4)), "int8"), + B: T.Tensor((T.int64(2),), "float16"), + C: T.Tensor((T.int64(2),), "int8"), + dequantized: T.Tensor((T.int64(2), T.int64(4)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -548,8 +548,8 @@ def main(data: R.Tensor((2, 4), "int8")) -> R.Tensor((2, 4), "float16"): class Expected: @Ts.prim_func(private=True) def dequantize( - A: T.Buffer((T.int64(2), T.int64(4)), "int8"), - dequantized: T.Buffer((T.int64(2), T.int64(4)), "float16"), + A: T.Tensor((T.int64(2), T.int64(4)), "int8"), + dequantized: T.Tensor((T.int64(2), T.int64(4)), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): 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 5443eb63b090..980e28a79f73 100644 --- a/tests/python/relax/test_transform_legalize_ops_search_statistical.py +++ b/tests/python/relax/test_transform_legalize_ops_search_statistical.py @@ -44,7 +44,7 @@ def main(condition: R.Tensor((3, 2, 1), "bool"), x: R.Tensor((2, 3), "float32"), return gv @Ts.prim_func(private=True) - def where(rxplaceholder: T.Buffer((T.int64(3), T.int64(2), T.int64(1)), "bool"), rxplaceholder_1: T.Buffer((T.int64(2), T.int64(3)), "float32"), rxplaceholder_2: T.Buffer((T.int64(2), T.int64(1)), "float32"), T_where: T.Buffer((T.int64(3), T.int64(2), T.int64(3)), "float32")): + def where(rxplaceholder: T.Tensor((T.int64(3), T.int64(2), T.int64(1)), "bool"), rxplaceholder_1: T.Tensor((T.int64(2), T.int64(3)), "float32"), rxplaceholder_2: T.Tensor((T.int64(2), T.int64(1)), "float32"), T_where: T.Tensor((T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_where"): @@ -86,7 +86,7 @@ def main(condition: R.Tensor((a_main, b_main, 1), "bool"), x: R.Tensor((b_main, return gv @Ts.prim_func(private=True) - def where(rxplaceholder: T.Buffer([a_where, b_where, T.int64(1)], dtype='bool'), rxplaceholder_1: T.Buffer([b_where, c_where], dtype='float32'), rxplaceholder_2: T.Buffer([b_where, T.int64(1)], dtype='float32'), T_where: T.Buffer([a_where, b_where, c_where], dtype='float32')): + def where(rxplaceholder: T.Tensor([a_where, b_where, T.int64(1)], dtype='bool'), rxplaceholder_1: T.Tensor([b_where, c_where], dtype='float32'), rxplaceholder_2: T.Tensor([b_where, T.int64(1)], dtype='float32'), T_where: T.Tensor([a_where, b_where, c_where], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2 in T.grid(a_where, b_where, c_where): @@ -118,7 +118,7 @@ def main(x: R.Tensor((2, 3, 4, 5), dtype="float32")) -> R.Tensor((2, 4, 5), dtyp return gv @Ts.prim_func(private=True) - def argmax(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Buffer((T.int64(2), T.int64(4), T.int64(5)), "int64")): + def argmax(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Tensor((T.int64(2), T.int64(4), T.int64(5)), "int64")): T.func_attr({"tirx.noalias": True}) rxplaceholder_red_temp_v0 = Ts.sblock_alloc_buffer((T.int64(2), T.int64(4), T.int64(5)), "int64") rxplaceholder_red_temp_v1 = Ts.sblock_alloc_buffer((T.int64(2), T.int64(4), T.int64(5))) @@ -177,7 +177,7 @@ def main(x: R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) -> R.Te return gv @Ts.prim_func(private=True) - def argmax(rxplaceholder: T.Buffer((a_argmax, b_argmax, c_argmax, d_argmax)), rxplaceholder_red: T.Buffer((a_argmax, T.int64(1), c_argmax, d_argmax), 'int64')): + def argmax(rxplaceholder: T.Tensor((a_argmax, b_argmax, c_argmax, d_argmax)), rxplaceholder_red: T.Tensor((a_argmax, T.int64(1), c_argmax, d_argmax), 'int64')): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -219,7 +219,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32")) -> R.Tensor((), "int64"): @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def argmin(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Buffer((), "int64")): + def argmin(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Tensor((), "int64")): T.func_attr({"tirx.noalias": True}) rxplaceholder_red_temp_v0 = Ts.sblock_alloc_buffer((), "int64") rxplaceholder_red_temp_v1 = Ts.sblock_alloc_buffer(()) @@ -277,7 +277,7 @@ def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((1, 1, 1, 1), "int64" @tvm.script.ir_module class Expected: @Ts.prim_func(private=True) - def argmin(rxplaceholder: T.Buffer((a_argmin, b_argmin, c_argmin, d_argmin)), rxplaceholder_red: T.Buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "int64")): + def argmin(rxplaceholder: T.Tensor((a_argmin, b_argmin, c_argmin, d_argmin)), rxplaceholder_red: T.Tensor((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "int64")): T.func_attr({"tirx.noalias": True}) rxplaceholder_red_temp_v0 = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "int64") @@ -331,7 +331,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32")) -> R.Tensor((2, 5), "float32"): return gv @Ts.prim_func(private=True) - def max(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Buffer((T.int64(2), T.int64(5)), "float32")): + def max(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Tensor((T.int64(2), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(5), T.int64(3), T.int64(4)): with Ts.sblock("rxplaceholder_red"): @@ -378,7 +378,7 @@ def main(x: R.Tensor((a_main, b_main, c_main, d_main), "float32")) -> R.Tensor(( return gv @Ts.prim_func(private=True) - def max(rxplaceholder: T.Buffer([a_max, b_max, c_max, d_max], dtype='float32'), rxplaceholder_red: T.Buffer([a_max, d_max], dtype='float32')): + def max(rxplaceholder: T.Tensor([a_max, b_max, c_max, d_max], dtype='float32'), rxplaceholder_red: T.Tensor([a_max, d_max], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_max, d_max, b_max, c_max): @@ -412,7 +412,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32")) -> R.Tensor((2, 1, 1, 5), "float3 return gv @Ts.prim_func(private=True) - def min(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Buffer((T.int64(2), T.int64(1), T.int64(1), T.int64(5)), "float32")): + def min(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Tensor((T.int64(2), T.int64(1), T.int64(1), T.int64(5)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3, i4, i5 in T.grid(T.int64(2), T.int64(1), T.int64(1), T.int64(5), T.int64(3), T.int64(4)): with Ts.sblock("rxplaceholder_red"): @@ -459,7 +459,7 @@ def main(x: R.Tensor((a_main, b_main, c_main, d_main), "float32")) -> R.Tensor(( return gv @Ts.prim_func(private=True) - def min(rxplaceholder: T.Buffer([a_min, b_min, c_min, d_min], dtype='float32'), rxplaceholder_red: T.Buffer([a_min, T.int64(1), T.int64(1), d_min], dtype='float32')): + def min(rxplaceholder: T.Tensor([a_min, b_min, c_min, d_min], dtype='float32'), rxplaceholder_red: T.Tensor([a_min, T.int64(1), T.int64(1), d_min], dtype='float32')): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3, i4, i5 in T.grid(a_min, T.int64(1), T.int64(1), d_min, b_min, c_min): @@ -493,7 +493,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32")) -> R.Tensor((), "float32"): return gv @Ts.prim_func(private=True) - def sum(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Buffer((), "float32")): + def sum(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Tensor((), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(2), T.int64(3), T.int64(4), T.int64(5)): with Ts.sblock("rxplaceholder_red"): @@ -540,7 +540,7 @@ def main(x: R.Tensor((a_main, b_main, c_main, d_main), "float32")) -> R.Tensor(( return gv @Ts.prim_func(private=True) - def sum(rxplaceholder: T.Buffer([a_sum, b_sum, c_sum, d_sum], dtype='float32'), rxplaceholder_red: T.Buffer((), "float32")): + def sum(rxplaceholder: T.Tensor([a_sum, b_sum, c_sum, d_sum], dtype='float32'), rxplaceholder_red: T.Tensor((), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(a_sum, b_sum, c_sum, d_sum): @@ -574,7 +574,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32")) -> R.Tensor((1, 1, 1, 1), "float3 return gv @Ts.prim_func(private=True) - def prod(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "float32")): + def prod(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), rxplaceholder_red: T.Tensor((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3, i4, i5, i6, i7 in T.grid(T.int64(1), T.int64(1), T.int64(1), T.int64(1), T.int64(2), T.int64(3), T.int64(4), T.int64(5)): with Ts.sblock("rxplaceholder_red"): @@ -607,7 +607,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "bool")) -> R.Tensor((1, 1, 1, 1), "bool"): return gv @Ts.prim_func(private=True) - def prod(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "bool"), rxplaceholder_red: T.Buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "bool")): + def prod(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "bool"), rxplaceholder_red: T.Tensor((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "bool")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3, i4, i5, i6, i7 in T.grid(T.int64(1), T.int64(1), T.int64(1), T.int64(1), T.int64(2), T.int64(3), T.int64(4), T.int64(5)): with Ts.sblock("rxplaceholder_red"): @@ -654,7 +654,7 @@ def main(x: R.Tensor((a_main, b_main, c_main, d_main), "float32")) -> R.Tensor(( return gv @Ts.prim_func(private=True) - def prod(rxplaceholder: T.Buffer([a_prod, b_prod, c_prod, d_prod], dtype='float32'), rxplaceholder_red: T.Buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "float32")): + def prod(rxplaceholder: T.Tensor([a_prod, b_prod, c_prod, d_prod], dtype='float32'), rxplaceholder_red: T.Tensor((T.int64(1), T.int64(1), T.int64(1), T.int64(1)), "float32")): T.func_attr({"tirx.noalias": True}) 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): @@ -752,7 +752,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32")) -> R.Tensor((3, 4), "float32"): return gv @Ts.prim_func(private=True) - def mean(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), T_divide: T.Buffer((T.int64(3), T.int64(4)), "float32")): + def mean(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), T_divide: T.Tensor((T.int64(3), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) rxplaceholder_red = Ts.sblock_alloc_buffer([T.int64(3), T.int64(4)], dtype="float32") for i0, i1, i2, i3 in T.grid(T.int64(3), T.int64(4), T.int64(2), T.int64(5)): @@ -806,7 +806,7 @@ def main(x: R.Tensor((a_main, b_main, c_main, d_main), dtype="float32")) -> R.Te return gv @Ts.prim_func(private=True) - def mean(rxplaceholder: T.Buffer([a_mean, b_mean, c_mean, d_mean], dtype='float32'), T_divide: T.Buffer([b_mean, c_mean], dtype='float32')): + def mean(rxplaceholder: T.Tensor([a_mean, b_mean, c_mean, d_mean], dtype='float32'), T_divide: T.Tensor([b_mean, c_mean], dtype='float32')): T.func_attr({"tirx.noalias": True}) rxplaceholder_red = Ts.sblock_alloc_buffer([b_mean, c_mean], dtype="float32") @@ -847,7 +847,7 @@ def main(x: R.Tensor((2, 3, 4, 5), dtype="float32")) -> R.Tuple(R.Tensor((3, 4, return gv @Ts.prim_func(private=True) - def median(data_buf: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), align=8), T_squeeze: T.Buffer((T.int64(3), T.int64(4), T.int64(5)), "float32"), T_squeeze_1: T.Buffer((T.int64(3), T.int64(4), T.int64(5)), "int64")): + def median(data_buf: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), align=8), T_squeeze: T.Tensor((T.int64(3), T.int64(4), T.int64(5)), "float32"), T_squeeze_1: T.Tensor((T.int64(3), T.int64(4), T.int64(5)), "int64")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -920,7 +920,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32")) -> R.Tensor((), "float32"): @I.ir_module class Expected: @Ts.prim_func(private=True) - def std(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), compute: T.Buffer((), "float32")): + def std(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), compute: T.Tensor((), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): rxplaceholder_red = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(1), T.int64(1))) @@ -1011,7 +1011,7 @@ def main(x: R.Tensor((a, b, c, d), "float32")) -> R.Tensor((), "float32"): @I.ir_module class Expected: @Ts.prim_func(private=True) - def std(rxplaceholder: T.Buffer((a_std, b_std, c_std, d_std)), compute: T.Buffer((), "float32")): + def std(rxplaceholder: T.Tensor((a_std, b_std, c_std, d_std)), compute: T.Tensor((), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1094,7 +1094,7 @@ def main(x: R.Tensor((2, 3, 4, 5), dtype="float32")) -> R.Tensor((1, 3, 4, 1), d return gv @Ts.prim_func(private=True) - def variance(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), T_divide: T.Buffer((T.int64(1), T.int64(3), T.int64(4), T.int64(1)), "float32")): + def variance(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), T_divide: T.Tensor((T.int64(1), T.int64(3), T.int64(4), T.int64(1)), "float32")): T.func_attr({"tirx.noalias": True}) rxplaceholder_red = Ts.sblock_alloc_buffer([T.int64(1), T.int64(3), T.int64(4), T.int64(1)], dtype="float32") T_divide_1 = Ts.sblock_alloc_buffer([T.int64(1), T.int64(3), T.int64(4), T.int64(1)], dtype="float32") @@ -1178,7 +1178,7 @@ def main(x: R.Tensor((a_main, b_main, c_main, d_main), "float32")) -> R.Tensor(( return gv @Ts.prim_func(private=True) - def variance(rxplaceholder: T.Buffer([a_variance, b_variance, c_variance, d_variance], dtype='float32'), T_divide: T.Buffer([T.int64(1), b_variance, c_variance, T.int64(1)], dtype='float32')): + def variance(rxplaceholder: T.Tensor([a_variance, b_variance, c_variance, d_variance], dtype='float32'), T_divide: T.Tensor([T.int64(1), b_variance, c_variance, T.int64(1)], dtype='float32')): T.func_attr({"tirx.noalias": True}) rxplaceholder_red = Ts.sblock_alloc_buffer([T.int64(1), b_variance, c_variance, T.int64(1)], dtype="float32") @@ -1244,7 +1244,7 @@ def main(x: R.Tensor((2, 3, 4, 5), "float32")) -> R.Tensor((3, 4), "float32"): @I.ir_module class Expected: @Ts.prim_func(private=True) - def variance(rxplaceholder: T.Buffer((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), T_divide: T.Buffer((T.int64(3), T.int64(4)), "float32")): + def variance(rxplaceholder: T.Tensor((T.int64(2), T.int64(3), T.int64(4), T.int64(5)), "float32"), T_divide: T.Tensor((T.int64(3), T.int64(4)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): rxplaceholder_red = Ts.sblock_alloc_buffer((T.int64(1), T.int64(3), T.int64(4), T.int64(1))) @@ -1323,7 +1323,7 @@ def main(x: R.Tensor((), dtype="float32")) -> R.Tensor((), dtype="float32"): return gv @Ts.prim_func(private=True) - def max(x: T.Buffer((), "float32"), x_red: T.Buffer((), "float32")): + def max(x: T.Tensor((), "float32"), x_red: T.Tensor((), "float32")): T.func_attr({"tirx.noalias": True}) with Ts.sblock("x_red"): vi = Ts.axis.spatial(T.int64(1), T.int64(0)) diff --git a/tests/python/relax/test_transform_lift_transform_params.py b/tests/python/relax/test_transform_lift_transform_params.py index cdba0dad7bdb..782fc75fd5ef 100644 --- a/tests/python/relax/test_transform_lift_transform_params.py +++ b/tests/python/relax/test_transform_lift_transform_params.py @@ -40,7 +40,7 @@ def test_basic(consume_params): class Before: @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((3, 16, 3, 3), "float32"), out: T.Buffer((16, 3, 3, 3), "float32") + w1: T.Tensor((3, 16, 3, 3), "float32"), out: T.Tensor((16, 3, 3, 3), "float32") ) -> None: for ax0, ax1, ax2, ax3 in T.grid(16, 3, 3, 3): with Ts.sblock("layout_transform"): @@ -105,7 +105,7 @@ def main( @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((3, 16, 3, 3), "float32"), out: T.Buffer((16, 3, 3, 3), "float32") + w1: T.Tensor((3, 16, 3, 3), "float32"), out: T.Tensor((16, 3, 3, 3), "float32") ): for ax0, ax1, ax2, ax3 in T.grid(16, 3, 3, 3): with Ts.sblock("layout_transform"): @@ -176,7 +176,7 @@ def main( @Ts.prim_func def transform_layout_IOHW_to_OIHW( - w1: T.Buffer((3, 16, 3, 3), "float32"), out: T.Buffer((16, 3, 3, 3), "float32") + w1: T.Tensor((3, 16, 3, 3), "float32"), out: T.Tensor((16, 3, 3, 3), "float32") ): for ax0, ax1, ax2, ax3 in T.grid(16, 3, 3, 3): with Ts.sblock("layout_transform"): @@ -1442,7 +1442,7 @@ def test_symbolic_var_2(): @I.ir_module class Before: @Ts.prim_func - def zeros(T_full: T.Buffer((n_zeros, n_zeros))): + def zeros(T_full: T.Tensor((n_zeros, n_zeros))): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(n_zeros, n_zeros): @@ -1469,7 +1469,7 @@ def main(shape: R.Shape([n_main])) -> R.Shape([n_main]): @I.ir_module class Expected: @Ts.prim_func - def zeros(T_full: T.Buffer((n_zeros, n_zeros))): + def zeros(T_full: T.Tensor((n_zeros, n_zeros))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1531,9 +1531,9 @@ def main( @Ts.prim_func(private=True) def slice( - Input_2d: T.Buffer(shape=[16, 16], dtype="int32"), + Input_2d: T.Tensor(shape=[16, 16], dtype="int32"), slice_index: T.int64, - Output_Slice: T.Buffer(shape=[16], dtype="int32"), + Output_Slice: T.Tensor(shape=[16], dtype="int32"), ): T.func_attr({"tirx.noalias": True}) for j in range(16): @@ -1586,9 +1586,9 @@ def main_transform_params( @Ts.prim_func(private=True) def slice( - Input_2d: T.Buffer(shape=[16, 16], dtype="int32"), + Input_2d: T.Tensor(shape=[16, 16], dtype="int32"), slice_index: T.int64, - Output_Slice: T.Buffer(shape=[16], dtype="int32"), + Output_Slice: T.Tensor(shape=[16], dtype="int32"), ): T.func_attr({"tirx.noalias": True}) for j in range(16): diff --git a/tests/python/relax/test_transform_merge_composite_functions.py b/tests/python/relax/test_transform_merge_composite_functions.py index 5fa0f5a63079..d1fb9f8b0eb3 100644 --- a/tests/python/relax/test_transform_merge_composite_functions.py +++ b/tests/python/relax/test_transform_merge_composite_functions.py @@ -1148,8 +1148,8 @@ def fused_relax_nn_relu( @Ts.prim_func(private=True) def relu( - Input: T.Buffer(T.int64(10), "float32"), - Output: T.Buffer(T.int64(10), "float32"), + Input: T.Tensor(T.int64(10), "float32"), + Output: T.Tensor(T.int64(10), "float32"), ): T.func_attr({"tirx.noalias": True}) for i in range(T.int64(10)): @@ -1200,8 +1200,8 @@ def composite_lambda( @Ts.prim_func(private=True) def relu( - Input: T.Buffer(T.int64(10), "float32"), - Output: T.Buffer(T.int64(10), "float32"), + Input: T.Tensor(T.int64(10), "float32"), + Output: T.Tensor(T.int64(10), "float32"), ): T.func_attr({"tirx.noalias": True}) for i in range(T.int64(10)): diff --git a/tests/python/relax/test_transform_meta_schedule_apply_database.py b/tests/python/relax/test_transform_meta_schedule_apply_database.py index 779eb6d00e52..9f881f306728 100644 --- a/tests/python/relax/test_transform_meta_schedule_apply_database.py +++ b/tests/python/relax/test_transform_meta_schedule_apply_database.py @@ -31,7 +31,7 @@ def test_apply_to_func_with_different_block_name(): @I.ir_module class RecordModule: @Ts.prim_func - def main(A: T.Buffer((2,), "float32"), B: T.Buffer((2,), "float32")): + def main(A: T.Tensor((2,), "float32"), B: T.Tensor((2,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i in T.serial(2): with Ts.sblock("block"): @@ -41,7 +41,7 @@ def main(A: T.Buffer((2,), "float32"), B: T.Buffer((2,), "float32")): @I.ir_module class BlockRenamedModule: @Ts.prim_func - def main(A: T.Buffer((2,), "float32"), B: T.Buffer((2,), "float32")): + def main(A: T.Tensor((2,), "float32"), B: T.Tensor((2,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i in T.serial(2): with Ts.sblock("renamed_block"): @@ -51,7 +51,7 @@ def main(A: T.Buffer((2,), "float32"), B: T.Buffer((2,), "float32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((2,), "float32"), B: T.Buffer((2,), "float32")): + def main(A: T.Tensor((2,), "float32"), B: T.Tensor((2,), "float32")): T.func_attr( { "tirx.is_scheduled": True, diff --git a/tests/python/relax/test_transform_meta_schedule_tuning.py b/tests/python/relax/test_transform_meta_schedule_tuning.py index 51cb63eed48d..f3d9643f30f4 100644 --- a/tests/python/relax/test_transform_meta_schedule_tuning.py +++ b/tests/python/relax/test_transform_meta_schedule_tuning.py @@ -53,7 +53,7 @@ @tvm.script.ir_module class InputModule: @Ts.prim_func - def tir_matmul(A: T.Buffer((32, 32)), B: T.Buffer((32, 32)), C: T.Buffer((32, 32))) -> None: + def tir_matmul(A: T.Tensor((32, 32)), B: T.Tensor((32, 32)), C: T.Tensor((32, 32))) -> None: T.func_attr({"global_symbol": "tir_matmul"}) for i0, j0, k0 in T.grid(32, 32, 32): @@ -64,7 +64,7 @@ def tir_matmul(A: T.Buffer((32, 32)), B: T.Buffer((32, 32)), C: T.Buffer((32, 32 C[i, j] += A[i, k] * B[j, k] @Ts.prim_func - def tir_relu(A: T.Buffer((32, 32)), B: T.Buffer((32, 32))): + def tir_relu(A: T.Tensor((32, 32)), B: T.Tensor((32, 32))): T.func_attr({"global_symbol": "tir_relu"}) for i, j in T.grid(32, 32): @@ -166,9 +166,9 @@ def test_ms_tuning_primfunc(): class DefaultScheduledModule: @Ts.prim_func def tir_matmul( - A: T.Buffer((32, 32), "float32"), - B: T.Buffer((32, 32), "float32"), - C: T.Buffer((32, 32), "float32"), + A: T.Tensor((32, 32), "float32"), + B: T.Tensor((32, 32), "float32"), + C: T.Tensor((32, 32), "float32"), ): T.func_attr({"global_symbol": "tir_matmul", "tirx.is_scheduled": True}) # with Ts.sblock("root"): @@ -186,7 +186,7 @@ def tir_matmul( C[i, j] = C[i, j] + A[i, k] * B[j, k] @Ts.prim_func - def tir_relu(A: T.Buffer((32, 32), "float32"), B: T.Buffer((32, 32), "float32")): + def tir_relu(A: T.Tensor((32, 32), "float32"), B: T.Tensor((32, 32), "float32")): T.func_attr({"global_symbol": "tir_relu", "tirx.is_scheduled": True}) # with Ts.sblock("root"): for i_j_fused_0 in T.thread_binding(1, thread="blockIdx.x"): diff --git a/tests/python/relax/test_transform_normalize_global_var.py b/tests/python/relax/test_transform_normalize_global_var.py index cd5694e03605..48a19f83f8f3 100644 --- a/tests/python/relax/test_transform_normalize_global_var.py +++ b/tests/python/relax/test_transform_normalize_global_var.py @@ -67,7 +67,7 @@ def test_normalize_tir_function(): @I.ir_module(check_well_formed=False) class Before: @Ts.prim_func(private=True) - def f(x: T.Buffer((1,), "int32")): + def f(x: T.Tensor((1,), "int32")): x[0] = T.int32(0) @R.function @@ -80,7 +80,7 @@ def f1(): @I.ir_module class Expected: @Ts.prim_func(private=True) - def f1(x: T.Buffer((1,), "int32")): + def f1(x: T.Tensor((1,), "int32")): x[0] = 0 @R.function diff --git a/tests/python/relax/test_transform_operator_specific_normalization.py b/tests/python/relax/test_transform_operator_specific_normalization.py index 5c6b142a12be..75c901de058a 100644 --- a/tests/python/relax/test_transform_operator_specific_normalization.py +++ b/tests/python/relax/test_transform_operator_specific_normalization.py @@ -189,7 +189,7 @@ def main(A: R.Tensor([16], "float32")): ) @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): + def multiply_by_two(A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32")): for i in range(16): B[i] = A[i] * 2.0 @@ -206,7 +206,7 @@ def main(A: R.Tensor([16], "float32")): ) @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): + def multiply_by_two(A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32")): for i in range(16): B[i] = A[i] * 2.0 @@ -236,7 +236,7 @@ def main(args: R.Tuple([R.Tensor([16], "float32")])): ) @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): + def multiply_by_two(A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32")): for i in range(16): B[i] = A[i] * 2.0 @@ -252,7 +252,7 @@ def main(args: R.Tuple([R.Tensor([16], "float32")])): ) @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): + def multiply_by_two(A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32")): for i in range(16): B[i] = A[i] * 2.0 @@ -283,7 +283,7 @@ def main(A: R.Tensor([16], "float32")): ) @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer(16, "float32")): + def multiply_by_two(A: T.Tensor(16, "float32")): for i in range(16): A[i] = A[i] * 2.0 @@ -303,7 +303,7 @@ def main(A: R.Tensor([16], "float32")): ) @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer(16, "float32")): + def multiply_by_two(A: T.Tensor(16, "float32")): for i in range(16): A[i] = A[i] * 2.0 @@ -334,13 +334,13 @@ def main(A: R.Tensor([16], "float32")): ) @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): + def multiply_by_two(A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32")): for i in range(16): B[i] = A[i] * 2.0 @Ts.prim_func(private=True) def f_grad( - A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32"), Grad: T.Buffer(16, "float32") + A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32"), Grad: T.Tensor(16, "float32") ): for i in range(16): Grad[i] = 2.0 @@ -361,13 +361,13 @@ def main(A: R.Tensor([16], "float32")): ) @Ts.prim_func(private=True) - def multiply_by_two(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): + def multiply_by_two(A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32")): for i in range(16): B[i] = A[i] * 2.0 @Ts.prim_func(private=True) def f_grad( - A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32"), Grad: T.Buffer(16, "float32") + A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32"), Grad: T.Tensor(16, "float32") ): for i in range(16): Grad[i] = 2.0 diff --git a/tests/python/relax/test_transform_rewrite_cuda_graph.py b/tests/python/relax/test_transform_rewrite_cuda_graph.py index 1e160e4040a4..af6f3f3fa5db 100644 --- a/tests/python/relax/test_transform_rewrite_cuda_graph.py +++ b/tests/python/relax/test_transform_rewrite_cuda_graph.py @@ -40,7 +40,7 @@ def test_rewrite_cuda_graph(): @I.ir_module class Before: @Ts.prim_func - def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def exp(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4)), "float32")): # function attr dict T.func_attr({"tirx.noalias": True, "global_symbol": "exp"}) for i0_i1_fused_0 in T.thread_binding(T.int64(1), thread="blockIdx.x"): @@ -80,7 +80,7 @@ def main(x: R.Tensor((2, 4), dtype="float32")) -> R.Tensor((2,4), dtype="float32 @I.ir_module class Expected: @Ts.prim_func - def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def exp(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4)), "float32")): # function attr dict T.func_attr({"tirx.noalias": True, "global_symbol": "exp"}) # body @@ -150,7 +150,7 @@ def test_tuple(): @I.ir_module class Before: @Ts.prim_func - def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def exp(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4)), "float32")): # function attr dict T.func_attr({"tirx.noalias": True, "global_symbol": "exp"}) # body @@ -192,7 +192,7 @@ def main(x: R.Tensor((2, 4), dtype="float32")) -> R.Tensor((2, 4), dtype="float3 @I.ir_module class Expected: @Ts.prim_func - def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def exp(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4)), "float32")): T.func_attr({"global_symbol": "exp", "tirx.noalias": True}) # with Ts.sblock("root"): for i0_i1_fused_0 in T.thread_binding(T.int64(1), thread="blockIdx.x"): @@ -257,7 +257,7 @@ def test_vm_builtin(): @I.ir_module class Before: @Ts.prim_func - def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def exp(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4)), "float32")): # function attr dict T.func_attr({"tirx.noalias": True, "global_symbol": "exp"}) for i0_i1_fused_0 in T.thread_binding(T.int64(1), thread="blockIdx.x"): @@ -292,7 +292,7 @@ def main(x: R.Tensor((2, 4), dtype="float32")) -> R.Tensor((2,4), dtype="float32 @I.ir_module class Expected: @Ts.prim_func - def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def exp(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4)), "float32")): T.func_attr({"global_symbol": "exp", "tirx.noalias": True}) # with Ts.sblock("root"): for i0_i1_fused_0 in T.thread_binding(T.int64(1), thread="blockIdx.x"): @@ -392,9 +392,9 @@ def main( class Expected: @Ts.prim_func def fused_conv2d_relu( - data: T.Buffer((T.int64(16), T.int64(32), T.int64(32), T.int64(16)), "float16"), - weight1: T.Buffer((T.int64(16), T.int64(3), T.int64(3), T.int64(16)), "float16"), - var_compute_intermediate: T.Buffer( + data: T.Tensor((T.int64(16), T.int64(32), T.int64(32), T.int64(16)), "float16"), + weight1: T.Tensor((T.int64(16), T.int64(3), T.int64(3), T.int64(16)), "float16"), + var_compute_intermediate: T.Tensor( (T.int64(16), T.int64(32), T.int64(32), T.int64(16)), "float16" ), ): @@ -455,10 +455,10 @@ def fused_conv2d_relu( @Ts.prim_func def layer_norm( - A: T.Buffer((T.int64(16), T.int64(32), T.int64(32), T.int64(16)), "float16"), - B: T.Buffer((T.int64(16),), "float16"), - C: T.Buffer((T.int64(16),), "float16"), - T_layer_norm: T.Buffer((T.int64(16), T.int64(32), T.int64(32), T.int64(16)), "float16"), + A: T.Tensor((T.int64(16), T.int64(32), T.int64(32), T.int64(16)), "float16"), + B: T.Tensor((T.int64(16),), "float16"), + C: T.Tensor((T.int64(16),), "float16"), + T_layer_norm: T.Tensor((T.int64(16), T.int64(32), T.int64(32), T.int64(16)), "float16"), ): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -765,7 +765,7 @@ def test_dynamic_capture(): @I.ir_module class Before: @Ts.prim_func - def add_one(x: T.Buffer((m_add_one,), "float32"), y: T.Buffer((m_add_one,), "float32")): + def add_one(x: T.Tensor((m_add_one,), "float32"), y: T.Tensor((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_add_one): @@ -803,7 +803,7 @@ def main(x: R.Tensor((m_main,), "float32")) -> R.Tensor((m_main,), "float32"): @I.ir_module class Expected: @Ts.prim_func - def add_one(x: T.Buffer((m_add_one,)), y: T.Buffer((m_add_one,))): + def add_one(x: T.Tensor((m_add_one,)), y: T.Tensor((m_add_one,))): # with Ts.sblock("root"): for i in T.serial(T.int64(0), m_add_one): with Ts.sblock("add"): diff --git a/tests/python/relax/test_transform_rewrite_dataflow_reshape.py b/tests/python/relax/test_transform_rewrite_dataflow_reshape.py index 971f0822e7c6..9dceb0413b02 100644 --- a/tests/python/relax/test_transform_rewrite_dataflow_reshape.py +++ b/tests/python/relax/test_transform_rewrite_dataflow_reshape.py @@ -30,8 +30,8 @@ def test_reshape_expand_dims(): class Module: @Ts.prim_func def reshape( - rxplaceholder: T.Buffer((T.int64(8), T.int64(3)), "float32"), - T_reshape: T.Buffer((T.int64(2), T.int64(4), T.int64(3)), "float32"), + rxplaceholder: T.Tensor((T.int64(8), T.int64(3)), "float32"), + T_reshape: T.Tensor((T.int64(2), T.int64(4), T.int64(3)), "float32"), ): for ax0, ax1, ax2 in T.grid(T.int64(2), T.int64(4), T.int64(3)): with Ts.sblock("T_reshape"): @@ -50,8 +50,8 @@ def reshape( @Ts.prim_func def expand_dims( - rxplaceholder: T.Buffer((T.int64(2), T.int64(4), T.int64(3)), "float32"), - expand_dims: T.Buffer( + rxplaceholder: T.Tensor((T.int64(2), T.int64(4), T.int64(3)), "float32"), + expand_dims: T.Tensor( (T.int64(2), T.int64(1), T.int64(4), T.int64(1), T.int64(3)), "float32" ), ): @@ -79,8 +79,8 @@ def main(x: R.Tensor((8, 3), dtype="float32")) -> R.Tensor( class Expected: @Ts.prim_func def reshape( - rxplaceholder: T.Buffer((T.int64(8), T.int64(3)), "float32"), - T_reshape: T.Buffer((T.int64(2), T.int64(4), T.int64(3)), "float32"), + rxplaceholder: T.Tensor((T.int64(8), T.int64(3)), "float32"), + T_reshape: T.Tensor((T.int64(2), T.int64(4), T.int64(3)), "float32"), ): for ax0, ax1, ax2 in T.grid(T.int64(2), T.int64(4), T.int64(3)): with Ts.sblock("T_reshape"): @@ -99,8 +99,8 @@ def reshape( @Ts.prim_func def expand_dims( - rxplaceholder: T.Buffer((T.int64(2), T.int64(4), T.int64(3)), "float32"), - expand_dims: T.Buffer( + rxplaceholder: T.Tensor((T.int64(2), T.int64(4), T.int64(3)), "float32"), + expand_dims: T.Tensor( (T.int64(2), T.int64(1), T.int64(4), T.int64(1), T.int64(3)), "float32" ), ): @@ -138,7 +138,7 @@ def test_reshape_pattern_detect(): @tvm.script.ir_module class Module: @Ts.prim_func - def reshape(rxplaceholder: T.Buffer((T.int64(2), T.int64(4096), T.int64(320)), "float32"), T_reshape: T.Buffer((T.int64(2), T.int64(4096), T.int64(5), T.int64(64)), "float32")): + def reshape(rxplaceholder: T.Tensor((T.int64(2), T.int64(4096), T.int64(320)), "float32"), T_reshape: T.Tensor((T.int64(2), T.int64(4096), T.int64(5), T.int64(64)), "float32")): for ax0_ax1_ax2_ax3_fused_1 in T.thread_binding(T.int64(256), thread="blockIdx.x"): for ax0_ax1_ax2_ax3_fused_2 in T.thread_binding(T.int64(1024), thread="threadIdx.x"): for ax0_ax1_ax2_ax3_fused_0 in range(T.int64(10)): @@ -153,8 +153,8 @@ def reshape(rxplaceholder: T.Buffer((T.int64(2), T.int64(4096), T.int64(320)), " @Ts.prim_func def expand_dims( - rxplaceholder: T.Buffer((T.int64(2), T.int64(4096), T.int64(5), T.int64(64)), "float32"), - expand_dims: T.Buffer( + rxplaceholder: T.Tensor((T.int64(2), T.int64(4096), T.int64(5), T.int64(64)), "float32"), + expand_dims: T.Tensor( (T.int64(2), T.int64(1), T.int64(4096), T.int64(1), T.int64(5), T.int64(64)), "float32", ), @@ -184,7 +184,7 @@ def main( @tvm.script.ir_module class Expected: @Ts.prim_func - def expand_dims(rxplaceholder: T.Buffer((T.int64(2), T.int64(4096), T.int64(5), T.int64(64)), "float32"), expand_dims_1: T.Buffer((T.int64(2), T.int64(1), T.int64(4096), T.int64(1), T.int64(5), T.int64(64)), "float32")): + def expand_dims(rxplaceholder: T.Tensor((T.int64(2), T.int64(4096), T.int64(5), T.int64(64)), "float32"), expand_dims_1: T.Tensor((T.int64(2), T.int64(1), T.int64(4096), T.int64(1), T.int64(5), T.int64(64)), "float32")): # with Ts.sblock("root"): for i0, i1, i2, i3, i4, i5 in T.grid(T.int64(2), T.int64(1), T.int64(4096), T.int64(1), T.int64(5), T.int64(64)): with Ts.sblock("expand_dims"): @@ -194,7 +194,7 @@ def expand_dims(rxplaceholder: T.Buffer((T.int64(2), T.int64(4096), T.int64(5), expand_dims_1[i0_1, i1_1, i2_1, i3_1, i4_1, i5_1] = rxplaceholder[i0_1, i2_1, i4_1, i5_1] @Ts.prim_func - def reshape(rxplaceholder: T.Buffer((T.int64(2), T.int64(4096), T.int64(320)), "float32"), T_reshape: T.Buffer((T.int64(2), T.int64(4096), T.int64(5), T.int64(64)), "float32")): + def reshape(rxplaceholder: T.Tensor((T.int64(2), T.int64(4096), T.int64(320)), "float32"), T_reshape: T.Tensor((T.int64(2), T.int64(4096), T.int64(5), T.int64(64)), "float32")): # with Ts.sblock("root"): for ax0_ax1_ax2_ax3_fused_1 in T.thread_binding(T.int64(256), thread="blockIdx.x"): for ax0_ax1_ax2_ax3_fused_2 in T.thread_binding(T.int64(1024), thread="threadIdx.x"): @@ -230,7 +230,7 @@ def test_reshape_dynamic_shape(): class Module: @Ts.prim_func(private=True) def reshape( - A: T.Buffer((n, 16, 128), "float16"), T_reshape: T.Buffer((1, n, 16, 128), "float16") + A: T.Tensor((n, 16, 128), "float16"), T_reshape: T.Tensor((1, n, 16, 128), "float16") ): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) @@ -272,7 +272,7 @@ def main(x: R.Tensor((8, 16, 128), dtype="float16")) -> R.Tensor( class Expected: @Ts.prim_func(private=True) def reshape( - A: T.Buffer((n, 16, 128), "float16"), T_reshape: T.Buffer((1, n, 16, 128), "float16") + A: T.Tensor((n, 16, 128), "float16"), T_reshape: T.Tensor((1, n, 16, 128), "float16") ): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) @@ -319,8 +319,8 @@ def test_reshape_non_dataflow(): class Module: @Ts.prim_func def reshape( - rxplaceholder: T.Buffer((T.int64(8), T.int64(3)), "float32"), - T_reshape: T.Buffer((T.int64(2), T.int64(4), T.int64(3)), "float32"), + rxplaceholder: T.Tensor((T.int64(8), T.int64(3)), "float32"), + T_reshape: T.Tensor((T.int64(2), T.int64(4), T.int64(3)), "float32"), ): for ax0, ax1, ax2 in T.grid(T.int64(2), T.int64(4), T.int64(3)): with Ts.sblock("T_reshape"): @@ -354,10 +354,10 @@ def test_tuple_get_reshape(): class Module: @Ts.prim_func def fused_reshape5( - lv2_0: T.Buffer((T.int64(2), T.int64(4096), T.int64(320)), "float16"), - lv2_1: T.Buffer((T.int64(2), T.int64(4096), T.int64(320)), "float16"), - lv2_2: T.Buffer((T.int64(2), T.int64(4096), T.int64(320)), "float16"), - T_reshape_handle_intermediate: T.Buffer( + lv2_0: T.Tensor((T.int64(2), T.int64(4096), T.int64(320)), "float16"), + lv2_1: T.Tensor((T.int64(2), T.int64(4096), T.int64(320)), "float16"), + lv2_2: T.Tensor((T.int64(2), T.int64(4096), T.int64(320)), "float16"), + T_reshape_handle_intermediate: T.Tensor( (T.int64(2), T.int64(4096), T.int64(8), T.int64(40)), "float16" ), ): @@ -415,10 +415,10 @@ def main( class Expected: @Ts.prim_func def fused_reshape5( - lv2_0: T.Buffer((T.int64(2), T.int64(4096), T.int64(320)), "float16"), - lv2_1: T.Buffer((T.int64(2), T.int64(4096), T.int64(320)), "float16"), - lv2_2: T.Buffer((T.int64(2), T.int64(4096), T.int64(320)), "float16"), - T_reshape_handle_intermediate: T.Buffer( + lv2_0: T.Tensor((T.int64(2), T.int64(4096), T.int64(320)), "float16"), + lv2_1: T.Tensor((T.int64(2), T.int64(4096), T.int64(320)), "float16"), + lv2_2: T.Tensor((T.int64(2), T.int64(4096), T.int64(320)), "float16"), + T_reshape_handle_intermediate: T.Tensor( (T.int64(2), T.int64(4096), T.int64(8), T.int64(40)), "float16" ), ): @@ -481,8 +481,8 @@ class Module: # of the input. @Ts.prim_func def strided_slice( - A: T.Buffer((T.int64(1), T.int64(1024)), "int32"), - T_strided_slice: T.Buffer((T.int64(1), T.int64(1000)), "int32"), + A: T.Tensor((T.int64(1), T.int64(1024)), "int32"), + T_strided_slice: T.Tensor((T.int64(1), T.int64(1000)), "int32"), ): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(1), T.int64(1000)): @@ -494,7 +494,7 @@ def strided_slice( @Ts.prim_func def add_one( - A: T.Buffer((T.int64(1), T.int64(1000)), "int32"), + A: T.Tensor((T.int64(1), T.int64(1000)), "int32"), T_add_one: T.buffer((T.int64(1), T.int64(1000)), "int32"), ): for ax0, ax1 in T.grid(T.int64(1), T.int64(1000)): @@ -550,9 +550,9 @@ def main(x: R.Tensor((), dtype="float32")) -> R.Tensor((1,), dtype="float32"): class Expected: @Ts.prim_func(private=True) def add( - A: T.Buffer((T.int64(1),), "float32"), - B: T.Buffer((T.int64(1),), "float32"), - T_add: T.Buffer((T.int64(1),), "float32"), + A: T.Tensor((T.int64(1),), "float32"), + B: T.Tensor((T.int64(1),), "float32"), + T_add: T.Tensor((T.int64(1),), "float32"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -564,7 +564,7 @@ def add( T_add[v_ax0] = A[v_ax0] + B[v_ax0] @Ts.prim_func(private=True) - def reshape(A: T.Buffer((), "float32"), T_reshape: T.Buffer((T.int64(1),), "float32")): + def reshape(A: T.Tensor((), "float32"), T_reshape: T.Tensor((T.int64(1),), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for ax0 in range(T.int64(1)): @@ -614,9 +614,9 @@ def main(x: R.Tensor((256,), dtype="float32")): @Ts.prim_func(private=True) def add( - y1: T.Buffer((T.int64(64), T.int64(4)), "float32"), - y2: T.Buffer((T.int64(64), T.int64(4)), "float32"), - z: T.Buffer((T.int64(64), T.int64(4)), "float32"), + y1: T.Tensor((T.int64(64), T.int64(4)), "float32"), + y2: T.Tensor((T.int64(64), T.int64(4)), "float32"), + z: T.Tensor((T.int64(64), T.int64(4)), "float32"), ): T.func_attr({"tirx.noalias": True}) @@ -672,10 +672,10 @@ def add( # @T.prim_func(private=True) # def add( -# y1: T.Buffer([N // 4, 4], "float32"), -# y2: T.Buffer([N // 4, 4], "float32"), +# y1: T.Tensor([N // 4, 4], "float32"), +# y2: T.Tensor([N // 4, 4], "float32"), # N: T.int64, -# z: T.Buffer([N // 4, 4], "float32"), +# z: T.Tensor([N // 4, 4], "float32"), # ): @@ -737,10 +737,10 @@ def main(x: R.Tensor([N, 16], dtype="float32")): @Ts.prim_func(private=True) def add( - y1: T.Buffer([add_N * 4, T.int64(4)], "float32"), - y2: T.Buffer([add_N * 4, T.int64(4)], "float32"), + y1: T.Tensor([add_N * 4, T.int64(4)], "float32"), + y2: T.Tensor([add_N * 4, T.int64(4)], "float32"), N: add_N, - z: T.Buffer([add_N * 4, T.int64(4)], "float32"), + z: T.Tensor([add_N * 4, T.int64(4)], "float32"), ): T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/relax/test_transform_specialize_primfunc_based_on_callsite.py b/tests/python/relax/test_transform_specialize_primfunc_based_on_callsite.py index 60b2b42963b3..4ec647a8588a 100644 --- a/tests/python/relax/test_transform_specialize_primfunc_based_on_callsite.py +++ b/tests/python/relax/test_transform_specialize_primfunc_based_on_callsite.py @@ -99,8 +99,8 @@ class Input: @Ts.prim_func(private=True) def max_pool2d_opencl( - gv: T.Buffer((T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32"), - pool_max: T.Buffer( + gv: T.Tensor((T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32"), + pool_max: T.Tensor( (T.int64(2), T.int64(1), T.int64(13), T.int64(13), T.int64(4)), "float32" ), ): @@ -140,8 +140,8 @@ def max_pool2d_opencl( @Ts.prim_func(private=True) def te_layout_transform( - x: T.Buffer((T.int64(2), T.int64(4), T.int64(26), T.int64(26)), "float32"), - te_layout_transform: T.Buffer( + x: T.Tensor((T.int64(2), T.int64(4), T.int64(26), T.int64(26)), "float32"), + te_layout_transform: T.Tensor( (T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32" ), ): @@ -161,10 +161,10 @@ def te_layout_transform( @Ts.prim_func(private=True) def te_layout_transform2( - lv2: T.Buffer( + lv2: T.Tensor( (T.int64(2), T.int64(1), T.int64(13), T.int64(13), T.int64(4)), "float32" ), - te_layout_transform: T.Buffer( + te_layout_transform: T.Tensor( (T.int64(2), T.int64(4), T.int64(13), T.int64(13)), "float32" ), ): @@ -228,9 +228,9 @@ class Input: @Ts.prim_func(private=True) def conv2d_NCHWc_OIHWo_opencl( - lv: T.Buffer((T.int64(2), T.int64(4), T.int64(28), T.int64(28), T.int64(4)), "float32"), - lv1: T.Buffer((T.int64(1), T.int64(16), T.int64(3), T.int64(3), T.int64(4)), "float32"), - conv2d_NCHWc_OIHWo: T.Buffer( + lv: T.Tensor((T.int64(2), T.int64(4), T.int64(28), T.int64(28), T.int64(4)), "float32"), + lv1: T.Tensor((T.int64(1), T.int64(16), T.int64(3), T.int64(3), T.int64(4)), "float32"), + conv2d_NCHWc_OIHWo: T.Tensor( (T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32" ), ): @@ -238,11 +238,11 @@ def conv2d_NCHWc_OIHWo_opencl( @Ts.prim_func(private=True) def fused_relu_concatenate_split( - gv: T.Buffer((T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32"), - T_split_sections_intermediate: T.Buffer( + gv: T.Tensor((T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32"), + T_split_sections_intermediate: T.Tensor( (T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32" ), - T_split_sections_intermediate_1: T.Buffer( + T_split_sections_intermediate_1: T.Tensor( (T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32" ), ): @@ -251,8 +251,8 @@ def fused_relu_concatenate_split( @Ts.prim_func(private=True) def te_layout_transform( - x: T.Buffer((T.int64(2), T.int64(16), T.int64(28), T.int64(28)), "float32"), - te_layout_transform: T.Buffer( + x: T.Tensor((T.int64(2), T.int64(16), T.int64(28), T.int64(28)), "float32"), + te_layout_transform: T.Tensor( (T.int64(2), T.int64(4), T.int64(28), T.int64(28), T.int64(4)), "float32" ), ): @@ -260,8 +260,8 @@ def te_layout_transform( @Ts.prim_func(private=True) def te_layout_transform1( - w: T.Buffer((T.int64(4), T.int64(16), T.int64(3), T.int64(3)), "float32"), - te_layout_transform: T.Buffer( + w: T.Tensor((T.int64(4), T.int64(16), T.int64(3), T.int64(3)), "float32"), + te_layout_transform: T.Tensor( (T.int64(1), T.int64(16), T.int64(3), T.int64(3), T.int64(4)), "float32" ), ): @@ -269,10 +269,10 @@ def te_layout_transform1( @Ts.prim_func(private=True) def te_layout_transform2( - lv3: T.Buffer( + lv3: T.Tensor( (T.int64(2), T.int64(1), T.int64(26), T.int64(26), T.int64(4)), "float32" ), - te_layout_transform: T.Buffer( + te_layout_transform: T.Tensor( (T.int64(2), T.int64(4), T.int64(26), T.int64(26)), "float32" ), ): diff --git a/tests/python/relax/test_transform_split_layout_rewrite_preproc.py b/tests/python/relax/test_transform_split_layout_rewrite_preproc.py index 43b6f7d3e9e2..bd4f995ba391 100644 --- a/tests/python/relax/test_transform_split_layout_rewrite_preproc.py +++ b/tests/python/relax/test_transform_split_layout_rewrite_preproc.py @@ -28,9 +28,9 @@ def test_single_buffer(): class Before: @Ts.prim_func(private=True) def tir_func( - X: T.Buffer((224, 224), "float32"), - W: T.Buffer((224, 224), "float32"), - Out: T.Buffer((224, 224), "float32"), + X: T.Tensor((224, 224), "float32"), + W: T.Tensor((224, 224), "float32"), + Out: T.Tensor((224, 224), "float32"), ): T.func_attr({"layout_free_buffers": [1]}) W_rewrite = Ts.sblock_alloc_buffer((4, 4, 56, 56)) @@ -61,9 +61,9 @@ def forward( class After: @Ts.prim_func(private=True) def tir_func_prepacked( - X: T.Buffer((224, 224), "float32"), - W_rewrite: T.Buffer((4, 4, 56, 56), "float32"), - Out: T.Buffer((224, 224), "float32"), + X: T.Tensor((224, 224), "float32"), + W_rewrite: T.Tensor((4, 4, 56, 56), "float32"), + Out: T.Tensor((224, 224), "float32"), ): for i0, j0, i1, j1 in T.grid(4, 4, 56, 56): with Ts.sblock("Out"): @@ -73,8 +73,8 @@ def tir_func_prepacked( @Ts.prim_func(private=True) def tir_func_weight_prepack( - W: T.Buffer((224, 224), "float32"), - W_rewrite: T.Buffer((4, 4, 56, 56), "float32"), + W: T.Tensor((224, 224), "float32"), + W_rewrite: T.Tensor((4, 4, 56, 56), "float32"), ): for i, j in T.grid(224, 224): with Ts.sblock("W_rewrite"): @@ -108,10 +108,10 @@ def test_multiple_buffers(): class Before: @Ts.prim_func(private=True) def tir_func( - X: T.Buffer((224, 224), "float32"), - W1: T.Buffer((224, 224), "float32"), - W2: T.Buffer((224, 224), "float32"), - Out: T.Buffer((224, 224), "float32"), + X: T.Tensor((224, 224), "float32"), + W1: T.Tensor((224, 224), "float32"), + W2: T.Tensor((224, 224), "float32"), + Out: T.Tensor((224, 224), "float32"), ): W1_rewrite = Ts.sblock_alloc_buffer((4, 4, 56, 56)) W2_rewrite = Ts.sblock_alloc_buffer((4, 4, 56, 56)) @@ -154,10 +154,10 @@ def forward( class After: @Ts.prim_func(private=True) def tir_func_prepacked( - X: T.Buffer((224, 224), "float32"), - W1_rewrite: T.Buffer((4, 4, 56, 56), "float32"), - W2_rewrite: T.Buffer((4, 4, 56, 56), "float32"), - Out: T.Buffer((224, 224), "float32"), + X: T.Tensor((224, 224), "float32"), + W1_rewrite: T.Tensor((4, 4, 56, 56), "float32"), + W2_rewrite: T.Tensor((4, 4, 56, 56), "float32"), + Out: T.Tensor((224, 224), "float32"), ): for i0, j0, i1, j1 in T.grid(4, 4, 56, 56): with Ts.sblock("Out"): @@ -171,10 +171,10 @@ def tir_func_prepacked( @Ts.prim_func(private=True) def tir_func_weight_prepack( - W1: T.Buffer((224, 224), "float32"), - W2: T.Buffer((224, 224), "float32"), - W1_rewrite: T.Buffer((4, 4, 56, 56), "float32"), - W2_rewrite: T.Buffer((4, 4, 56, 56), "float32"), + W1: T.Tensor((224, 224), "float32"), + W2: T.Tensor((224, 224), "float32"), + W1_rewrite: T.Tensor((4, 4, 56, 56), "float32"), + W2_rewrite: T.Tensor((4, 4, 56, 56), "float32"), ): for i, j in T.grid(224, 224): with Ts.sblock("W1_rewrite"): @@ -220,9 +220,9 @@ def test_attr_inheritance(): class Before: @Ts.prim_func(private=True) def tir_func( - X: T.Buffer((224, 224), "float32"), - W: T.Buffer((224, 224), "float32"), - Out: T.Buffer((224, 224), "float32"), + X: T.Tensor((224, 224), "float32"), + W: T.Tensor((224, 224), "float32"), + Out: T.Tensor((224, 224), "float32"), ): T.func_attr({"layout_free_buffers": [1], "tirx.noalias": True}) W_rewrite = Ts.sblock_alloc_buffer((4, 4, 56, 56)) @@ -253,9 +253,9 @@ def forward( class After: @Ts.prim_func(private=True) def tir_func_prepacked( - X: T.Buffer((224, 224), "float32"), - W_rewrite: T.Buffer((4, 4, 56, 56), "float32"), - Out: T.Buffer((224, 224), "float32"), + X: T.Tensor((224, 224), "float32"), + W_rewrite: T.Tensor((4, 4, 56, 56), "float32"), + Out: T.Tensor((224, 224), "float32"), ): T.func_attr({"tirx.noalias": True}) for i0, j0, i1, j1 in T.grid(4, 4, 56, 56): @@ -266,8 +266,8 @@ def tir_func_prepacked( @Ts.prim_func(private=True) def tir_func_weight_prepack( - W: T.Buffer((224, 224), "float32"), - W_rewrite: T.Buffer((4, 4, 56, 56), "float32"), + W: T.Tensor((224, 224), "float32"), + W_rewrite: T.Tensor((4, 4, 56, 56), "float32"), ): T.func_attr({"tirx.noalias": True}) for i, j in T.grid(224, 224): 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 9b4e9c2941d1..04f0c1a2e930 100644 --- a/tests/python/relax/test_transform_static_plan_block_memory.py +++ b/tests/python/relax/test_transform_static_plan_block_memory.py @@ -32,27 +32,27 @@ def test_basic(): @tvm.script.ir_module class Module: @Ts.prim_func - def add(rxplaceholder: T.Buffer(T.int64(8), "float32"), rxplaceholder_1: T.Buffer((), "float32"), T_add: T.Buffer(T.int64(8), "float32")): + def add(rxplaceholder: T.Tensor(T.int64(8), "float32"), rxplaceholder_1: T.Tensor((), "float32"), T_add: T.Tensor(T.int64(8), "float32")): T.evaluate(0) @Ts.prim_func - def reshape(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), T_reshape: T.Buffer(T.int64(8), "float32")): + def reshape(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), T_reshape: T.Tensor(T.int64(8), "float32")): T.evaluate(0) @Ts.prim_func - def relu(rxplaceholder: T.Buffer(T.int64(8), "float32"), compute: T.Buffer(T.int64(8), "float32")): + def relu(rxplaceholder: T.Tensor(T.int64(8), "float32"), compute: T.Tensor(T.int64(8), "float32")): T.evaluate(0) @Ts.prim_func - def log(rxplaceholder: T.Buffer(T.int64(10), "float32"), compute: T.Buffer(T.int64(10), "float32")): + def log(rxplaceholder: T.Tensor(T.int64(10), "float32"), compute: T.Tensor(T.int64(10), "float32")): T.evaluate(0) @Ts.prim_func - def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def exp(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4)), "float32")): T.evaluate(0) @Ts.prim_func - def pad(rxplaceholder: T.Buffer(T.int64(8), "float32"), PadInput: T.Buffer(T.int64(10), "float32")): + def pad(rxplaceholder: T.Tensor(T.int64(8), "float32"), PadInput: T.Tensor(T.int64(10), "float32")): T.evaluate(0) @R.function @@ -81,27 +81,27 @@ def main(x: R.Tensor((2, 4), dtype="float32")) -> R.Tensor((10,), dtype="float32 @tvm.script.ir_module class Expected: @Ts.prim_func - def add(rxplaceholder: T.Buffer(T.int64(8), "float32"), rxplaceholder_1: T.Buffer((), "float32"), T_add: T.Buffer(T.int64(8), "float32")): + def add(rxplaceholder: T.Tensor(T.int64(8), "float32"), rxplaceholder_1: T.Tensor((), "float32"), T_add: T.Tensor(T.int64(8), "float32")): T.evaluate(0) @Ts.prim_func - def reshape(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), T_reshape: T.Buffer(T.int64(8), "float32")): + def reshape(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), T_reshape: T.Tensor(T.int64(8), "float32")): T.evaluate(0) @Ts.prim_func - def relu(rxplaceholder: T.Buffer(T.int64(8), "float32"), compute: T.Buffer(T.int64(8), "float32")): + def relu(rxplaceholder: T.Tensor(T.int64(8), "float32"), compute: T.Tensor(T.int64(8), "float32")): T.evaluate(0) @Ts.prim_func - def log(rxplaceholder: T.Buffer(T.int64(10), "float32"), compute: T.Buffer(T.int64(10), "float32")): + def log(rxplaceholder: T.Tensor(T.int64(10), "float32"), compute: T.Tensor(T.int64(10), "float32")): T.evaluate(0) @Ts.prim_func - def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def exp(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4)), "float32")): T.evaluate(0) @Ts.prim_func - def pad(rxplaceholder: T.Buffer(T.int64(8), "float32"), PadInput: T.Buffer(T.int64(10), "float32")): + def pad(rxplaceholder: T.Tensor(T.int64(8), "float32"), PadInput: T.Tensor(T.int64(10), "float32")): T.evaluate(0) @R.function @@ -131,27 +131,27 @@ def main(x: R.Tensor((2, 4), dtype="float32")) -> R.Tensor((10,), dtype="float32 @I.ir_module class ExpectedLowered: @Ts.prim_func - def add(rxplaceholder: T.Buffer((T.int64(8),), "float32"), rxplaceholder_1: T.Buffer((), "float32"), T_add: T.Buffer((T.int64(8),), "float32")): + def add(rxplaceholder: T.Tensor((T.int64(8),), "float32"), rxplaceholder_1: T.Tensor((), "float32"), T_add: T.Tensor((T.int64(8),), "float32")): T.evaluate(0) @Ts.prim_func - def exp(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), compute: T.Buffer((T.int64(2), T.int64(4)), "float32")): + def exp(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), compute: T.Tensor((T.int64(2), T.int64(4)), "float32")): T.evaluate(0) @Ts.prim_func - def log(rxplaceholder: T.Buffer((T.int64(10),), "float32"), compute: T.Buffer((T.int64(10),), "float32")): + def log(rxplaceholder: T.Tensor((T.int64(10),), "float32"), compute: T.Tensor((T.int64(10),), "float32")): T.evaluate(0) @Ts.prim_func - def pad(rxplaceholder: T.Buffer((T.int64(8),), "float32"), PadInput: T.Buffer((T.int64(10),), "float32")): + def pad(rxplaceholder: T.Tensor((T.int64(8),), "float32"), PadInput: T.Tensor((T.int64(10),), "float32")): T.evaluate(0) @Ts.prim_func - def relu(rxplaceholder: T.Buffer((T.int64(8),), "float32"), compute: T.Buffer((T.int64(8),), "float32")): + def relu(rxplaceholder: T.Tensor((T.int64(8),), "float32"), compute: T.Tensor((T.int64(8),), "float32")): T.evaluate(0) @Ts.prim_func - def reshape(rxplaceholder: T.Buffer((T.int64(2), T.int64(4)), "float32"), T_reshape: T.Buffer((T.int64(8),), "float32")): + def reshape(rxplaceholder: T.Tensor((T.int64(2), T.int64(4)), "float32"), T_reshape: T.Tensor((T.int64(8),), "float32")): T.evaluate(0) @R.function @@ -196,17 +196,17 @@ def test_different_dtype(): class Module: @Ts.prim_func def add( - A: T.Buffer((T.int64(2), T.int64(3)), "float32"), - B: T.Buffer((T.int64(2), T.int64(3)), "float32"), - C: T.Buffer((T.int64(2), T.int64(3)), "float32"), + A: T.Tensor((T.int64(2), T.int64(3)), "float32"), + B: T.Tensor((T.int64(2), T.int64(3)), "float32"), + C: T.Tensor((T.int64(2), T.int64(3)), "float32"), ): T.evaluate(0) @Ts.prim_func def add1( - A: T.Buffer((T.int64(2), T.int64(3)), "int32"), - B: T.Buffer((T.int64(2), T.int64(3)), "int32"), - C: T.Buffer((T.int64(2), T.int64(3)), "int32"), + A: T.Tensor((T.int64(2), T.int64(3)), "int32"), + B: T.Tensor((T.int64(2), T.int64(3)), "int32"), + C: T.Tensor((T.int64(2), T.int64(3)), "int32"), ): T.evaluate(0) @@ -232,17 +232,17 @@ def main( class Expected: @Ts.prim_func def add( - A: T.Buffer((T.int64(2), T.int64(3)), "float32"), - B: T.Buffer((T.int64(2), T.int64(3)), "float32"), - C: T.Buffer((T.int64(2), T.int64(3)), "float32"), + A: T.Tensor((T.int64(2), T.int64(3)), "float32"), + B: T.Tensor((T.int64(2), T.int64(3)), "float32"), + C: T.Tensor((T.int64(2), T.int64(3)), "float32"), ): T.evaluate(0) @Ts.prim_func def add1( - A: T.Buffer((T.int64(2), T.int64(3)), "int32"), - B: T.Buffer((T.int64(2), T.int64(3)), "int32"), - C: T.Buffer((T.int64(2), T.int64(3)), "int32"), + A: T.Tensor((T.int64(2), T.int64(3)), "int32"), + B: T.Tensor((T.int64(2), T.int64(3)), "int32"), + C: T.Tensor((T.int64(2), T.int64(3)), "int32"), ): T.evaluate(0) @@ -279,9 +279,9 @@ def test_dtype_bool(): class Module: @Ts.prim_func def add1( - A: T.Buffer((T.int64(2), T.int64(3)), "bool"), - B: T.Buffer((T.int64(2), T.int64(3)), "bool"), - C: T.Buffer((T.int64(2), T.int64(3)), "bool"), + A: T.Tensor((T.int64(2), T.int64(3)), "bool"), + B: T.Tensor((T.int64(2), T.int64(3)), "bool"), + C: T.Tensor((T.int64(2), T.int64(3)), "bool"), ): T.evaluate(0) @@ -300,9 +300,9 @@ def main(y: R.Tensor((2, 3), dtype="bool")) -> R.Tensor((2, 3), dtype="bool"): class Expected: @Ts.prim_func def add1( - A: T.Buffer((T.int64(2), T.int64(3)), "bool"), - B: T.Buffer((T.int64(2), T.int64(3)), "bool"), - C: T.Buffer((T.int64(2), T.int64(3)), "bool"), + A: T.Tensor((T.int64(2), T.int64(3)), "bool"), + B: T.Tensor((T.int64(2), T.int64(3)), "bool"), + C: T.Tensor((T.int64(2), T.int64(3)), "bool"), ): T.evaluate(0) @@ -329,9 +329,9 @@ def test_same_dtype(): class Module: @Ts.prim_func def add( - A: T.Buffer((T.int64(2), T.int64(3)), "float32"), - B: T.Buffer((T.int64(2), T.int64(3)), "float32"), - C: T.Buffer((T.int64(2), T.int64(3)), "float32"), + A: T.Tensor((T.int64(2), T.int64(3)), "float32"), + B: T.Tensor((T.int64(2), T.int64(3)), "float32"), + C: T.Tensor((T.int64(2), T.int64(3)), "float32"), ): T.evaluate(0) @@ -357,9 +357,9 @@ def main( class Expected: @Ts.prim_func def add( - A: T.Buffer((T.int64(2), T.int64(3)), "float32"), - B: T.Buffer((T.int64(2), T.int64(3)), "float32"), - C: T.Buffer((T.int64(2), T.int64(3)), "float32"), + A: T.Tensor((T.int64(2), T.int64(3)), "float32"), + B: T.Tensor((T.int64(2), T.int64(3)), "float32"), + C: T.Tensor((T.int64(2), T.int64(3)), "float32"), ): T.evaluate(0) @@ -392,11 +392,11 @@ def test_if_cond(): @tvm.script.ir_module class Module: @Ts.prim_func - def all_less_than_zero(A: T.Buffer((2, 3), "float32"), B: T.Buffer((), "bool")): + def all_less_than_zero(A: T.Tensor((2, 3), "float32"), B: T.Tensor((), "bool")): T.evaluate(0) @Ts.prim_func - def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")): + def exp(A: T.Tensor((2, 3), "float32"), B: T.Tensor((2, 3), "float32")): T.evaluate(0) @R.function @@ -428,7 +428,7 @@ def test_if_then_else(): @tvm.script.ir_module class Module: @Ts.prim_func - def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")): + def exp(A: T.Tensor((2, 3), "float32"), B: T.Tensor((2, 3), "float32")): T.evaluate(0) @R.function @@ -457,7 +457,7 @@ def test_cross_block_use(): @tvm.script.ir_module class Module: @Ts.prim_func - def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")): + def exp(A: T.Tensor((2, 3), "float32"), B: T.Tensor((2, 3), "float32")): T.evaluate(0) @R.function @@ -496,7 +496,7 @@ def test_nested_tuple(): @tvm.script.ir_module class Module: @Ts.prim_func - def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")): + def exp(A: T.Tensor((2, 3), "float32"), B: T.Tensor((2, 3), "float32")): T.evaluate(0) @R.function @@ -552,7 +552,7 @@ def main(x: R.Tensor((2, 3), dtype="float32")) -> R.Tensor((2, 3), dtype="float3 @tvm.script.ir_module class Expected: @Ts.prim_func - def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")): + def exp(A: T.Tensor((2, 3), "float32"), B: T.Tensor((2, 3), "float32")): T.evaluate(0) @R.function @@ -689,7 +689,7 @@ def test_symbolic_shape(): @tvm.script.ir_module class Module: @Ts.prim_func - def exp(A: T.Buffer((m_exp, n_exp), "float32"), B: T.Buffer((m_exp, n_exp), "float32")): + def exp(A: T.Tensor((m_exp, n_exp), "float32"), B: T.Tensor((m_exp, n_exp), "float32")): T.evaluate(0) @R.function @@ -710,7 +710,7 @@ def main(x: R.Tensor((m_main, n_main), "float32")): @tvm.script.ir_module class Expected: @Ts.prim_func - def exp(A: T.Buffer((m_exp, n_exp), "float32"), B: T.Buffer((m_exp, n_exp), "float32")): + def exp(A: T.Tensor((m_exp, n_exp), "float32"), B: T.Tensor((m_exp, n_exp), "float32")): T.evaluate(0) @R.function @@ -769,9 +769,9 @@ def test_reshape_param(): class Module: @Ts.prim_func def add( - A: T.Buffer((T.int64(2), T.int64(25), T.int64(2)), "float32"), - B: T.Buffer((T.int64(2), T.int64(25), T.int64(2)), "float32"), - C: T.Buffer((T.int64(2), T.int64(25), T.int64(2)), "float32"), + A: T.Tensor((T.int64(2), T.int64(25), T.int64(2)), "float32"), + B: T.Tensor((T.int64(2), T.int64(25), T.int64(2)), "float32"), + C: T.Tensor((T.int64(2), T.int64(25), T.int64(2)), "float32"), ): T.evaluate(0) @@ -799,17 +799,17 @@ def test_multiple_functions(): class Module: @Ts.prim_func def add( - A: T.Buffer((T.int64(2), T.int64(3)), "float32"), - B: T.Buffer((T.int64(2), T.int64(3)), "float32"), - C: T.Buffer((T.int64(2), T.int64(3)), "float32"), + A: T.Tensor((T.int64(2), T.int64(3)), "float32"), + B: T.Tensor((T.int64(2), T.int64(3)), "float32"), + C: T.Tensor((T.int64(2), T.int64(3)), "float32"), ): T.evaluate(0) @Ts.prim_func def add1( - A: T.Buffer((T.int64(2), T.int64(3)), "int32"), - B: T.Buffer((T.int64(2), T.int64(3)), "int32"), - C: T.Buffer((T.int64(2), T.int64(3)), "int32"), + A: T.Tensor((T.int64(2), T.int64(3)), "int32"), + B: T.Tensor((T.int64(2), T.int64(3)), "int32"), + C: T.Tensor((T.int64(2), T.int64(3)), "int32"), ): T.evaluate(0) @@ -853,17 +853,17 @@ def func2( class Expected: @Ts.prim_func def add( - A: T.Buffer((T.int64(2), T.int64(3)), "float32"), - B: T.Buffer((T.int64(2), T.int64(3)), "float32"), - C: T.Buffer((T.int64(2), T.int64(3)), "float32"), + A: T.Tensor((T.int64(2), T.int64(3)), "float32"), + B: T.Tensor((T.int64(2), T.int64(3)), "float32"), + C: T.Tensor((T.int64(2), T.int64(3)), "float32"), ): T.evaluate(0) @Ts.prim_func def add1( - A: T.Buffer((T.int64(2), T.int64(3)), "int32"), - B: T.Buffer((T.int64(2), T.int64(3)), "int32"), - C: T.Buffer((T.int64(2), T.int64(3)), "int32"), + A: T.Tensor((T.int64(2), T.int64(3)), "int32"), + B: T.Tensor((T.int64(2), T.int64(3)), "int32"), + C: T.Tensor((T.int64(2), T.int64(3)), "int32"), ): T.evaluate(0) @@ -1763,14 +1763,14 @@ def test_match_cast_preserves_storage_liveness(): @I.ir_module class Before: @Ts.prim_func - def copy(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def copy(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.evaluate(0) @Ts.prim_func def add( - A: T.Buffer((16,), "float32"), - B: T.Buffer((16,), "float32"), - C: T.Buffer((16,), "float32"), + A: T.Tensor((16,), "float32"), + B: T.Tensor((16,), "float32"), + C: T.Tensor((16,), "float32"), ): T.evaluate(0) @@ -1803,7 +1803,7 @@ def test_builtin_reshape_preserves_storage_liveness(): @I.ir_module class Before: @Ts.prim_func - def copy(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def copy(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.evaluate(0) @R.function @@ -1894,7 +1894,7 @@ def test_if_branches_do_not_share_storage_var(): @I.ir_module class Before: @Ts.prim_func - def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")): + def exp(A: T.Tensor((2, 3), "float32"), B: T.Tensor((2, 3), "float32")): T.evaluate(0) @R.function @@ -1926,7 +1926,7 @@ def test_if_branch_storage_not_reused_after_if(): @I.ir_module class Before: @Ts.prim_func - def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")): + def exp(A: T.Tensor((2, 3), "float32"), B: T.Tensor((2, 3), "float32")): T.evaluate(0) @R.function @@ -1958,7 +1958,7 @@ def test_if_branches_share_storage_allocated_before_if(): @I.ir_module class Before: @Ts.prim_func - def exp(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")): + def exp(A: T.Tensor((2, 3), "float32"), B: T.Tensor((2, 3), "float32")): T.evaluate(0) @R.function diff --git a/tests/python/relax/test_transform_to_mixed_precision.py b/tests/python/relax/test_transform_to_mixed_precision.py index 14815fcc7cd9..2fc16a3b5a00 100644 --- a/tests/python/relax/test_transform_to_mixed_precision.py +++ b/tests/python/relax/test_transform_to_mixed_precision.py @@ -1068,8 +1068,8 @@ def main(A: R.Tensor([64], "float16")): @Ts.prim_func def tir_identity( - Input: T.Buffer(64, "float16"), - Output: T.Buffer(64, "float16"), + Input: T.Tensor(64, "float16"), + Output: T.Tensor(64, "float16"), ): for i in range(64): with Ts.sblock("copy"): diff --git a/tests/python/relax/test_vm_alloc_storage_with_scope.py b/tests/python/relax/test_vm_alloc_storage_with_scope.py index 1d7ccb08e63b..23b908c74d44 100644 --- a/tests/python/relax/test_vm_alloc_storage_with_scope.py +++ b/tests/python/relax/test_vm_alloc_storage_with_scope.py @@ -31,9 +31,9 @@ class Module: @Ts.prim_func def add( - arg0: T.Buffer((2, 2), "float32"), - arg1: T.Buffer((2, 2), "float32"), - output: T.Buffer((2, 2), "float32"), + arg0: T.Tensor((2, 2), "float32"), + arg1: T.Tensor((2, 2), "float32"), + output: T.Tensor((2, 2), "float32"), ): T.func_attr({"operator_name": "relax.add"}) for ax0 in range(2): diff --git a/tests/python/relax/test_vm_build.py b/tests/python/relax/test_vm_build.py index 602fc27e4b4d..04525e477870 100644 --- a/tests/python/relax/test_vm_build.py +++ b/tests/python/relax/test_vm_build.py @@ -226,9 +226,9 @@ def test_vm_compile_e2e_func_param_with_shape(): class TestVMCompileE2E2: @Ts.prim_func def tir_matmul( - A: T.Buffer((m_tir_matmul, n_tir_matmul)), - B: T.Buffer((n_tir_matmul, k_tir_matmul)), - C: T.Buffer((m_tir_matmul, k_tir_matmul)), + A: T.Tensor((m_tir_matmul, n_tir_matmul)), + B: T.Tensor((n_tir_matmul, k_tir_matmul)), + C: T.Tensor((m_tir_matmul, k_tir_matmul)), ) -> None: T.func_attr({"global_symbol": "tir_matmul"}) @@ -265,10 +265,10 @@ def test_call_tir_inplace_e2e_simple(): class TestCallTIRInplaceE2ESimple: @Ts.prim_func def copy( - A: T.Buffer((2, 3), "int32"), - B: T.Buffer((2, 3), "int32"), - C: T.Buffer((2, 3), "int32"), - out1: T.Buffer((2, 3), "int32"), + A: T.Tensor((2, 3), "int32"), + B: T.Tensor((2, 3), "int32"), + C: T.Tensor((2, 3), "int32"), + out1: T.Tensor((2, 3), "int32"), ): # copies the contents of C into A, B, and out1 T.func_attr({"tirx.noalias": True}) @@ -323,7 +323,7 @@ def test_call_tir_inplace_e2e_rw(): @tvm.script.ir_module class TestCallTIRInplaceE2ERW: @Ts.prim_func - def inplace_add(A: T.Buffer((2, 3), "int32"), B: T.Buffer((2, 3), "int32")): + def inplace_add(A: T.Tensor((2, 3), "int32"), B: T.Tensor((2, 3), "int32")): # sums A and B, storing the result in A T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): @@ -691,7 +691,7 @@ def main(x: R.Tensor((2, 3), dtype="float32")): return y @Ts.prim_func - def copy(A: T.Buffer((2, 3), "float32"), B: T.Buffer((2, 3), "float32")): + def copy(A: T.Tensor((2, 3), "float32"), B: T.Tensor((2, 3), "float32")): for i0, i1 in T.grid(2, 3): with Ts.sblock("block"): vi0, vi1 = Ts.axis.remap("SS", [i0, i1]) @@ -714,7 +714,7 @@ def test_sub_func_call(): @tvm.script.ir_module class TestVMSubFunction: @Ts.prim_func - def tir_matmul(A: T.Buffer((m, n)), B: T.Buffer((n, k)), C: T.Buffer((m, k))) -> None: + def tir_matmul(A: T.Tensor((m, n)), B: T.Tensor((n, k)), C: T.Tensor((m, k))) -> None: T.func_attr({"global_symbol": "tir_matmul"}) for i, j, k_index in T.grid(m, k, n): @@ -893,7 +893,7 @@ def main(x: R.Tensor((1,), "float32"), y: R.Tensor((1,), "float32")): @tvm.script.ir_module class TestVMSetInput: @Ts.prim_func - def test_vm_mul(A: T.Buffer((m, n)), B: T.Buffer((m, n)), C: T.Buffer((m, n))): + def test_vm_mul(A: T.Tensor((m, n)), B: T.Tensor((m, n)), C: T.Tensor((m, n))): T.func_attr({"global_symbol": "test_vm_mul"}) for i, j in T.grid(m, n): @@ -942,7 +942,7 @@ class ModA: I.module_attrs({"system_lib_prefix": "libA_"}) @Ts.prim_func - def tir_init(x: T.Buffer([N], "float32")): + def tir_init(x: T.Tensor([N], "float32")): for i in range(N): x[i] = T.float32(0) @@ -959,7 +959,7 @@ class ModB: I.module_attrs({"system_lib_prefix": "libB_"}) @Ts.prim_func - def tir_init(x: T.Buffer([N], "float32")): + def tir_init(x: T.Tensor([N], "float32")): for i in range(N): x[i] = T.float32(1) diff --git a/tests/python/relax/test_vm_codegen_only.py b/tests/python/relax/test_vm_codegen_only.py index 52c6b45d9ec2..b9125f506a92 100644 --- a/tests/python/relax/test_vm_codegen_only.py +++ b/tests/python/relax/test_vm_codegen_only.py @@ -355,7 +355,7 @@ def test_vm_kill_object(): @I.ir_module class TestKillObject: @Ts.prim_func - def full(T_full: T.Buffer((T.int64(4),), "float32")): + def full(T_full: T.Tensor((T.int64(4),), "float32")): T.func_attr({"global_symbol": "full", "tirx.noalias": True}) for ax0 in range(T.int64(4)): with Ts.sblock("T_full"): @@ -365,7 +365,7 @@ def full(T_full: T.Buffer((T.int64(4),), "float32")): T_full[v_ax0] = T.float32(0) @Ts.prim_func - def full1(T_full: T.Buffer((T.int64(4),), "float32")): + def full1(T_full: T.Tensor((T.int64(4),), "float32")): T.func_attr({"global_symbol": "full1", "tirx.noalias": True}) for ax0 in range(T.int64(4)): with Ts.sblock("T_full"): diff --git a/tests/python/relax/test_vm_cuda_graph.py b/tests/python/relax/test_vm_cuda_graph.py index 3a1e665863f0..c5bb40b9127d 100644 --- a/tests/python/relax/test_vm_cuda_graph.py +++ b/tests/python/relax/test_vm_cuda_graph.py @@ -55,7 +55,7 @@ def main(x: R.Tensor((16, 16), dtype="float32")) -> R.Tensor((16, 16), dtype="fl return lv5 @Ts.prim_func - def add(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")): + def add(A: T.Tensor((16, 16), "float32"), B: T.Tensor((16, 16), "float32")): T.func_attr({"global_symbol": "add"}) with Ts.sblock("root"): for i in T.thread_binding(16, thread="threadIdx.x"): diff --git a/tests/python/relax/texture/test_texture_nd.py b/tests/python/relax/texture/test_texture_nd.py index e04bae37fb78..e5b27a495c5d 100644 --- a/tests/python/relax/texture/test_texture_nd.py +++ b/tests/python/relax/texture/test_texture_nd.py @@ -124,7 +124,7 @@ def test_texture_copy(backend, dtype, channel_size, read_width): @I.ir_module class TextureCopy: @Ts.prim_func - def main(A: T.Buffer((M, N), dtype), B: T.Buffer((M, N), dtype)): + def main(A: T.Tensor((M, N), dtype), B: T.Tensor((M, N), dtype)): T.func_attr({"global_symbol": "main"}) for li, lj in T.grid(M, N): with Ts.sblock("Copy"): diff --git a/tests/python/runtime/test_executable.py b/tests/python/runtime/test_executable.py index 8b505bb33bdd..91f732ef425f 100644 --- a/tests/python/runtime/test_executable.py +++ b/tests/python/runtime/test_executable.py @@ -32,9 +32,9 @@ class MyModule: @Ts.prim_func def add( - A: T.Buffer((10,), "float32"), - B: T.Buffer((10,), "float32"), - C: T.Buffer((10,), "float32"), + A: T.Tensor((10,), "float32"), + B: T.Tensor((10,), "float32"), + C: T.Tensor((10,), "float32"), ): for i in range(10): C[i] = A[i] + B[i] diff --git a/tests/python/runtime/test_runtime_extension.py b/tests/python/runtime/test_runtime_extension.py index cc1bfac003e4..6d7029a0cc4c 100644 --- a/tests/python/runtime/test_runtime_extension.py +++ b/tests/python/runtime/test_runtime_extension.py @@ -28,7 +28,7 @@ def test_dltensor_compatible(): @I.ir_module class Module: @Ts.prim_func - def arange(Ab: T.Buffer((n,), "int64")): + def arange(Ab: T.Tensor((n,), "int64")): for i in T.serial(n - 1): Ab[i + 1] = Ab[i] + T.int64(1) diff --git a/tests/python/runtime/test_runtime_module_load.py b/tests/python/runtime/test_runtime_module_load.py index 91479f72b3a6..df667c8cf6c3 100644 --- a/tests/python/runtime/test_runtime_module_load.py +++ b/tests/python/runtime/test_runtime_module_load.py @@ -51,7 +51,7 @@ def test_dso_module_load(): def save_object(names): n = te.var("n") - Ab = tvm.tirx.decl_buffer((n,), dtype) + Ab = tvm.tirx.decl_tensor((n,), dtype) i = te.var("i") # for i in 0 to n-1: stmt = tvm.tirx.For( diff --git a/tests/python/s_tir/analysis/test_s_tir_analysis_calculate_allocated_memory.py b/tests/python/s_tir/analysis/test_s_tir_analysis_calculate_allocated_memory.py index 9e3b6efca3a2..ae758912f0bd 100644 --- a/tests/python/s_tir/analysis/test_s_tir_analysis_calculate_allocated_memory.py +++ b/tests/python/s_tir/analysis/test_s_tir_analysis_calculate_allocated_memory.py @@ -29,13 +29,13 @@ @tvm.script.ir_module class Module: @Ts.prim_func - def scale_by_two(a: T.Buffer((128,), "int8"), c: T.Buffer((128,), "int8")): + def scale_by_two(a: T.Tensor((128,), "int8"), c: T.Tensor((128,), "int8")): for i in T.serial(128): with Ts.sblock("C"): c[i] = a[i] * T.int8(2) @Ts.prim_func - def scale_by_two_three(a: T.Buffer((128,), "int8"), c: T.Buffer((128,), "int8")): + def scale_by_two_three(a: T.Tensor((128,), "int8"), c: T.Tensor((128,), "int8")): B = Ts.sblock_alloc_buffer([128], dtype="int8", scope="global.vtcm") for i in T.serial(128): with Ts.sblock("B"): @@ -71,9 +71,9 @@ def test_scale_by(primFunc, size): @Ts.prim_func def matmul_mix_scope( - A: T.Buffer([128, 128], scope="global"), - B: T.Buffer([128, 128], scope="global"), - C: T.Buffer([128, 128], scope="global"), + A: T.Tensor([128, 128], scope="global"), + B: T.Tensor([128, 128], scope="global"), + C: T.Tensor([128, 128], scope="global"), ) -> None: A_allocated = Ts.sblock_alloc_buffer([128, 128], dtype="float32", scope="global.texture") B_allocated = Ts.sblock_alloc_buffer([128, 128], dtype="float32", scope="global.texture") diff --git a/tests/python/s_tir/analysis/test_s_tir_analysis_estimate_tir_flops.py b/tests/python/s_tir/analysis/test_s_tir_analysis_estimate_tir_flops.py index dad8c09009ad..fb2313821e87 100644 --- a/tests/python/s_tir/analysis/test_s_tir_analysis_estimate_tir_flops.py +++ b/tests/python/s_tir/analysis/test_s_tir_analysis_estimate_tir_flops.py @@ -53,7 +53,7 @@ def test_te_workload(workload, flops): @Ts.prim_func -def flops_with_let(a: T.Buffer(16, "float32")): +def flops_with_let(a: T.Tensor(16, "float32")): for i in range(8): j = i + 8 a[j] = a[i] @@ -65,7 +65,7 @@ def test_flops_with_let(): @Ts.prim_func -def flops_with_if(a: T.Buffer(16, "float32"), b: T.Buffer(16, "float32")): +def flops_with_if(a: T.Tensor(16, "float32"), b: T.Tensor(16, "float32")): for i in range(16): if i % 2 == 0: a[i] = b[i] @@ -80,14 +80,14 @@ def test_flops_with_if(): @Ts.prim_func -def flops_with_forloop_as_expression(A: T.Buffer(1)): +def flops_with_forloop_as_expression(A: T.Tensor(1)): for i in T.serial(0, 16): for k in T.serial(0, i): A[0] = A[0] + 1 @Ts.prim_func -def flops_override(A: T.Buffer(16, "float32")): +def flops_override(A: T.Tensor(16, "float32")): T.func_attr({"estimated_flops": 32}) for i in range(16): A[0] = A[0] + 1 @@ -106,7 +106,7 @@ def test_estimate_flops_forloop_as_expression(): def test_estimate_flops_with_decl_buffer(): def make_func(use_decl_buffer): - buffer_func = T.decl_buffer if use_decl_buffer else T.Buffer + buffer_func = T.decl_tensor if use_decl_buffer else T.Tensor @Ts.prim_func def func(A_data: T.handle("float32")): @@ -122,7 +122,7 @@ def func(A_data: T.handle("float32")): @Ts.prim_func -def flops_with_nonint_extent(a: T.Buffer(16, "float32")): +def flops_with_nonint_extent(a: T.Tensor(16, "float32")): for i in range(4 + 4): a[i] = 2 * a[i] @@ -132,7 +132,7 @@ def test_flops_with_nonint_extent(): @Ts.prim_func -def flops_with_variable_extent(a: T.Buffer(16, "float32")): +def flops_with_variable_extent(a: T.Tensor(16, "float32")): for i in range(4 + 4): for j in range(i + 8): a[j] = 2 * a[i] diff --git a/tests/python/s_tir/analysis/test_s_tir_analysis_identify_memcpy.py b/tests/python/s_tir/analysis/test_s_tir_analysis_identify_memcpy.py index 794ccd554cb3..54c695fc54a6 100644 --- a/tests/python/s_tir/analysis/test_s_tir_analysis_identify_memcpy.py +++ b/tests/python/s_tir/analysis/test_s_tir_analysis_identify_memcpy.py @@ -52,7 +52,7 @@ def test_1d(): """Simplest test case""" @Ts.prim_func - def func(A: T.Buffer(1024, "float32"), B: T.Buffer(1024, "float32")): + def func(A: T.Tensor(1024, "float32"), B: T.Tensor(1024, "float32")): for i in T.serial(1024): B[i] = A[i] @@ -65,7 +65,7 @@ def test_1d_compute(): """Like test_1d, but a computation prevents this being a memcpy""" @Ts.prim_func - def func(A: T.Buffer(1024, "float32"), B: T.Buffer(1024, "float32")): + def func(A: T.Tensor(1024, "float32"), B: T.Tensor(1024, "float32")): for i in T.serial(1024): B[i] = A[i] + 1.0 @@ -77,7 +77,7 @@ def test_1d_conditional(): """Like test_1d, but a conditionals prevents this being a memcpy""" @Ts.prim_func - def func(A: T.Buffer(1024, "float32"), B: T.Buffer(1024, "float32")): + def func(A: T.Tensor(1024, "float32"), B: T.Tensor(1024, "float32")): for i in T.serial(1024): if i < 1024: B[i] = A[i] @@ -90,7 +90,7 @@ def test_1d_strided_input(): """Like test_1d, but strided input prevents this being a memcpy""" @Ts.prim_func - def func(A: T.Buffer(2048, "float32"), B: T.Buffer(1024, "float32")): + def func(A: T.Tensor(2048, "float32"), B: T.Tensor(1024, "float32")): for i in T.serial(1024): B[i] = A[i * 2] @@ -102,7 +102,7 @@ def test_1d_strided_output(): """Like test_1d, but strided output prevents this being a memcpy""" @Ts.prim_func - def func(A: T.Buffer(1024, "float32"), B: T.Buffer(2048, "float32")): + def func(A: T.Tensor(1024, "float32"), B: T.Tensor(2048, "float32")): for i in T.serial(1024): B[i * 2] = A[i] @@ -114,7 +114,7 @@ def test_1d_input_2d_output_fused_loop(): """Like test_1d, but the output is written as a 2-d buffer""" @Ts.prim_func - def func(A: T.Buffer(1024, "float32"), B: T.Buffer((32, 32), "float32")): + def func(A: T.Tensor(1024, "float32"), B: T.Tensor((32, 32), "float32")): for i in T.serial(1024): B[i // 32, i % 32] = A[i] @@ -127,7 +127,7 @@ def test_2d_input_1d_output_fused_loop(): """Like test_1d, but the input is written as a 2-d buffer""" @Ts.prim_func - def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer(1024, "float32")): + def func(A: T.Tensor((32, 32), "float32"), B: T.Tensor(1024, "float32")): for i in T.serial(1024): B[i] = A[i // 32, i % 32] @@ -146,7 +146,7 @@ def test_1d_input_1d_output_nested_loop(): """ @Ts.prim_func - def func(A: T.Buffer(1024, "float32"), B: T.Buffer(1024, "float32")): + def func(A: T.Tensor(1024, "float32"), B: T.Tensor(1024, "float32")): for i, j in T.grid(32, 32): B[i * 32 + j] = A[i * 32 + j] @@ -168,7 +168,7 @@ def test_1d_input_1d_output_nested_loop_equivalent_expressions(): """ @Ts.prim_func - def func(A: T.Buffer(1024, "float32"), B: T.Buffer(1024, "float32")): + def func(A: T.Tensor(1024, "float32"), B: T.Tensor(1024, "float32")): for i, j in T.grid(32, 32): B[i * 32 + j] = A[j + i * 32] @@ -185,7 +185,7 @@ def test_1d_input_2d_output_nested_loop(): """Like test_1d_input_1d_output_nested_loop, but with a 2-d output buffer""" @Ts.prim_func - def func(A: T.Buffer(1024, "float32"), B: T.Buffer((32, 32), "float32")): + def func(A: T.Tensor(1024, "float32"), B: T.Tensor((32, 32), "float32")): for i, j in T.grid(32, 32): B[i, j] = A[i * 32 + j] @@ -202,7 +202,7 @@ def test_2d_input_1d_output_nested_loop(): """Like test_1d_input_1d_output_nested_loop, but with a 2-d input buffer""" @Ts.prim_func - def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer(1024, "float32")): + def func(A: T.Tensor((32, 32), "float32"), B: T.Tensor(1024, "float32")): for i, j in T.grid(32, 32): B[i * 32 + j] = A[i, j] @@ -219,7 +219,7 @@ def test_2d_input_2d_output_nested_loop(): """Like test_1d_input_1d_output_nested_loop, but with 2-d input/output buffers""" @Ts.prim_func - def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer((32, 32), "float32")): + def func(A: T.Tensor((32, 32), "float32"), B: T.Tensor((32, 32), "float32")): for i, j in T.grid(32, 32): B[i, j] = A[i, j] @@ -239,7 +239,7 @@ def test_2d_input_2d_output_transpose_output(): """ @Ts.prim_func - def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer((32, 32), "float32")): + def func(A: T.Tensor((32, 32), "float32"), B: T.Tensor((32, 32), "float32")): for i, j in T.grid(32, 32): B[j, i] = A[i, j] @@ -257,7 +257,7 @@ def test_2d_input_2d_output_transpose_input(): """ @Ts.prim_func - def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer((32, 32), "float32")): + def func(A: T.Tensor((32, 32), "float32"), B: T.Tensor((32, 32), "float32")): for i, j in T.grid(32, 32): B[i, j] = A[j, i] @@ -278,7 +278,7 @@ def test_2d_input_2d_output_transpose_both(): """ @Ts.prim_func - def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer((32, 32), "float32")): + def func(A: T.Tensor((32, 32), "float32"), B: T.Tensor((32, 32), "float32")): for i, j in T.grid(32, 32): B[j, i] = A[j, i] @@ -298,7 +298,7 @@ def test_cache_read(): """ @Ts.prim_func - def func(A: T.Buffer((32, 32), "float32"), B: T.Buffer(32, "float32")): + def func(A: T.Tensor((32, 32), "float32"), B: T.Tensor(32, "float32")): for i, j in T.grid(32, 32): B[j] = A[i, j] @@ -319,7 +319,7 @@ def test_cache_write(): """ @Ts.prim_func - def func(A: T.Buffer(32, "float32"), B: T.Buffer((32, 32), "float32")): + def func(A: T.Tensor(32, "float32"), B: T.Tensor((32, 32), "float32")): for i, j in T.grid(32, 32): B[i, j] = A[j] diff --git a/tests/python/s_tir/analysis/test_s_tir_analysis_is_pure_function.py b/tests/python/s_tir/analysis/test_s_tir_analysis_is_pure_function.py index 506f7ec25c0e..bbc7f256e6c5 100644 --- a/tests/python/s_tir/analysis/test_s_tir_analysis_is_pure_function.py +++ b/tests/python/s_tir/analysis/test_s_tir_analysis_is_pure_function.py @@ -60,21 +60,21 @@ def func(N: T.int32, M: T.int32) -> T.int32: class TestReadBufferArgument(CheckPureFunction): @Ts.prim_func - def func(A: T.Buffer(16, "float32")) -> T.float32: + def func(A: T.Tensor(16, "float32")) -> T.float32: return A[0] class TestWriteToBufferArgument(CheckImpureFunction): @Ts.prim_func - def func(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): + def func(A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32")): for i in range(16): B[i] = A[i] class TestWriteToInternalAllocation(CheckPureFunction): @Ts.prim_func - def func(A: T.Buffer([16, 16], "float32")) -> T.float32: - Sum = T.decl_buffer([], "float32") + def func(A: T.Tensor([16, 16], "float32")) -> T.float32: + Sum = T.decl_tensor([], "float32") Sum[()] = 0.0 for i, j in T.grid(16, 16): Sum[()] = Sum[()] + A[i, j] diff --git a/tests/python/s_tir/analysis/test_s_tir_analysis_oob.py b/tests/python/s_tir/analysis/test_s_tir_analysis_oob.py index 0340bd49a060..bcf6514ef73c 100644 --- a/tests/python/s_tir/analysis/test_s_tir_analysis_oob.py +++ b/tests/python/s_tir/analysis/test_s_tir_analysis_oob.py @@ -22,29 +22,29 @@ @Ts.prim_func -def bad_load(A: T.Buffer((2, 3), "float32"), B: T.Buffer((3, 2), "float32")): +def bad_load(A: T.Tensor((2, 3), "float32"), B: T.Tensor((3, 2), "float32")): B[0, 0] = A[2, 2] @Ts.prim_func -def bad_load_loop(A: T.Buffer((2, 3), "float32"), B: T.Buffer((3, 2), "float32")): +def bad_load_loop(A: T.Tensor((2, 3), "float32"), B: T.Tensor((3, 2), "float32")): for i in range(3): B[i, 0] = A[i, 2] @Ts.prim_func -def bad_store(A: T.Buffer((2, 3), "float32"), B: T.Buffer((3, 2), "float32")): +def bad_store(A: T.Tensor((2, 3), "float32"), B: T.Tensor((3, 2), "float32")): B[0, 3] = A[1, 2] @Ts.prim_func -def bad_store_loop(A: T.Buffer((2, 3), "float32"), B: T.Buffer((3, 2), "float32")): +def bad_store_loop(A: T.Tensor((2, 3), "float32"), B: T.Tensor((3, 2), "float32")): for i in range(3): B[0, i] = A[1, i] @Ts.prim_func -def unknown_bounds(A: T.Buffer((2, 3), "float32"), B: T.Buffer((3, 2), "float32"), N: T.int32): +def unknown_bounds(A: T.Tensor((2, 3), "float32"), B: T.Tensor((3, 2), "float32"), N: T.int32): for i in range(3): B[0, N] = A[1, i] diff --git a/tests/python/s_tir/analysis/test_s_tir_analysis_verify_block_scope.py b/tests/python/s_tir/analysis/test_s_tir_analysis_verify_block_scope.py index 5ef814b3285d..7d52109f27f9 100644 --- a/tests/python/s_tir/analysis/test_s_tir_analysis_verify_block_scope.py +++ b/tests/python/s_tir/analysis/test_s_tir_analysis_verify_block_scope.py @@ -27,7 +27,7 @@ def test_buffer_region_bounds_are_visited(): data = tvm.tirx.Var( "data", tvm.ir.PointerType(tvm.ir.PrimType("int32"), storage_scope="global") ) - buffer = tvm.tirx.decl_buffer([4], "int32", data=data) + buffer = tvm.tirx.decl_tensor([4], "int32", data=data) undefined = tvm.tirx.Var("undefined", "int32") region = tvm.tirx.BufferRegion(buffer, [tvm.ir.Range.from_min_extent(undefined, 4)]) block = tvm.s_tir.SBlock([], [region], [], "region", tvm.tirx.Evaluate(0)) @@ -38,8 +38,8 @@ def test_buffer_region_bounds_are_visited(): def test_fail_use_out_loop_var(): @Ts.prim_func(check_well_formed=False) def element_wise( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), ): for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -56,7 +56,7 @@ def test_block_match_buffer_defines_buffer_obj(): @I.ir_module class mod: @Ts.prim_func - def func(A: T.Buffer([256, 256], "float32")): + def func(A: T.Tensor([256, 256], "float32")): for (*iters,) in T.grid(16, 16, 16, 16): with Ts.sblock("compute"): tile_i, tile_j, i, j = Ts.axis.remap("SSSS", iters) @@ -77,7 +77,7 @@ def test_block_match_buffer_defines_symbolic_variables(): @I.ir_module class mod: @Ts.prim_func - def func(A: T.Buffer([256, 256], "int32")): + def func(A: T.Tensor([256, 256], "int32")): for (*iters,) in T.grid(16, 16, 16, 16): with Ts.sblock("compute"): tile_i, tile_j, i, j = Ts.axis.remap("SSSS", iters) @@ -99,7 +99,7 @@ def test_match_buffer_in_block_is_well_formed(): @I.ir_module class mod: @Ts.prim_func - def func(A: T.Buffer((128, 128), "float32")): + def func(A: T.Tensor((128, 128), "float32")): for (*iters,) in T.grid(8, 8, 16, 16): with Ts.sblock("compute"): ti, tj, i, j = Ts.axis.remap("SSSS", iters) @@ -117,13 +117,13 @@ def test_error_undeclared_buffer_in_schedulable_tir(): # Manually construct a BufferStore that uses a buffer without any declaration # inside a block context. n = tvm.tirx.Var("n", "int32") - A = tvm.tirx.decl_buffer([n], "float32", name="A") + A = tvm.tirx.decl_tensor([n], "float32", name="A") i = tvm.tirx.Var("i", "int32") # Create an undeclared buffer using an explicit data pointer that is NOT - # a function parameter and NOT wrapped with DeclBuffer. + # a function parameter and NOT wrapped with DeclTensor. B_data = tvm.tirx.Var("B_data", tvm.ir.PointerType(tvm.ir.PrimType("float32"))) - B = tvm.tirx.decl_buffer([n], "float32", name="B", data=B_data) + B = tvm.tirx.decl_tensor([n], "float32", name="B", data=B_data) # Build a block that writes to B without any declaration of B. bi = tvm.tirx.Var("bi", "int32") @@ -144,11 +144,11 @@ def test_error_undeclared_buffer_in_schedulable_tir(): params=[A, B_data], body=tvm.tirx.For(i, 0, n, tvm.tirx.ForKind.SERIAL, block_realize), # Note: B is NOT a function parameter, so its declaration scope is only - # within a DeclBuffer node (which we intentionally omit here). + # within a DeclTensor node (which we intentionally omit here). ) # B is used in the block but was never declared — should fail. with pytest.raises( - (ValueError, tvm.error.InternalError), match="buffer B.*without a prior DeclBuffer" + (ValueError, tvm.error.InternalError), match="buffer B.*without a prior DeclTensor" ): tvm.s_tir.analysis.verify_well_formed(prim_func) diff --git a/tests/python/s_tir/analysis/test_sblock_access_region.py b/tests/python/s_tir/analysis/test_sblock_access_region.py index ea27472ebdf2..31a16caa3951 100644 --- a/tests/python/s_tir/analysis/test_sblock_access_region.py +++ b/tests/python/s_tir/analysis/test_sblock_access_region.py @@ -135,13 +135,13 @@ def opaque_access_with_tvm_access_ptr_func() -> None: @Ts.prim_func def decl_buffer_alias_func( - A: T.Buffer((16,), "float32"), - B: T.Buffer((16,), "float32"), + A: T.Tensor((16,), "float32"), + B: T.Tensor((16,), "float32"), ) -> None: with Ts.sblock("alias"): Ts.reads(A[0]) Ts.writes(B[0]) - A_view = T.decl_buffer((16,), "float32", data=A.data) + A_view = T.decl_tensor((16,), "float32", data=A.data) B[0] = A[0] + A_view[0] @@ -409,10 +409,10 @@ def test_access_of_decompose_reduction(): def test_buffer_access_with_let_binding(): @Ts.prim_func def func( - storage: T.Buffer((16, 16, 16), "float32"), - seq_slot_ids: T.Buffer((16,), "int32"), - history_slot_ids: T.Buffer((16,), "int32"), - output: T.Buffer((16, 16), "float32"), + storage: T.Tensor((16, 16, 16), "float32"), + seq_slot_ids: T.Tensor((16,), "int32"), + history_slot_ids: T.Tensor((16,), "int32"), + output: T.Tensor((16, 16), "float32"), ): for i, s in T.grid(16, 16): with Ts.sblock("copy"): @@ -437,9 +437,9 @@ def func( def test_buffer_access_with_nested_let_binding(): @Ts.prim_func def func( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ): for i, s in T.grid(16, 16): with Ts.sblock("copy"): @@ -497,8 +497,8 @@ def test_conditional_inequality_access_regions(case): "rounded": ([x], [(-20, 41)], tirx.all(x * 3 >= -7, x * 2 <= 9), [(-2, 7)]), } variables, domains, condition, expected = cases[case] - inside = tirx.decl_buffer([256] * len(variables), name="inside") - outside = tirx.decl_buffer([256] * len(variables), name="outside") + inside = tirx.decl_tensor([256] * len(variables), name="inside") + outside = tirx.decl_tensor([256] * len(variables), name="outside") body = tirx.SeqStmt( [ tirx.IfThenElse(condition, tirx.Evaluate(inside[tuple(variables)]), None), diff --git a/tests/python/s_tir/analysis/test_sblock_buffer_access_lca.py b/tests/python/s_tir/analysis/test_sblock_buffer_access_lca.py index 32b13f998166..5a118aea56e7 100644 --- a/tests/python/s_tir/analysis/test_sblock_buffer_access_lca.py +++ b/tests/python/s_tir/analysis/test_sblock_buffer_access_lca.py @@ -22,7 +22,7 @@ @Ts.prim_func def buffer_load_store_func( - A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32") ) -> None: C = Ts.sblock_alloc_buffer((128, 128), "float32") D = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -47,12 +47,12 @@ def buffer_load_store_func( @Ts.prim_func def buffer_opaque_access( - B: T.Buffer([16, 16], "float32"), C: T.Buffer([16, 16], "float32") + B: T.Tensor([16, 16], "float32"), C: T.Tensor([16, 16], "float32") ) -> None: with Ts.sblock(): Ts.reads([]) Ts.writes(B[0:16, 0:16]) - A = T.decl_buffer([256], "float32") + A = T.decl_tensor([256], "float32") for i, j in T.grid(16, 16): A[i * 16 + j] = 1 for i in range(0, 16): @@ -68,13 +68,13 @@ def buffer_opaque_access( @Ts.prim_func -def lca_is_func_root(A: T.Buffer([0, 0], "float32")) -> None: +def lca_is_func_root(A: T.Tensor([0, 0], "float32")) -> None: A[0, 0] = 1.0 @Ts.prim_func def match_buffer_func( - A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32") ) -> None: for i, j in T.grid(8, 8): with Ts.sblock("block"): @@ -94,7 +94,7 @@ def match_buffer_func( @Ts.prim_func def global_buffer_with_blockidx( - a: T.Buffer((1, 32), "int32"), b: T.Buffer((1, 32), "int32") + a: T.Tensor((1, 32), "int32"), b: T.Tensor((1, 32), "int32") ) -> None: for i0 in T.thread_binding(0, 1, thread="blockIdx.x"): for i1 in T.thread_binding(0, 32, thread="threadIdx.x"): diff --git a/tests/python/s_tir/base/test_compilation_pipeline.py b/tests/python/s_tir/base/test_compilation_pipeline.py index 0bada4c82fa0..d979df5f93d1 100644 --- a/tests/python/s_tir/base/test_compilation_pipeline.py +++ b/tests/python/s_tir/base/test_compilation_pipeline.py @@ -30,14 +30,14 @@ @pytest.mark.parametrize("mixed", [False, True]) def test_default_pipeline_selects_dialect(mixed): @Ts.prim_func - def scheduled(A: Ts.Buffer((16,), "float32"), B: Ts.Buffer((16,), "float32")): + def scheduled(A: Ts.Tensor((16,), "float32"), B: Ts.Tensor((16,), "float32")): for i in range(16): with Ts.sblock("copy"): vi = Ts.axis.spatial(16, i) B[vi] = A[vi] + Ts.float32(1) @T.prim_func - def native(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def native(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): for i in range(16): B[i] = A[i] + T.float32(2) diff --git a/tests/python/s_tir/base/test_s_tir_texture_scope.py b/tests/python/s_tir/base/test_s_tir_texture_scope.py index ff9b63f579bb..77922460e5bb 100644 --- a/tests/python/s_tir/base/test_s_tir_texture_scope.py +++ b/tests/python/s_tir/base/test_s_tir_texture_scope.py @@ -34,8 +34,8 @@ def test_texture_scope(): class PlusOneMultTwo: @Ts.prim_func def main( - A: T.Buffer((128, 128, 4), dtype="float32", scope="global.texture"), - C: T.Buffer((128, 128, 4), dtype="float32", scope="global.texture"), + A: T.Tensor((128, 128, 4), dtype="float32", scope="global.texture"), + C: T.Tensor((128, 128, 4), dtype="float32", scope="global.texture"), ) -> None: T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/s_tir/base/test_sblock_dependence_info.py b/tests/python/s_tir/base/test_sblock_dependence_info.py index 8ba912ed49b8..d036c9252115 100644 --- a/tests/python/s_tir/base/test_sblock_dependence_info.py +++ b/tests/python/s_tir/base/test_sblock_dependence_info.py @@ -36,7 +36,7 @@ @Ts.prim_func -def elementwise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")) -> None: +def elementwise(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -54,7 +54,7 @@ def elementwise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "flo @Ts.prim_func def war_dependency( - A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), B: T.Tensor((128, 128)), C: T.Tensor((128, 128)) ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("C"): @@ -66,7 +66,7 @@ def war_dependency( @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j in T.grid(128, 128): with Ts.sblock("init"): vi, vj = Ts.axis.remap("SS", [i, j]) diff --git a/tests/python/s_tir/base/test_tir_te_extern_primfunc.py b/tests/python/s_tir/base/test_tir_te_extern_primfunc.py index d239352a4f0a..02b8d89e1e58 100644 --- a/tests/python/s_tir/base/test_tir_te_extern_primfunc.py +++ b/tests/python/s_tir/base/test_tir_te_extern_primfunc.py @@ -33,7 +33,7 @@ @Ts.prim_func -def func_1(A: T.Buffer((16,), "float32"), C: T.Buffer((1,), "float32")): +def func_1(A: T.Tensor((16,), "float32"), C: T.Tensor((1,), "float32")): for i in T.serial( 0, 16, @@ -61,7 +61,7 @@ def verify_func_1(module): @Ts.prim_func def func_2( - C: T.Buffer((1,), "float32"), A: T.Buffer((16,), "float32"), D: T.Buffer((2,), "float32") + C: T.Tensor((1,), "float32"), A: T.Tensor((16,), "float32"), D: T.Tensor((2,), "float32") ): for i in T.serial( 0, @@ -91,11 +91,11 @@ def verify_func_2(module): @Ts.prim_func def func_3( - C: T.Buffer((1,), "float32"), - A: T.Buffer((16,), "float32"), - D: T.Buffer((2,), "float32"), - E: T.Buffer((16,), "float32"), - F: T.Buffer((16,), "float32"), + C: T.Tensor((1,), "float32"), + A: T.Tensor((16,), "float32"), + D: T.Tensor((2,), "float32"), + E: T.Tensor((16,), "float32"), + F: T.Tensor((16,), "float32"), ): for i in T.serial( 0, @@ -133,11 +133,11 @@ def verify_func_3(module): @Ts.prim_func def func_4( - C: T.Buffer((1,), "float32"), - A: T.Buffer((16,), "float32"), - F: T.Buffer((16,), "float32"), - D: T.Buffer((2,), "float32"), - E: T.Buffer((16,), "float32"), + C: T.Tensor((1,), "float32"), + A: T.Tensor((16,), "float32"), + F: T.Tensor((16,), "float32"), + D: T.Tensor((2,), "float32"), + E: T.Tensor((16,), "float32"), ): for i in T.serial( 0, diff --git a/tests/python/s_tir/dlight/test_benchmark.py b/tests/python/s_tir/dlight/test_benchmark.py index 8fc5603668bb..f5f488da202c 100644 --- a/tests/python/s_tir/dlight/test_benchmark.py +++ b/tests/python/s_tir/dlight/test_benchmark.py @@ -51,7 +51,7 @@ @I.ir_module(check_well_formed=False) class Module: @Ts.prim_func - def full1(T_full: T.Buffer((T.int64(1), T.int64(32), T.int64(1), full1_n), 'float16')): + def full1(T_full: T.Tensor((T.int64(1), T.int64(32), T.int64(1), full1_n), 'float16')): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -63,7 +63,7 @@ def full1(T_full: T.Buffer((T.int64(1), T.int64(32), T.int64(1), full1_n), 'floa T_full[v_ax0, v_ax1, v_ax2, v_ax3] = T.float16(1.0) @Ts.prim_func - def full2(T_full: T.Buffer((T.int64(1), T.int64(32), full2_n, T.int64(128)), 'float16')): + def full2(T_full: T.Tensor((T.int64(1), T.int64(32), full2_n, T.int64(128)), 'float16')): T.func_attr({"op_pattern": 0, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -75,7 +75,7 @@ def full2(T_full: T.Buffer((T.int64(1), T.int64(32), full2_n, T.int64(128)), 'fl T_full[v_ax0, v_ax1, v_ax2, v_ax3] = T.float16(1.0) @Ts.prim_func - def matmul1(A: T.Buffer((T.int64(1), T.int64(32), T.int64(1), matmul1_n), 'float16'), B: T.Buffer((T.int64(1), T.int64(32), matmul1_n, T.int64(128)), 'float16'), matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16")): + def matmul1(A: T.Tensor((T.int64(1), T.int64(32), T.int64(1), matmul1_n), 'float16'), B: T.Tensor((T.int64(1), T.int64(32), matmul1_n, T.int64(128)), 'float16'), matmul: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16")): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -105,7 +105,7 @@ def test(): m = T.dynamic("m") @Ts.prim_func -def cuda_workload(inp0: T.Buffer((T.int64(1), m, T.int64(4096))), inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), matmul: T.Buffer((T.int64(1), m, T.int64(4096)))): +def cuda_workload(inp0: T.Tensor((T.int64(1), m, T.int64(4096))), inp1: T.Tensor((T.int64(4096), T.int64(4096)), "float32"), matmul: T.Tensor((T.int64(1), m, T.int64(4096)))): T.func_attr({"tirx.is_scheduled": True}) # 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 07c2c26ca2f7..e0d8cb53e713 100644 --- a/tests/python/s_tir/dlight/test_cpu_gemv.py +++ b/tests/python/s_tir/dlight/test_cpu_gemv.py @@ -32,7 +32,7 @@ def test_gemv_basic(): n = T.dynamic("n", "int32") @Ts.prim_func(private=True) - def before(lv1637: T.Buffer((1, 32, 1, 128), "float16"), lv1638: T.Buffer((1, 32, n, 128), 'float16'), lv1614: T.Buffer((1, 1, 1, n), 'float16'), var_compute_intermediate: T.Buffer((1, 32, 1, n))): + def before(lv1637: T.Tensor((1, 32, 1, 128), "float16"), lv1638: T.Tensor((1, 32, n, 128), 'float16'), lv1614: T.Tensor((1, 1, 1, n), 'float16'), var_compute_intermediate: T.Tensor((1, 32, 1, n))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -76,7 +76,7 @@ def before(lv1637: T.Buffer((1, 32, 1, 128), "float16"), lv1638: T.Buffer((1, 32 n = T.dynamic("n", "int32") @Ts.prim_func(private=True) - def expected(lv1637: T.Buffer((1, 32, 1, 128), "float16"), lv1638: T.Buffer((1, 32, n, 128), 'float16'), lv1614: T.Buffer((1, 1, 1, n), 'float16'), var_compute_intermediate: T.Buffer((1, 32, 1, n))): + def expected(lv1637: T.Tensor((1, 32, 1, 128), "float16"), lv1638: T.Tensor((1, 32, n, 128), 'float16'), lv1614: T.Tensor((1, 1, 1, n), 'float16'), var_compute_intermediate: T.Tensor((1, 32, 1, n))): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -114,7 +114,7 @@ def expected(lv1637: T.Buffer((1, 32, 1, 128), "float16"), lv1638: T.Buffer((1, def test_decode_gemv_256_threads(): # fmt: off @Ts.prim_func(private=True) - def before(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128), "float16"), lv1654: T.Buffer((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Buffer((1, 1, 22016), "float16")): + def before(lv571: T.Tensor((22016, 512), "uint32"), lv572: T.Tensor((22016, 128), "float16"), lv1654: T.Tensor((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Tensor((1, 1, 22016), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): p_output0_intermediate = Ts.sblock_alloc_buffer((22016, 4096), "float16") @@ -134,7 +134,7 @@ def before(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128) var_NT_matmul_intermediate[v_i0, v_i1, v_i2] = var_NT_matmul_intermediate[v_i0, v_i1, v_i2] + lv1654[v_i0, v_i1, v_k] * p_output0_intermediate[v_i2, v_k] @Ts.prim_func(private=True) - def expected(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128), "float16"), lv1654: T.Buffer((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Buffer((1, 1, 22016), "float16")): + def expected(lv571: T.Tensor((22016, 512), "uint32"), lv572: T.Tensor((22016, 128), "float16"), lv1654: T.Tensor((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Tensor((1, 1, 22016), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): for u_fused in range(1): @@ -162,7 +162,7 @@ def test_decode_gemv1(): # fmt: off @Ts.prim_func(private=True) - def before(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128), "float16"), lv1654: T.Buffer((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Buffer((1, 1, 22016), "float16")): + def before(lv571: T.Tensor((22016, 512), "uint32"), lv572: T.Tensor((22016, 128), "float16"), lv1654: T.Tensor((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Tensor((1, 1, 22016), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): p_output0_intermediate = Ts.sblock_alloc_buffer((22016, 4096), "float16") @@ -182,7 +182,7 @@ def before(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128) var_NT_matmul_intermediate[v_i0, v_i1, v_i2] = var_NT_matmul_intermediate[v_i0, v_i1, v_i2] + lv1654[v_i0, v_i1, v_k] * p_output0_intermediate[v_i2, v_k] @Ts.prim_func(private=True) - def expected(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128), "float16"), lv1654: T.Buffer((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Buffer((1, 1, 22016), "float16")): + def expected(lv571: T.Tensor((22016, 512), "uint32"), lv572: T.Tensor((22016, 128), "float16"), lv1654: T.Tensor((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Tensor((1, 1, 22016), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): for u_fused in range(1): @@ -210,7 +210,7 @@ def test_decode_gemv2(): # fmt: off @Ts.prim_func(private=True) - def before(lv771: T.Buffer((32000, 512), "uint32"), lv772: T.Buffer((32000, 128), "float16"), lv3216: T.Buffer((1, 1, 4096), "float16"), p_output0_intermediate: T.Buffer((1, 1, 32000), "float32")): + def before(lv771: T.Tensor((32000, 512), "uint32"), lv772: T.Tensor((32000, 128), "float16"), lv3216: T.Tensor((1, 1, 4096), "float16"), p_output0_intermediate: T.Tensor((1, 1, 32000), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): p_output0_intermediate_1 = Ts.sblock_alloc_buffer((32000, 4096), "float16") @@ -237,7 +237,7 @@ def before(lv771: T.Buffer((32000, 512), "uint32"), lv772: T.Buffer((32000, 128) p_output0_intermediate[v_i0, v_i1, v_i2] = T.Cast("float32", var_NT_matmul_intermediate[v_i0, v_i1, v_i2]) @Ts.prim_func(private=True) - def expected(lv771: T.Buffer((32000, 512), "uint32"), lv772: T.Buffer((32000, 128), "float16"), lv3216: T.Buffer((1, 1, 4096), "float16"), p_output0_intermediate: T.Buffer((1, 1, 32000), "float32")): + def expected(lv771: T.Tensor((32000, 512), "uint32"), lv772: T.Tensor((32000, 128), "float16"), lv3216: T.Tensor((1, 1, 4096), "float16"), p_output0_intermediate: T.Tensor((1, 1, 32000), "float32")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): var_NT_matmul_intermediate = Ts.sblock_alloc_buffer((1, 1, 32000), "float16") @@ -272,7 +272,7 @@ def test_decode_gemv3(): # fmt: off @Ts.prim_func(private=True) - def before(lv575: T.Buffer((T.int64(4096), T.int64(1376)), "uint32"), lv576: T.Buffer((T.int64(4096), T.int64(344)), "float16"), lv574: T.Buffer((T.int64(1), T.int64(1), T.int64(11008)), "float16"), lv570: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16")): + def before(lv575: T.Tensor((T.int64(4096), T.int64(1376)), "uint32"), lv576: T.Tensor((T.int64(4096), T.int64(344)), "float16"), lv574: T.Tensor((T.int64(1), T.int64(1), T.int64(11008)), "float16"), lv570: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): p_output0_intermediate_1 = Ts.sblock_alloc_buffer((T.int64(4096), T.int64(11008)), "float16") @@ -299,7 +299,7 @@ def before(lv575: T.Buffer((T.int64(4096), T.int64(1376)), "uint32"), lv576: T.B p_output0_intermediate[v_ax0, v_ax1, v_ax2] = lv570[v_ax0, v_ax1, v_ax2] + var_NT_matmul_intermediate[v_ax0, v_ax1, v_ax2] @Ts.prim_func(private=True) - def expected(lv575: T.Buffer((T.int64(4096), T.int64(1376)), "uint32"), lv576: T.Buffer((T.int64(4096), T.int64(344)), "float16"), lv574: T.Buffer((T.int64(1), T.int64(1), T.int64(11008)), "float16"), lv570: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16")): + def expected(lv575: T.Tensor((T.int64(4096), T.int64(1376)), "uint32"), lv576: T.Tensor((T.int64(4096), T.int64(344)), "float16"), lv574: T.Tensor((T.int64(1), T.int64(1), T.int64(11008)), "float16"), lv570: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): var_NT_matmul_intermediate = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16") @@ -334,7 +334,7 @@ def expected(lv575: T.Buffer((T.int64(4096), T.int64(1376)), "uint32"), lv576: T def test_autogptq_decode_gemv(): # fmt: off @Ts.prim_func(private=True) - def func(lv9: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), lv10: T.Buffer((T.int64(32), T.int64(512)), "uint32"), lv11: T.Buffer((T.int64(32), T.int64(4096)), "float16"), lv12: T.Buffer((T.int64(4096),), "uint32"), lv8: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), lv1613: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16")): + def func(lv9: T.Tensor((T.int64(512), T.int64(4096)), "uint32"), lv10: T.Tensor((T.int64(32), T.int64(512)), "uint32"), lv11: T.Tensor((T.int64(32), T.int64(4096)), "float16"), lv12: T.Tensor((T.int64(4096),), "uint32"), lv8: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), lv1613: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): decode_intermediate = Ts.sblock_alloc_buffer((T.int64(4096), T.int64(4096)), "float16") @@ -373,11 +373,11 @@ def test_outer_reduction_adreno(): # fmt: off @Ts.prim_func(private=True) def before( - lv575: T.Buffer((1376, 4096), "uint32"), - lv576: T.Buffer((344, 4096), "float16"), - lv574: T.Buffer((1, 1, 11008), "float16"), - lv570: T.Buffer((1, 1, 4096), "float16"), - p_output0_intermediate: T.Buffer((1, 1, 4096), "float16"), + lv575: T.Tensor((1376, 4096), "uint32"), + lv576: T.Tensor((344, 4096), "float16"), + lv574: T.Tensor((1, 1, 11008), "float16"), + lv570: T.Tensor((1, 1, 4096), "float16"), + p_output0_intermediate: T.Tensor((1, 1, 4096), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -399,7 +399,7 @@ def before( p_output0_intermediate[v_ax0, v_ax1, v_ax2] = lv570[v_ax0, v_ax1, v_ax2] + var_matmul_intermediate[v_ax0, v_ax1, v_ax2] @Ts.prim_func(private=True) - def expected(lv575: T.Buffer((1376, 4096), "uint32"), lv576: T.Buffer((344, 4096), "float16"), lv574: T.Buffer((1, 1, 11008), "float16"), lv570: T.Buffer((1, 1, 4096), "float16"), p_output0_intermediate: T.Buffer((1, 1, 4096), "float16")): + def expected(lv575: T.Tensor((1376, 4096), "uint32"), lv576: T.Tensor((344, 4096), "float16"), lv574: T.Tensor((1, 1, 11008), "float16"), lv570: T.Tensor((1, 1, 4096), "float16"), p_output0_intermediate: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): p_output0_intermediate_1 = Ts.sblock_alloc_buffer((11008, 4096), "float16") @@ -436,7 +436,7 @@ def test_outer_reduction_adreno_dynamic(): v = T.dynamic("v") @Ts.prim_func(private=True) - def before(lv612: T.Buffer((T.int64(512), v), 'uint32'), lv613: T.Buffer((T.int64(128), v), 'float16'), lv1607: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Buffer((T.int64(1), T.int64(1), v))): + def before(lv612: T.Tensor((T.int64(512), v), 'uint32'), lv613: T.Tensor((T.int64(128), v), 'float16'), lv1607: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Tensor((T.int64(1), T.int64(1), v))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -466,7 +466,7 @@ def before(lv612: T.Buffer((T.int64(512), v), 'uint32'), lv613: T.Buffer((T.int6 v = T.dynamic("v") @Ts.prim_func(private=True) - def expected(lv612: T.Buffer((T.int64(512), v), 'uint32'), lv613: T.Buffer((T.int64(128), v), 'float16'), lv1607: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Buffer((T.int64(1), T.int64(1), v))): + def expected(lv612: T.Tensor((T.int64(512), v), 'uint32'), lv613: T.Tensor((T.int64(128), v), 'float16'), lv1607: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Tensor((T.int64(1), T.int64(1), v))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -503,7 +503,7 @@ def expected(lv612: T.Buffer((T.int64(512), v), 'uint32'), lv613: T.Buffer((T.in def test_blockized_gemv(): # fmt: off @Ts.prim_func(private=True) - def before(x: T.Buffer((1, 4096), "float16"), w: T.Buffer((8, 16384, 4096), "float16"), indptr: T.Buffer((2,), "int32"), o: T.Buffer((2, 16384), "float16")): + def before(x: T.Tensor((1, 4096), "float16"), w: T.Tensor((8, 16384, 4096), "float16"), indptr: T.Tensor((2,), "int32"), o: T.Tensor((2, 16384), "float16")): # with Ts.sblock("root"): for expert_id in T.thread_binding(2, thread="blockIdx.y"): with Ts.sblock("gemv_o"): @@ -522,7 +522,7 @@ def before(x: T.Buffer((1, 4096), "float16"), w: T.Buffer((8, 16384, 4096), "flo o[v_expert_id_o, vi_i] = o[v_expert_id_o, vi_i] + x[0, vj_i] * w[indptr[v_expert_id_o], vi_i, vj_i] @Ts.prim_func(private=True) - def expected(x: T.Buffer((1, 4096), "float16"), w: T.Buffer((8, 16384, 4096), "float16"), indptr: T.Buffer((2,), "int32"), o: T.Buffer((2, 16384), "float16")): + def expected(x: T.Tensor((1, 4096), "float16"), w: T.Tensor((8, 16384, 4096), "float16"), indptr: T.Tensor((2,), "int32"), o: T.Tensor((2, 16384), "float16")): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): for expert_id in T.thread_binding(2, thread="blockIdx.y"): @@ -555,8 +555,8 @@ def expected(x: T.Buffer((1, 4096), "float16"), w: T.Buffer((8, 16384, 4096), "f def test_func_to_skip(): @Ts.prim_func def before( - data_buf: T.Buffer((seq_len * T.int64(8),), "int32", align=8), # noqa: F821 - output_buf: T.Buffer((seq_len * T.int64(8),), "int32", align=8), # noqa: F821 + data_buf: T.Tensor((seq_len * T.int64(8),), "int32", align=8), # noqa: F821 + output_buf: T.Tensor((seq_len * T.int64(8),), "int32", align=8), # noqa: F821 seq_len: T.int64, ): with Ts.sblock("exclusive_scan_thrust"): diff --git a/tests/python/s_tir/dlight/test_gpu_conv.py b/tests/python/s_tir/dlight/test_gpu_conv.py index 9652121ce34f..7cfb1058cab8 100644 --- a/tests/python/s_tir/dlight/test_gpu_conv.py +++ b/tests/python/s_tir/dlight/test_gpu_conv.py @@ -28,9 +28,9 @@ def test_conv3d(): # fmt: off @Ts.prim_func(private=True) def before( - A: T.Buffer((14308, 3, 2, 14, 14), "float16"), - W: T.Buffer((1280, 3, 2, 14, 14), "float16"), - C: T.Buffer((14308, 1280, 1, 1, 1), "float16"), + A: T.Tensor((14308, 3, 2, 14, 14), "float16"), + W: T.Tensor((1280, 3, 2, 14, 14), "float16"), + C: T.Tensor((14308, 1280, 1, 1, 1), "float16"), ): pad_A = Ts.sblock_alloc_buffer((14308, 3, 2, 14, 14), "float16") for i0, i1, i2, i3, i4 in T.grid(14308, 3, 2, 14, 14): @@ -45,7 +45,7 @@ def before( C[v_nn, v_ff, v_yy, v_xx, v_zz] += pad_A[v_nn, v_rc, v_yy * 2 + v_ry, v_xx * 14 + v_rx, v_zz * 14 + v_rz]* W[v_ff, v_rc, v_ry, v_rx, v_rz] @Ts.prim_func(private=True) - def expected(A: T.Buffer((14308, 3, 2, 14, 14), "float16"), W: T.Buffer((1280, 3, 2, 14, 14), "float16"), C: T.Buffer((14308, 1280, 1, 1, 1), "float16")): + def expected(A: T.Tensor((14308, 3, 2, 14, 14), "float16"), W: T.Tensor((1280, 3, 2, 14, 14), "float16"), C: T.Tensor((14308, 1280, 1, 1, 1), "float16")): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): C_reindex_pad_local = Ts.sblock_alloc_buffer((1, 14336, 1280), "float16", scope="local") diff --git a/tests/python/s_tir/dlight/test_gpu_fallback.py b/tests/python/s_tir/dlight/test_gpu_fallback.py index 28978b089187..eec89b81f3b4 100644 --- a/tests/python/s_tir/dlight/test_gpu_fallback.py +++ b/tests/python/s_tir/dlight/test_gpu_fallback.py @@ -33,8 +33,8 @@ def test_fallback(): class Before: @Ts.prim_func def main( - A: T.Buffer((1, 32, 1, 128), "float16"), - C: T.Buffer((1, 1, 4096), "float16"), + A: T.Tensor((1, 32, 1, 128), "float16"), + C: T.Tensor((1, 1, 4096), "float16"), ): B = Ts.sblock_alloc_buffer((1, 1, 32, 128), "float16") for i, j, k, l in T.grid(1, 1, 32, 128): @@ -50,8 +50,8 @@ def main( class After: @Ts.prim_func def main( - A: T.Buffer((1, 32, 1, 128), "float16"), - C: T.Buffer((1, 1, 4096), "float16"), + A: T.Tensor((1, 32, 1, 128), "float16"), + C: T.Tensor((1, 1, 4096), "float16"), ): T.func_attr({"tirx.is_scheduled": True}) for ax0_fused_0 in T.thread_binding(4, thread="blockIdx.x"): @@ -74,7 +74,7 @@ def test_fallback_skips_zero_extent_spatial(): @I.ir_module class Module: @Ts.prim_func - def main(A: T.Buffer((4, 0), "float32"), B: T.Buffer((4, 0), "float32")): + def main(A: T.Tensor((4, 0), "float32"), B: T.Tensor((4, 0), "float32")): for i, j in T.grid(4, 0): with Ts.sblock("copy"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -91,7 +91,7 @@ def test_fallback_reduction(): @I.ir_module class Module: @Ts.prim_func - def main(A: T.Buffer((1, 6144), "float32"), B: T.Buffer((1,), "float32")): + def main(A: T.Tensor((1, 6144), "float32"), B: T.Tensor((1,), "float32")): for ax0, ax1 in T.grid(1, 6144): with Ts.sblock("block"): v0 = Ts.axis.spatial(1, ax0) @@ -105,7 +105,7 @@ def main(A: T.Buffer((1, 6144), "float32"), B: T.Buffer((1,), "float32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((1, 6144), "float32"), B: T.Buffer((1,), "float32")): + def main(A: T.Tensor((1, 6144), "float32"), B: T.Tensor((1,), "float32")): T.func_attr({"tirx.is_scheduled": True}) for ax0_fused_0 in T.thread_binding(T.int64(1), thread="blockIdx.x"): for ax0_fused_1 in T.thread_binding(T.int64(1024), thread="threadIdx.x"): @@ -142,10 +142,10 @@ def test_fallback_irregular_spatial(): @Ts.prim_func(private=True) def func( - pages: T.Buffer((num_total_pages, nlayer, nhead, page_size), "float16"), - page_table_indptr: T.Buffer((num_total_seqs_plus_1,), "int32"), - page_table_values: T.Buffer((npage,), "int32"), - values: T.Buffer((nlayer, nhead, seqlen), "float16"), + pages: T.Tensor((num_total_pages, nlayer, nhead, page_size), "float16"), + page_table_indptr: T.Tensor((num_total_seqs_plus_1,), "int32"), + page_table_values: T.Tensor((npage,), "int32"), + values: T.Tensor((nlayer, nhead, seqlen), "float16"), seq_id: T.int32, ): for l, h, pos in T.grid(nlayer, nhead, seqlen): @@ -168,7 +168,7 @@ def func( num_total_seqs_plus_1 = T.dynamic("num_total_seqs_plus_1", "int32") @Ts.prim_func(private=True) - def expected(pages: T.Buffer((num_total_pages, nlayer, nhead, page_size), 'float16'), page_table_indptr: T.Buffer((num_total_seqs_plus_1,), 'int32'), page_table_values: T.Buffer((npage,), 'int32'), values: T.Buffer((nlayer, nhead, seqlen), 'float16'), seq_id: T.int32): + def expected(pages: T.Tensor((num_total_pages, nlayer, nhead, page_size), 'float16'), page_table_indptr: T.Tensor((num_total_seqs_plus_1,), 'int32'), page_table_values: T.Tensor((npage,), 'int32'), values: T.Tensor((nlayer, nhead, seqlen), 'float16'), seq_id: T.int32): T.func_attr({"tirx.is_scheduled": True}) for ax0_ax1_ax2_fused_0 in T.thread_binding((nlayer * nhead * seqlen + 1023) // 1024, thread="blockIdx.x"): @@ -199,8 +199,8 @@ class Before: # using the `Target.current`. @Ts.prim_func def gpu_func( - A: T.Buffer((1, 32, 1, 128), "float16"), - C: T.Buffer((1, 1, 4096), "float16"), + A: T.Tensor((1, 32, 1, 128), "float16"), + C: T.Tensor((1, 1, 4096), "float16"), ): B = Ts.sblock_alloc_buffer((1, 1, 32, 128), "float16") for i, j, k, l in T.grid(1, 1, 32, 128): @@ -217,8 +217,8 @@ def gpu_func( # based on the annotation's target. @Ts.prim_func def cpu_func( - A: T.Buffer((1, 32, 1, 128), "float16"), - C: T.Buffer((1, 1, 4096), "float16"), + A: T.Tensor((1, 32, 1, 128), "float16"), + C: T.Tensor((1, 1, 4096), "float16"), ): T.func_attr({"target": T.target("llvm")}) B = Ts.sblock_alloc_buffer((1, 1, 32, 128), "float16") @@ -235,8 +235,8 @@ def cpu_func( class After: @Ts.prim_func def gpu_func( - A: T.Buffer((1, 32, 1, 128), "float16"), - C: T.Buffer((1, 1, 4096), "float16"), + A: T.Tensor((1, 32, 1, 128), "float16"), + C: T.Tensor((1, 1, 4096), "float16"), ): T.func_attr({"tirx.is_scheduled": True}) for ax0_fused_0 in T.thread_binding(4, thread="blockIdx.x"): @@ -249,8 +249,8 @@ def gpu_func( @Ts.prim_func def cpu_func( - A: T.Buffer((1, 32, 1, 128), "float16"), - C: T.Buffer((1, 1, 4096), "float16"), + A: T.Tensor((1, 32, 1, 128), "float16"), + C: T.Tensor((1, 1, 4096), "float16"), ): T.func_attr({"target": T.target("llvm")}) B = Ts.sblock_alloc_buffer((1, 1, 32, 128), "float16") @@ -275,7 +275,7 @@ def test_schedule_error_propagates_from_rule(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((128,), "float32"), C: T.Buffer((128,), "float32")): + def main(A: T.Tensor((128,), "float32"), C: T.Tensor((128,), "float32")): for i in range(128): with Ts.sblock("copy"): vi = Ts.axis.remap("S", [i]) diff --git a/tests/python/s_tir/dlight/test_gpu_gemv.py b/tests/python/s_tir/dlight/test_gpu_gemv.py index d5d35caa7e0d..c17fe5be3bc9 100644 --- a/tests/python/s_tir/dlight/test_gpu_gemv.py +++ b/tests/python/s_tir/dlight/test_gpu_gemv.py @@ -30,9 +30,9 @@ def test_gemv_rejects_composite_normalized_axis(): @Ts.prim_func(private=True) def before( - data: T.Buffer((1, 64, n), "float32"), # noqa: F821 - weight: T.Buffer((64, 1, 512), "float32"), - output: T.Buffer((1, 1, n * 256), "float32"), # noqa: F821 + data: T.Tensor((1, 64, n), "float32"), # noqa: F821 + weight: T.Tensor((64, 1, 512), "float32"), + output: T.Tensor((1, 1, n * 256), "float32"), # noqa: F821 n: T.int64, ): for w, rc, rw in T.grid(n * 256, 64, 512): @@ -57,7 +57,7 @@ def test_gemv_basic(): n = T.dynamic("n", "int32") @Ts.prim_func(private=True) - def before(lv1637: T.Buffer((1, 32, 1, 128), "float16"), lv1638: T.Buffer((1, 32, n, 128), 'float16'), lv1614: T.Buffer((1, 1, 1, n), 'float16'), var_compute_intermediate: T.Buffer((1, 32, 1, n))): + def before(lv1637: T.Tensor((1, 32, 1, 128), "float16"), lv1638: T.Tensor((1, 32, n, 128), 'float16'), lv1614: T.Tensor((1, 1, 1, n), 'float16'), var_compute_intermediate: T.Tensor((1, 32, 1, n))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -101,7 +101,7 @@ def before(lv1637: T.Buffer((1, 32, 1, 128), "float16"), lv1638: T.Buffer((1, 32 n = T.dynamic("n", "int32") @Ts.prim_func(private=True) - def expected(lv1637: T.Buffer((1, 32, 1, 128), "float16"), lv1638: T.Buffer((1, 32, n, 128), 'float16'), lv1614: T.Buffer((1, 1, 1, n), 'float16'), var_compute_intermediate: T.Buffer((1, 32, 1, n))): + def expected(lv1637: T.Tensor((1, 32, 1, 128), "float16"), lv1638: T.Tensor((1, 32, n, 128), 'float16'), lv1614: T.Tensor((1, 1, 1, n), 'float16'), var_compute_intermediate: T.Tensor((1, 32, 1, n))): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -207,7 +207,7 @@ def expected(lv1637: T.Buffer((1, 32, 1, 128), "float16"), lv1638: T.Buffer((1, def test_decode_gemv_256_threads(): # fmt: off @Ts.prim_func(private=True) - def before(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128), "float16"), lv1654: T.Buffer((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Buffer((1, 1, 22016), "float16")): + def before(lv571: T.Tensor((22016, 512), "uint32"), lv572: T.Tensor((22016, 128), "float16"), lv1654: T.Tensor((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Tensor((1, 1, 22016), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): p_output0_intermediate = Ts.sblock_alloc_buffer((22016, 4096), "float16") @@ -227,7 +227,7 @@ def before(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128) var_NT_matmul_intermediate[v_i0, v_i1, v_i2] = var_NT_matmul_intermediate[v_i0, v_i1, v_i2] + lv1654[v_i0, v_i1, v_k] * p_output0_intermediate[v_i2, v_k] @Ts.prim_func(private=True) - def expected(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128), "float16"), lv1654: T.Buffer((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Buffer((1, 1, 22016), "float16")): + def expected(lv571: T.Tensor((22016, 512), "uint32"), lv572: T.Tensor((22016, 128), "float16"), lv1654: T.Tensor((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Tensor((1, 1, 22016), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): var_NT_matmul_intermediate_rf_local = Ts.sblock_alloc_buffer((16, 1, 1, 22016), "float16", scope="local") @@ -303,7 +303,7 @@ def test_decode_gemv1(): # fmt: off @Ts.prim_func(private=True) - def before(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128), "float16"), lv1654: T.Buffer((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Buffer((1, 1, 22016), "float16")): + def before(lv571: T.Tensor((22016, 512), "uint32"), lv572: T.Tensor((22016, 128), "float16"), lv1654: T.Tensor((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Tensor((1, 1, 22016), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): p_output0_intermediate = Ts.sblock_alloc_buffer((22016, 4096), "float16") @@ -323,7 +323,7 @@ def before(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128) var_NT_matmul_intermediate[v_i0, v_i1, v_i2] = var_NT_matmul_intermediate[v_i0, v_i1, v_i2] + lv1654[v_i0, v_i1, v_k] * p_output0_intermediate[v_i2, v_k] @Ts.prim_func(private=True) - def expected(lv571: T.Buffer((22016, 512), "uint32"), lv572: T.Buffer((22016, 128), "float16"), lv1654: T.Buffer((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Buffer((1, 1, 22016), "float16")): + def expected(lv571: T.Tensor((22016, 512), "uint32"), lv572: T.Tensor((22016, 128), "float16"), lv1654: T.Tensor((1, 1, 4096), "float16"), var_NT_matmul_intermediate: T.Tensor((1, 1, 22016), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): var_NT_matmul_intermediate_rf_local = Ts.sblock_alloc_buffer((128, 1, 1, 22016), "float16", scope="local") @@ -411,7 +411,7 @@ def test_decode_gemv2(): # fmt: off @Ts.prim_func(private=True) - def before(lv771: T.Buffer((32000, 512), "uint32"), lv772: T.Buffer((32000, 128), "float16"), lv3216: T.Buffer((1, 1, 4096), "float16"), p_output0_intermediate: T.Buffer((1, 1, 32000), "float32")): + def before(lv771: T.Tensor((32000, 512), "uint32"), lv772: T.Tensor((32000, 128), "float16"), lv3216: T.Tensor((1, 1, 4096), "float16"), p_output0_intermediate: T.Tensor((1, 1, 32000), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): p_output0_intermediate_1 = Ts.sblock_alloc_buffer((32000, 4096), "float16") @@ -438,7 +438,7 @@ def before(lv771: T.Buffer((32000, 512), "uint32"), lv772: T.Buffer((32000, 128) p_output0_intermediate[v_i0, v_i1, v_i2] = T.Cast("float32", var_NT_matmul_intermediate[v_i0, v_i1, v_i2]) @Ts.prim_func(private=True) - def expected(lv771: T.Buffer((32000, 512), "uint32"), lv772: T.Buffer((32000, 128), "float16"), lv3216: T.Buffer((1, 1, 4096), "float16"), p_output0_intermediate: T.Buffer((1, 1, 32000), "float32")): + def expected(lv771: T.Tensor((32000, 512), "uint32"), lv772: T.Tensor((32000, 128), "float16"), lv3216: T.Tensor((1, 1, 4096), "float16"), p_output0_intermediate: T.Tensor((1, 1, 32000), "float32")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): var_NT_matmul_intermediate_local = Ts.sblock_alloc_buffer((1, 1, 32000), "float16", scope="local") @@ -534,7 +534,7 @@ def test_decode_gemv3(): # fmt: off @Ts.prim_func(private=True) - def before(lv575: T.Buffer((T.int64(4096), T.int64(1376)), "uint32"), lv576: T.Buffer((T.int64(4096), T.int64(344)), "float16"), lv574: T.Buffer((T.int64(1), T.int64(1), T.int64(11008)), "float16"), lv570: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16")): + def before(lv575: T.Tensor((T.int64(4096), T.int64(1376)), "uint32"), lv576: T.Tensor((T.int64(4096), T.int64(344)), "float16"), lv574: T.Tensor((T.int64(1), T.int64(1), T.int64(11008)), "float16"), lv570: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): p_output0_intermediate_1 = Ts.sblock_alloc_buffer((T.int64(4096), T.int64(11008)), "float16") @@ -561,7 +561,7 @@ def before(lv575: T.Buffer((T.int64(4096), T.int64(1376)), "uint32"), lv576: T.B p_output0_intermediate[v_ax0, v_ax1, v_ax2] = lv570[v_ax0, v_ax1, v_ax2] + var_NT_matmul_intermediate[v_ax0, v_ax1, v_ax2] @Ts.prim_func(private=True) - def expected(lv575: T.Buffer((T.int64(4096), T.int64(1376)), "uint32"), lv576: T.Buffer((T.int64(4096), T.int64(344)), "float16"), lv574: T.Buffer((T.int64(1), T.int64(1), T.int64(11008)), "float16"), lv570: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16")): + def expected(lv575: T.Tensor((T.int64(4096), T.int64(1376)), "uint32"), lv576: T.Tensor((T.int64(4096), T.int64(344)), "float16"), lv574: T.Tensor((T.int64(1), T.int64(1), T.int64(11008)), "float16"), lv570: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): var_NT_matmul_intermediate_local = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16", scope="local") @@ -657,7 +657,7 @@ def expected(lv575: T.Buffer((T.int64(4096), T.int64(1376)), "uint32"), lv576: T def test_autogptq_decode_gemv(): # fmt: off @Ts.prim_func(private=True) - def func(lv9: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), lv10: T.Buffer((T.int64(32), T.int64(512)), "uint32"), lv11: T.Buffer((T.int64(32), T.int64(4096)), "float16"), lv12: T.Buffer((T.int64(4096),), "uint32"), lv8: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), lv1613: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16")): + def func(lv9: T.Tensor((T.int64(512), T.int64(4096)), "uint32"), lv10: T.Tensor((T.int64(32), T.int64(512)), "uint32"), lv11: T.Tensor((T.int64(32), T.int64(4096)), "float16"), lv12: T.Tensor((T.int64(4096),), "uint32"), lv8: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), lv1613: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): decode_intermediate = Ts.sblock_alloc_buffer((T.int64(4096), T.int64(4096)), "float16") @@ -696,11 +696,11 @@ def test_outer_reduction_adreno(): # fmt: off @Ts.prim_func(private=True) def before( - lv575: T.Buffer((1376, 4096), "uint32"), - lv576: T.Buffer((344, 4096), "float16"), - lv574: T.Buffer((1, 1, 11008), "float16"), - lv570: T.Buffer((1, 1, 4096), "float16"), - p_output0_intermediate: T.Buffer((1, 1, 4096), "float16"), + lv575: T.Tensor((1376, 4096), "uint32"), + lv576: T.Tensor((344, 4096), "float16"), + lv574: T.Tensor((1, 1, 11008), "float16"), + lv570: T.Tensor((1, 1, 4096), "float16"), + p_output0_intermediate: T.Tensor((1, 1, 4096), "float16"), ): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -722,7 +722,7 @@ def before( p_output0_intermediate[v_ax0, v_ax1, v_ax2] = lv570[v_ax0, v_ax1, v_ax2] + var_matmul_intermediate[v_ax0, v_ax1, v_ax2] @Ts.prim_func(private=True) - def expected(lv575: T.Buffer((1376, 4096), "uint32"), lv576: T.Buffer((344, 4096), "float16"), lv574: T.Buffer((1, 1, 11008), "float16"), lv570: T.Buffer((1, 1, 4096), "float16"), p_output0_intermediate: T.Buffer((1, 1, 4096), "float16")): + def expected(lv575: T.Tensor((1376, 4096), "uint32"), lv576: T.Tensor((344, 4096), "float16"), lv574: T.Tensor((1, 1, 11008), "float16"), lv570: T.Tensor((1, 1, 4096), "float16"), p_output0_intermediate: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): var_matmul_intermediate_local = Ts.sblock_alloc_buffer((1, 1, 4096), "float16", scope="local") @@ -809,7 +809,7 @@ def test_outer_reduction_adreno_dynamic(): v = T.dynamic("v") @Ts.prim_func(private=True) - def before(lv612: T.Buffer((T.int64(512), v), 'uint32'), lv613: T.Buffer((T.int64(128), v), 'float16'), lv1607: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Buffer((T.int64(1), T.int64(1), v))): + def before(lv612: T.Tensor((T.int64(512), v), 'uint32'), lv613: T.Tensor((T.int64(128), v), 'float16'), lv1607: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Tensor((T.int64(1), T.int64(1), v))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -839,7 +839,7 @@ def before(lv612: T.Buffer((T.int64(512), v), 'uint32'), lv613: T.Buffer((T.int6 v = T.dynamic("v") @Ts.prim_func(private=True) - def expected(lv612: T.Buffer((T.int64(512), v), 'uint32'), lv613: T.Buffer((T.int64(128), v), 'float16'), lv1607: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Buffer((T.int64(1), T.int64(1), v))): + def expected(lv612: T.Tensor((T.int64(512), v), 'uint32'), lv613: T.Tensor((T.int64(128), v), 'float16'), lv1607: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), p_output0_intermediate: T.Tensor((T.int64(1), T.int64(1), v))): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -947,7 +947,7 @@ def expected(lv612: T.Buffer((T.int64(512), v), 'uint32'), lv613: T.Buffer((T.in def test_blockized_gemv(): # fmt: off @Ts.prim_func(private=True) - def before(x: T.Buffer((1, 4096), "float16"), w: T.Buffer((8, 16384, 4096), "float16"), indptr: T.Buffer((2,), "int32"), o: T.Buffer((2, 16384), "float16")): + def before(x: T.Tensor((1, 4096), "float16"), w: T.Tensor((8, 16384, 4096), "float16"), indptr: T.Tensor((2,), "int32"), o: T.Tensor((2, 16384), "float16")): # with Ts.sblock("root"): for expert_id in T.thread_binding(2, thread="blockIdx.y"): with Ts.sblock("gemv_o"): @@ -966,7 +966,7 @@ def before(x: T.Buffer((1, 4096), "float16"), w: T.Buffer((8, 16384, 4096), "flo o[v_expert_id_o, vi_i] = o[v_expert_id_o, vi_i] + x[0, vj_i] * w[indptr[v_expert_id_o], vi_i, vj_i] @Ts.prim_func(private=True) - def expected(x: T.Buffer((1, 4096), "float16"), w: T.Buffer((8, 16384, 4096), "float16"), indptr: T.Buffer((2,), "int32"), o: T.Buffer((2, 16384), "float16")): + def expected(x: T.Tensor((1, 4096), "float16"), w: T.Tensor((8, 16384, 4096), "float16"), indptr: T.Tensor((2,), "int32"), o: T.Tensor((2, 16384), "float16")): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): for expert_id in T.thread_binding(2, thread="blockIdx.y"): @@ -1049,8 +1049,8 @@ def expected(x: T.Buffer((1, 4096), "float16"), w: T.Buffer((8, 16384, 4096), "f def test_func_to_skip(): @Ts.prim_func def before( - data_buf: T.Buffer((seq_len * T.int64(8),), "int32", align=8), # noqa: F821 - output_buf: T.Buffer((seq_len * T.int64(8),), "int32", align=8), # noqa: F821 + data_buf: T.Tensor((seq_len * T.int64(8),), "int32", align=8), # noqa: F821 + output_buf: T.Tensor((seq_len * T.int64(8),), "int32", align=8), # noqa: F821 seq_len: T.int64, ): with Ts.sblock("exclusive_scan_thrust"): @@ -1083,9 +1083,9 @@ def test_gemv_cuda_target_without_max_shared_memory_per_block(): # fmt: off @Ts.prim_func(private=True) def before( - A: T.Buffer((1, 1, 1, 128), "float16"), - B: T.Buffer((1, 1, 64, 128), "float16"), - C: T.Buffer((1, 1, 1, 64), "float16"), + A: T.Tensor((1, 1, 1, 128), "float16"), + B: T.Tensor((1, 1, 64, 128), "float16"), + C: T.Tensor((1, 1, 1, 64), "float16"), ): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3, k in T.grid(1, 1, 1, 64, 128): @@ -1114,9 +1114,9 @@ def before( def test_gemv_rank_one_vector_input(): @Ts.prim_func(private=True) def before( - matrix: T.Buffer((2, 2), "float32"), - vector: T.Buffer((2,), "float32"), - output: T.Buffer((2,), "float32"), + matrix: T.Tensor((2, 2), "float32"), + vector: T.Tensor((2,), "float32"), + output: T.Tensor((2,), "float32"), ): T.func_attr({"tirx.noalias": True}) for i, k in T.grid(2, 2): @@ -1145,9 +1145,9 @@ def test_gemv_broadcast_epilogue(): @Ts.prim_func(private=True) def before( - A: T.Buffer((1, 32, 1, 128), "float16"), - B: T.Buffer((1, 32, n, 128), 'float16'), - C: T.Buffer((1, 32, 2, 3, n), 'float32'), + A: T.Tensor((1, 32, 1, 128), "float16"), + B: T.Tensor((1, 32, n, 128), 'float16'), + C: T.Tensor((1, 32, 2, 3, n), 'float32'), ): T.func_attr({"tirx.noalias": True}) 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 c857e83cd5d9..76d0dde1feb0 100644 --- a/tests/python/s_tir/dlight/test_gpu_general_reduction.py +++ b/tests/python/s_tir/dlight/test_gpu_general_reduction.py @@ -40,7 +40,7 @@ def _make_scalar_argmin(length): @I.ir_module class Before: @Ts.prim_func - def main(x: T.Buffer((T.int64(length),), "float32"), x_red: T.Buffer((), "int64")): + def main(x: T.Tensor((T.int64(length),), "float32"), x_red: T.Tensor((), "int64")): T.func_attr({"tirx.noalias": True}) x_red_temp_v0 = Ts.sblock_alloc_buffer((), "int64") x_red_temp_v1 = Ts.sblock_alloc_buffer(()) @@ -97,7 +97,7 @@ def test_softmax_1(): @I.ir_module class Before: @Ts.prim_func - def main(lv44: T.Buffer((T.int64(1), T.int64(32), n, m)), var_compute_intermediate: T.Buffer((T.int64(1), T.int64(32), n, m), 'float16')): + def main(lv44: T.Tensor((T.int64(1), T.int64(32), n, m)), var_compute_intermediate: T.Tensor((T.int64(1), T.int64(32), n, m), 'float16')): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -147,7 +147,7 @@ def main(lv44: T.Buffer((T.int64(1), T.int64(32), n, m)), var_compute_intermedia @I.ir_module class After: @Ts.prim_func - def main(lv44: T.Buffer((T.int64(1), T.int64(32), n, m)), var_compute_intermediate: T.Buffer((T.int64(1), T.int64(32), n, m), 'float16')): + def main(lv44: T.Tensor((T.int64(1), T.int64(32), n, m)), var_compute_intermediate: T.Tensor((T.int64(1), T.int64(32), n, m), 'float16')): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -199,7 +199,7 @@ def test_softmax_2(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((T.int64(1), T.int64(1), T.int64(32000)), "float32"), T_softmax_norm: T.Buffer((T.int64(1), T.int64(1), T.int64(32000)), "float32")): + def main(A: T.Tensor((T.int64(1), T.int64(1), T.int64(32000)), "float32"), T_softmax_norm: T.Tensor((T.int64(1), T.int64(1), T.int64(32000)), "float32")): # with Ts.sblock("root"): T_softmax_maxelem = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1))) T_softmax_exp = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(32000))) @@ -237,7 +237,7 @@ def main(A: T.Buffer((T.int64(1), T.int64(1), T.int64(32000)), "float32"), T_sof @I.ir_module class After: @Ts.prim_func - def main(A: T.Buffer((T.int64(1), T.int64(1), T.int64(32000)), "float32"), T_softmax_norm: T.Buffer((T.int64(1), T.int64(1), T.int64(32000)), "float32")): + def main(A: T.Tensor((T.int64(1), T.int64(1), T.int64(32000)), "float32"), T_softmax_norm: T.Tensor((T.int64(1), T.int64(1), T.int64(32000)), "float32")): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): T_softmax_maxelem_shared = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1)), scope="shared") @@ -284,7 +284,7 @@ def test_softmax_3(): @I.ir_module class Before: @Ts.prim_func - def main(input: T.Buffer((T.int64(1), T.int64(4), T.int64(32), T.int64(8192)), "float32"), T_softmax_norm: T.Buffer((T.int64(1), T.int64(4), T.int64(32), T.int64(8192)), "float32")): + def main(input: T.Tensor((T.int64(1), T.int64(4), T.int64(32), T.int64(8192)), "float32"), T_softmax_norm: T.Tensor((T.int64(1), T.int64(4), T.int64(32), T.int64(8192)), "float32")): # with Ts.sblock("root"): T_softmax_maxelem = Ts.sblock_alloc_buffer((T.int64(1), T.int64(4), T.int64(8192))) T_softmax_exp = Ts.sblock_alloc_buffer((T.int64(1), T.int64(4), T.int64(32), T.int64(8192))) @@ -322,7 +322,7 @@ def main(input: T.Buffer((T.int64(1), T.int64(4), T.int64(32), T.int64(8192)), " @I.ir_module class After: @Ts.prim_func - def main(input: T.Buffer((T.int64(1), T.int64(4), T.int64(32), T.int64(8192)), "float32"), T_softmax_norm: T.Buffer((T.int64(1), T.int64(4), T.int64(32), T.int64(8192)), "float32")): + def main(input: T.Tensor((T.int64(1), T.int64(4), T.int64(32), T.int64(8192)), "float32"), T_softmax_norm: T.Tensor((T.int64(1), T.int64(4), T.int64(32), T.int64(8192)), "float32")): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): T_softmax_maxelem_shared = Ts.sblock_alloc_buffer((T.int64(1), T.int64(4), T.int64(8192)), scope="shared") @@ -376,7 +376,7 @@ def test_layer_norm(): @I.ir_module class Before: @Ts.prim_func - def main(lv6: T.Buffer((T.int64(1), n, T.int64(2560))), weight1: T.Buffer((T.int64(2560),), "float32"), bias: T.Buffer((T.int64(2560),), "float32"), var_compute_intermediate: T.Buffer((T.int64(1), n, T.int64(2560)), 'float16')): + def main(lv6: T.Tensor((T.int64(1), n, T.int64(2560))), weight1: T.Tensor((T.int64(2560),), "float32"), bias: T.Tensor((T.int64(2560),), "float32"), var_compute_intermediate: T.Tensor((T.int64(1), n, T.int64(2560)), 'float16')): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -413,7 +413,7 @@ def main(lv6: T.Buffer((T.int64(1), n, T.int64(2560))), weight1: T.Buffer((T.int @I.ir_module class After: @Ts.prim_func - def main(lv6: T.Buffer((T.int64(1), n, T.int64(2560))), weight1: T.Buffer((T.int64(2560),), "float32"), bias: T.Buffer((T.int64(2560),), "float32"), var_compute_intermediate: T.Buffer((T.int64(1), n, T.int64(2560)), 'float16')): + def main(lv6: T.Tensor((T.int64(1), n, T.int64(2560))), weight1: T.Tensor((T.int64(2560),), "float32"), bias: T.Tensor((T.int64(2560),), "float32"), var_compute_intermediate: T.Tensor((T.int64(1), n, T.int64(2560)), 'float16')): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -454,7 +454,7 @@ def test_rms_norm(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((T.int64(1), n, T.int64(4096)), 'float16'), B: T.Buffer((T.int64(4096),), "float16"), rms_norm_1: T.Buffer((T.int64(1), n, T.int64(4096)), 'float16')): + def main(A: T.Tensor((T.int64(1), n, T.int64(4096)), 'float16'), B: T.Tensor((T.int64(4096),), "float16"), rms_norm_1: T.Tensor((T.int64(1), n, T.int64(4096)), 'float16')): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -479,7 +479,7 @@ def main(A: T.Buffer((T.int64(1), n, T.int64(4096)), 'float16'), B: T.Buffer((T. @I.ir_module class After: @Ts.prim_func - def main(A: T.Buffer((T.int64(1), n, T.int64(4096)), 'float16'), B: T.Buffer((T.int64(4096),), "float16"), rms_norm_1: T.Buffer((T.int64(1), n, T.int64(4096)), 'float16')): + def main(A: T.Tensor((T.int64(1), n, T.int64(4096)), 'float16'), B: T.Tensor((T.int64(4096),), "float16"), rms_norm_1: T.Tensor((T.int64(1), n, T.int64(4096)), 'float16')): T.func_attr({"op_pattern": 4, "tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -513,7 +513,7 @@ def test_group_norm(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((1, 2048), "float32"), B: T.Buffer((2048,), "float32"), C: T.Buffer((2048,), "float32"), T_reshape: T.Buffer((1, 2048), "float32")): + def main(A: T.Tensor((1, 2048), "float32"), B: T.Tensor((2048,), "float32"), C: T.Tensor((2048,), "float32"), T_reshape: T.Tensor((1, 2048), "float32")): T.func_attr({"tirx.noalias": True}) T_reshape_1 = Ts.sblock_alloc_buffer((1, 32, 64)) A_red_temp_v0 = Ts.sblock_alloc_buffer((1, 32)) @@ -567,7 +567,7 @@ def main(A: T.Buffer((1, 2048), "float32"), B: T.Buffer((2048,), "float32"), C: @I.ir_module class After: @Ts.prim_func - def main(A: T.Buffer((1, 2048), "float32"), B: T.Buffer((2048,), "float32"), C: T.Buffer((2048,), "float32"), T_reshape: T.Buffer((1, 2048), "float32")): + def main(A: T.Tensor((1, 2048), "float32"), B: T.Tensor((2048,), "float32"), C: T.Tensor((2048,), "float32"), T_reshape: T.Tensor((1, 2048), "float32")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): A_red_temp_v0_shared = Ts.sblock_alloc_buffer((1, 32), scope="shared") @@ -609,8 +609,8 @@ def test_logsumexp(): class Before: @Ts.prim_func def compute_lse( - A: T.Buffer((batch_size, vocab_size), dtype="float32"), - blocked_lse: T.Buffer((batch_size, num_chunks), dtype="float32"), + A: T.Tensor((batch_size, vocab_size), dtype="float32"), + blocked_lse: T.Tensor((batch_size, num_chunks), dtype="float32"), ): T.func_attr({"tirx.noalias": True}) @@ -658,7 +658,7 @@ def compute_lse( class After: @Ts.prim_func def compute_lse( - A: T.Buffer((batch_size, vocab_size)), blocked_lse: T.Buffer((batch_size, num_chunks)) + A: T.Tensor((batch_size, vocab_size)), blocked_lse: T.Tensor((batch_size, num_chunks)) ): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) 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 21d907eaeddc..fb05d0e450cf 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 @@ -33,7 +33,7 @@ def test_batch_decode_gemv(): 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"), lv807: T.Buffer((batch_size, T.int64(1), T.int64(28672)), 'float16'), NT_matmul_intermediate: T.Buffer((batch_size, T.int64(1), T.int64(4096)), 'float16')): + def before(lv429: T.Tensor((T.int64(4096), T.int64(3584)), "uint32"), lv430: T.Tensor((T.int64(4096), T.int64(896)), "float16"), lv807: T.Tensor((batch_size, T.int64(1), T.int64(28672)), 'float16'), NT_matmul_intermediate: T.Tensor((batch_size, T.int64(1), T.int64(4096)), 'float16')): T.func_attr({"tirx.noalias": True, "tirx.HoistIfThenElseExprWithBlock": 1}) # with Ts.sblock("root"): @@ -63,7 +63,7 @@ def before(lv429: T.Buffer((T.int64(4096), T.int64(3584)), "uint32"), lv430: T.B 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"), lv807: T.Buffer((batch_size, T.int64(1), T.int64(28672)), 'float16'), NT_matmul_intermediate: T.Buffer((batch_size, T.int64(1), T.int64(4096)), 'float16')): + def expected(lv429: T.Tensor((T.int64(4096), T.int64(3584)), "uint32"), lv430: T.Tensor((T.int64(4096), T.int64(896)), "float16"), lv807: T.Tensor((batch_size, T.int64(1), T.int64(28672)), 'float16'), NT_matmul_intermediate: T.Tensor((batch_size, T.int64(1), T.int64(4096)), 'float16')): T.func_attr({"tirx.HoistIfThenElseExprWithBlock": 1, "tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -161,7 +161,7 @@ def test_batch_gemv(): batch_size = T.dynamic("batch_size") @Ts.prim_func(private=True) - def before(A: T.Buffer((batch_size, T.int64(1), T.int64(K)), 'float16'), B: T.Buffer((T.int64(N), T.int64(K)), "float16"), NT_matmul: T.Buffer((batch_size, T.int64(1), T.int64(N)), 'float16')): + def before(A: T.Tensor((batch_size, T.int64(1), T.int64(K)), 'float16'), B: T.Tensor((T.int64(N), T.int64(K)), "float16"), NT_matmul: T.Tensor((batch_size, T.int64(1), T.int64(N)), 'float16')): T.func_attr({"tirx.noalias": True, "tirx.HoistIfThenElseExprWithBlock": 1}) # with Ts.sblock("root"): @@ -177,7 +177,7 @@ def before(A: T.Buffer((batch_size, T.int64(1), T.int64(K)), 'float16'), B: T.Bu batch_size = T.dynamic("batch_size") @Ts.prim_func(private=True) - def expected(A: T.Buffer((batch_size, T.int64(1), T.int64(4096)), 'float16'), B: T.Buffer((T.int64(4096), T.int64(4096)), "float16"), NT_matmul: T.Buffer((batch_size, T.int64(1), T.int64(4096)), 'float16')): + def expected(A: T.Tensor((batch_size, T.int64(1), T.int64(4096)), 'float16'), B: T.Tensor((T.int64(4096), T.int64(4096)), "float16"), NT_matmul: T.Tensor((batch_size, T.int64(1), T.int64(4096)), 'float16')): T.func_attr({"tirx.HoistIfThenElseExprWithBlock": 1, "tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -262,7 +262,7 @@ def test_reduction_symbolic_var(): kv_seq_len = T.dynamic("kv_seq_len") @Ts.prim_func(private=True) - def before(A: T.Buffer((T.int64(1), T.int64(32), T.int64(1), kv_seq_len)), B: T.Buffer((T.int64(1), T.int64(32), kv_seq_len, T.int64(128))), matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float32")): + def before(A: T.Tensor((T.int64(1), T.int64(32), T.int64(1), kv_seq_len)), B: T.Tensor((T.int64(1), T.int64(32), kv_seq_len, T.int64(128))), matmul: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -286,9 +286,9 @@ def test_small_spatial_axis(): @Ts.prim_func(private=True) def func( - A: T.Buffer((batch_size, T.int64(4096)), "float16"), - B: T.Buffer((T.int64(8), T.int64(4096)), "float16"), - C: T.Buffer((batch_size, T.int64(8)), "float16"), + A: T.Tensor((batch_size, T.int64(4096)), "float16"), + B: T.Tensor((T.int64(8), T.int64(4096)), "float16"), + C: T.Tensor((batch_size, T.int64(8)), "float16"), ): T.func_attr({"tirx.noalias": True}) @@ -305,7 +305,7 @@ def func( batch_size = T.dynamic("batch_size") @Ts.prim_func(private=True) - def expected(A: T.Buffer((batch_size, T.int64(4096)), 'float16'), B: T.Buffer((T.int64(8), T.int64(4096)), "float16"), C: T.Buffer((batch_size, T.int64(8)), 'float16')): + def expected(A: T.Tensor((batch_size, T.int64(4096)), 'float16'), B: T.Tensor((T.int64(8), T.int64(4096)), "float16"), C: T.Tensor((batch_size, T.int64(8)), 'float16')): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -397,10 +397,10 @@ def test_outer_reduction(): @Ts.prim_func(private=True) def before( - B0: T.Buffer((512, 6144), "uint32"), - B1: T.Buffer((128, 6144), "float16"), - A: T.Buffer((batch_size, 1, 4096), 'float16'), - C: T.Buffer((batch_size, 1, 6144), 'float16') + B0: T.Tensor((512, 6144), "uint32"), + B1: T.Tensor((128, 6144), "float16"), + A: T.Tensor((batch_size, 1, 4096), 'float16'), + C: T.Tensor((batch_size, 1, 6144), 'float16') ): compute = Ts.sblock_alloc_buffer((4096, 6144), "float16") @@ -423,7 +423,7 @@ def before( 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"), A: T.Buffer((batch_size, 1, 4096), 'float16'), C: T.Buffer((batch_size, 1, 6144), 'float16')): + def expected(B0: T.Tensor((512, 6144), "uint32"), B1: T.Tensor((128, 6144), "float16"), A: T.Tensor((batch_size, 1, 4096), 'float16'), C: T.Tensor((batch_size, 1, 6144), 'float16')): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): @@ -542,7 +542,7 @@ def test_low_batch_gemv_cuda_target_without_max_shared_memory_per_block(): batch_size = T.dynamic("batch_size") @Ts.prim_func(private=True) - def before(A: T.Buffer((batch_size, T.int64(1), T.int64(128)), 'float16'), B: T.Buffer((T.int64(128), T.int64(128)), "float16"), C: T.Buffer((batch_size, T.int64(1), T.int64(128)), 'float16')): + def before(A: T.Tensor((batch_size, T.int64(1), T.int64(128)), 'float16'), B: T.Tensor((T.int64(128), T.int64(128)), "float16"), C: T.Tensor((batch_size, T.int64(1), T.int64(128)), 'float16')): T.func_attr({"tir.noalias": True}) for i0, i1, i2, k in T.grid(batch_size, T.int64(1), T.int64(128), T.int64(128)): @@ -569,9 +569,9 @@ def test_low_batch_gemv_rejects_non_einsum_buffer_access(): @Ts.prim_func(private=True) def before( - A: T.Buffer((batch_size, 8), "float16"), - B: T.Buffer((4, batch_size + 8), "float16"), - C: T.Buffer((batch_size, 4), "float16"), + A: T.Tensor((batch_size, 8), "float16"), + B: T.Tensor((4, batch_size + 8), "float16"), + C: T.Tensor((batch_size, 4), "float16"), ): for i, j, k in T.grid(batch_size, 4, 8): with Ts.sblock("attention_score"): @@ -593,9 +593,9 @@ def test_low_batch_gemv_broadcast_epilogue(): @Ts.prim_func(private=True) def before( - A: T.Buffer((T.int64(1), batch_size, T.int64(1), T.int64(128)), 'float16'), - B: T.Buffer((T.int64(128), T.int64(128)), "float16"), - C: T.Buffer((T.int64(1), batch_size, T.int64(2), T.int64(3), T.int64(128)), 'float32'), + A: T.Tensor((T.int64(1), batch_size, T.int64(1), T.int64(128)), 'float16'), + B: T.Tensor((T.int64(128), T.int64(128)), "float16"), + C: T.Tensor((T.int64(1), batch_size, T.int64(2), T.int64(3), T.int64(128)), 'float32'), ): T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/s_tir/dlight/test_gpu_matmul.py b/tests/python/s_tir/dlight/test_gpu_matmul.py index 951798701897..6063ea3f8531 100644 --- a/tests/python/s_tir/dlight/test_gpu_matmul.py +++ b/tests/python/s_tir/dlight/test_gpu_matmul.py @@ -29,7 +29,7 @@ def test_matmul(): m = T.dynamic("m") @Ts.prim_func(private=True) - def before(inp0: T.Buffer((T.int64(1), m, T.int64(4096))), inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), matmul: T.Buffer((T.int64(1), m, T.int64(4096)))): + def before(inp0: T.Tensor((T.int64(1), m, T.int64(4096))), inp1: T.Tensor((T.int64(4096), T.int64(4096)), "float32"), matmul: T.Tensor((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)): with Ts.sblock("matmul"): @@ -41,7 +41,7 @@ def before(inp0: T.Buffer((T.int64(1), m, T.int64(4096))), inp1: T.Buffer((T.int m = T.dynamic("m") @Ts.prim_func(private=True) - def expected(inp0: T.Buffer((T.int64(1), m, T.int64(4096))), inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), matmul: T.Buffer((T.int64(1), m, T.int64(4096)))): + def expected(inp0: T.Tensor((T.int64(1), m, T.int64(4096))), inp1: T.Tensor((T.int64(4096), T.int64(4096)), "float32"), matmul: T.Tensor((T.int64(1), m, T.int64(4096)))): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): @@ -121,7 +121,7 @@ def test_matmul_int32(): m = T.dynamic("m", "int32") @Ts.prim_func(private=True) - def func(inp0: T.Buffer((1, m, 4096)), inp1: T.Buffer((4096, 4096), "float32"), matmul: T.Buffer((1, m, 4096))): + def func(inp0: T.Tensor((1, m, 4096)), inp1: T.Tensor((4096, 4096), "float32"), matmul: T.Tensor((1, m, 4096))): for i0, i1, i2, k in T.grid(1, m, 4096, 4096): with Ts.sblock("matmul"): @@ -133,7 +133,7 @@ def func(inp0: T.Buffer((1, m, 4096)), inp1: T.Buffer((4096, 4096), "float32"), m = T.dynamic("m", "int32") @Ts.prim_func(private=True) - def expected(inp0: T.Buffer((1, m, 4096)), inp1: T.Buffer((4096, 4096), "float32"), matmul: T.Buffer((1, m, 4096))): + def expected(inp0: T.Tensor((1, m, 4096)), inp1: T.Tensor((4096, 4096), "float32"), matmul: T.Tensor((1, m, 4096))): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): @@ -211,7 +211,7 @@ def expected(inp0: T.Buffer((1, m, 4096)), inp1: T.Buffer((4096, 4096), "float32 def test_fused_matmul(): # fmt: off @Ts.prim_func(private=True) - def before(W: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), S: T.Buffer((T.int64(128), T.int64(4096)), "uint32"), A: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32"), C: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32"), Out: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32")): + def before(W: T.Tensor((T.int64(512), T.int64(4096)), "uint32"), S: T.Tensor((T.int64(128), T.int64(4096)), "uint32"), A: T.Tensor((T.int64(1), T.int64(32), T.int64(4096)), "float32"), C: T.Tensor((T.int64(1), T.int64(32), T.int64(4096)), "float32"), Out: T.Tensor((T.int64(1), T.int64(32), T.int64(4096)), "float32")): var_decode_intermediate = Ts.sblock_alloc_buffer((T.int64(4096), T.int64(4096))) var_matmul_intermediate = Ts.sblock_alloc_buffer((T.int64(1), T.int64(32), T.int64(4096))) for i, j in T.grid(T.int64(4096), T.int64(4096)): @@ -236,7 +236,7 @@ def before(W: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), S: T.Buffer((T. Out[v_ax0, v_ax1, v_ax2] = C[v_ax0, v_ax1, v_ax2] + var_matmul_intermediate[v_ax0, v_ax1, v_ax2] @Ts.prim_func(private=True) - def expected(W: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), S: T.Buffer((T.int64(128), T.int64(4096)), "uint32"), A: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32"), C: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32"), Out: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32")): + def expected(W: T.Tensor((T.int64(512), T.int64(4096)), "uint32"), S: T.Tensor((T.int64(128), T.int64(4096)), "uint32"), A: T.Tensor((T.int64(1), T.int64(32), T.int64(4096)), "float32"), C: T.Tensor((T.int64(1), T.int64(32), T.int64(4096)), "float32"), Out: T.Tensor((T.int64(1), T.int64(32), T.int64(4096)), "float32")): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): var_matmul_intermediate_reindex_local = Ts.sblock_alloc_buffer((T.int64(1), T.int64(32), T.int64(4096)), scope="local") @@ -313,7 +313,7 @@ def expected(W: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), S: T.Buffer(( def test_skip_gemv(): # fmt: off @Ts.prim_func(private=True) - def before(W: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), S: T.Buffer((T.int64(128), T.int64(4096)), "uint32"), A: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float32"), C: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float32"), Out: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float32")): + def before(W: T.Tensor((T.int64(512), T.int64(4096)), "uint32"), S: T.Tensor((T.int64(128), T.int64(4096)), "uint32"), A: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float32"), C: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float32"), Out: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float32")): T.func_attr({"tirx.noalias": True}) var_decode_intermediate = Ts.sblock_alloc_buffer((T.int64(4096), T.int64(4096))) var_matmul_intermediate = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1), T.int64(4096))) @@ -353,7 +353,7 @@ def test_output_fp32(): 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"), lv48: T.Buffer((T.int64(1), n, T.int64(4096)), 'float16'), lv13_1: T.Buffer((T.int64(4096),), "float16"), lv3: T.Buffer((T.int64(1), n, T.int64(4096)), 'float16'), p_output0_intermediate: T.Buffer((T.int64(1), n, T.int64(4096)), 'float16')): + def before(lv13: T.Tensor((T.int64(4096), T.int64(512)), "uint32"), lv14: T.Tensor((T.int64(4096), T.int64(128)), "float16"), lv48: T.Tensor((T.int64(1), n, T.int64(4096)), 'float16'), lv13_1: T.Tensor((T.int64(4096),), "float16"), lv3: T.Tensor((T.int64(1), n, T.int64(4096)), 'float16'), p_output0_intermediate: T.Tensor((T.int64(1), n, T.int64(4096)), 'float16')): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -404,7 +404,7 @@ def before(lv13: T.Buffer((T.int64(4096), T.int64(512)), "uint32"), lv14: T.Buff 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"), lv48: T.Buffer((T.int64(1), n, T.int64(4096)), 'float16'), lv13_1: T.Buffer((T.int64(4096),), "float16"), lv3: T.Buffer((T.int64(1), n, T.int64(4096)), 'float16'), p_output0_intermediate: T.Buffer((T.int64(1), n, T.int64(4096)), 'float16')): + def expected(lv13: T.Tensor((T.int64(4096), T.int64(512)), "uint32"), lv14: T.Tensor((T.int64(4096), T.int64(128)), "float16"), lv48: T.Tensor((T.int64(1), n, T.int64(4096)), 'float16'), lv13_1: T.Tensor((T.int64(4096),), "float16"), lv3: T.Tensor((T.int64(1), n, T.int64(4096)), 'float16'), p_output0_intermediate: T.Tensor((T.int64(1), n, T.int64(4096)), 'float16')): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -485,7 +485,7 @@ def test_inline_consumer_chain(): n = T.dynamic("n") @Ts.prim_func(private=True) - def before(lv26: T.Buffer((n, T.int64(2048)), 'float16'), lv9: T.Buffer((T.int64(2048), T.int64(2048)), "float16"), lv52: T.Buffer((T.int64(1), n, T.int64(2048))), var_T_multiply_intermediate: T.Buffer((n, T.int64(2048)), 'float16')): + def before(lv26: T.Tensor((n, T.int64(2048)), 'float16'), lv9: T.Tensor((T.int64(2048), T.int64(2048)), "float16"), lv52: T.Tensor((T.int64(1), n, T.int64(2048))), var_T_multiply_intermediate: T.Tensor((n, T.int64(2048)), 'float16')): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -536,7 +536,7 @@ def before(lv26: T.Buffer((n, T.int64(2048)), 'float16'), lv9: T.Buffer((T.int64 n = T.dynamic("n") @Ts.prim_func(private=True) - def expected(lv26: T.Buffer((n, T.int64(2048)), 'float16'), lv9: T.Buffer((T.int64(2048), T.int64(2048)), "float16"), lv52: T.Buffer((T.int64(1), n, T.int64(2048))), var_T_multiply_intermediate: T.Buffer((n, T.int64(2048)), 'float16')): + def expected(lv26: T.Tensor((n, T.int64(2048)), 'float16'), lv9: T.Tensor((T.int64(2048), T.int64(2048)), "float16"), lv52: T.Tensor((T.int64(1), n, T.int64(2048))), var_T_multiply_intermediate: T.Tensor((n, T.int64(2048)), 'float16')): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -617,7 +617,7 @@ def test_matmul_android(): m = T.dynamic("m") @Ts.prim_func(private=True) - def before(inp0: T.Buffer((T.int64(1), m, T.int64(4096))), inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), matmul: T.Buffer((T.int64(1), m, T.int64(4096)))): + def before(inp0: T.Tensor((T.int64(1), m, T.int64(4096))), inp1: T.Tensor((T.int64(4096), T.int64(4096)), "float32"), matmul: T.Tensor((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)): with Ts.sblock("matmul"): @@ -629,7 +629,7 @@ def before(inp0: T.Buffer((T.int64(1), m, T.int64(4096))), inp1: T.Buffer((T.int m = T.dynamic("m") @Ts.prim_func(private=True) - def expected(inp0: T.Buffer((T.int64(1), m, T.int64(4096))), inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), matmul: T.Buffer((T.int64(1), m, T.int64(4096)))): + def expected(inp0: T.Tensor((T.int64(1), m, T.int64(4096))), inp1: T.Tensor((T.int64(4096), T.int64(4096)), "float32"), matmul: T.Tensor((T.int64(1), m, T.int64(4096)))): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): @@ -710,7 +710,7 @@ def test_fused_dequant_matmul_android(): 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"), rms_norm130: T.Buffer((T.int64(1), seq_len, T.int64(4096)), 'float16'), transformer_h_0_attn_c_attn_bias3: T.Buffer((T.int64(12288),), "float16"), T_add_intermediate_intermediate: T.Buffer((T.int64(1), seq_len, T.int64(12288)), 'float16')): + def before(lv452: T.Tensor((T.int64(512), T.int64(12288)), "uint32"), lv453: T.Tensor((T.int64(128), T.int64(12288)), "float16"), rms_norm130: T.Tensor((T.int64(1), seq_len, T.int64(4096)), 'float16'), transformer_h_0_attn_c_attn_bias3: T.Tensor((T.int64(12288),), "float16"), T_add_intermediate_intermediate: T.Tensor((T.int64(1), seq_len, T.int64(12288)), 'float16')): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -747,7 +747,7 @@ def before(lv452: T.Buffer((T.int64(512), T.int64(12288)), "uint32"), lv453: T.B 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"), rms_norm130: T.Buffer((T.int64(1), seq_len, T.int64(4096)), 'float16'), transformer_h_0_attn_c_attn_bias3: T.Buffer((T.int64(12288),), "float16"), T_add_intermediate_intermediate: T.Buffer((T.int64(1), seq_len, T.int64(12288)), 'float16')): + def expected(lv452: T.Tensor((T.int64(512), T.int64(12288)), "uint32"), lv453: T.Tensor((T.int64(128), T.int64(12288)), "float16"), rms_norm130: T.Tensor((T.int64(1), seq_len, T.int64(4096)), 'float16'), transformer_h_0_attn_c_attn_bias3: T.Tensor((T.int64(12288),), "float16"), T_add_intermediate_intermediate: T.Tensor((T.int64(1), seq_len, T.int64(12288)), 'float16')): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # 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 25fa04845dcc..a59ec2c817e3 100644 --- a/tests/python/s_tir/dlight/test_gpu_matmul_tensorize.py +++ b/tests/python/s_tir/dlight/test_gpu_matmul_tensorize.py @@ -28,7 +28,7 @@ def test_matmul_tensorize(): # fmt: off @Ts.prim_func(private=True) - def before(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float16"), compute: T.Buffer((256, 256), "float16")): + def before(X: T.Tensor((256, 256), "float16"), W: T.Tensor((256, 256), "float16"), compute: T.Tensor((256, 256), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i, j, k in T.grid(256, 256, 256): @@ -62,7 +62,7 @@ def before(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float16" 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")): + def expected(X: T.Tensor((256, 256), "float16"), W: T.Tensor((256, 256), "float16"), compute: T.Tensor((256, 256), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): X_reindex_shared_dyn = Ts.sblock_alloc_buffer((1, 256, 256), "float16", scope="shared.dyn") @@ -190,7 +190,7 @@ def test_matmul_tensorize_too_small(): m = T.dynamic("m", "int32") @Ts.prim_func(private=True) - def before(X: T.Buffer((m, 256), 'float16'), W: T.Buffer((15, 256), "float16"), compute: T.Buffer((m, 15))): + def before(X: T.Tensor((m, 256), 'float16'), W: T.Tensor((15, 256), "float16"), compute: T.Tensor((m, 15))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -206,7 +206,7 @@ def before(X: T.Buffer((m, 256), 'float16'), W: T.Buffer((15, 256), "float16"), m = T.dynamic("m", "int32") @Ts.prim_func(private=True) - def expected(X: T.Buffer((m, 256), 'float16'), W: T.Buffer((15, 256), "float16"), compute: T.Buffer((m, 15))): + def expected(X: T.Tensor((m, 256), 'float16'), W: T.Tensor((15, 256), "float16"), compute: T.Tensor((m, 15))): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -286,7 +286,7 @@ def test_matmul_tensorize_epilogue(): 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"), lv42: T.Buffer((T.int32(1), n, T.int32(2048)), 'float16'), lv3: T.Buffer((T.int32(1), n, T.int32(4096)), 'float16'), p_output0_intermediate: T.Buffer((T.int32(1), n, T.int32(4096)), 'float16')): + def before(lv686: T.Tensor((T.int32(4096), T.int32(256)), "uint32"), lv687: T.Tensor((T.int32(4096), T.int32(64)), "float16"), lv42: T.Tensor((T.int32(1), n, T.int32(2048)), 'float16'), lv3: T.Tensor((T.int32(1), n, T.int32(4096)), 'float16'), p_output0_intermediate: T.Tensor((T.int32(1), n, T.int32(4096)), 'float16')): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -343,7 +343,7 @@ def before(lv686: T.Buffer((T.int32(4096), T.int32(256)), "uint32"), lv687: T.Bu 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"), lv42: T.Buffer((1, n, 2048), 'float16'), lv3: T.Buffer((1, n, 4096), 'float16'), p_output0_intermediate: T.Buffer((1, n, 4096), 'float16')): + def expected(lv686: T.Tensor((4096, 256), "uint32"), lv687: T.Tensor((4096, 64), "float16"), lv42: T.Tensor((1, n, 2048), 'float16'), lv3: T.Tensor((1, n, 4096), 'float16'), p_output0_intermediate: T.Tensor((1, n, 4096), 'float16')): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -470,7 +470,7 @@ def expected(lv686: T.Buffer((4096, 256), "uint32"), lv687: T.Buffer((4096, 64), def test_matmul_int8_tensorize(): # fmt: off @Ts.prim_func(private=True) - def before(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), compute: T.Buffer((256, 256), "int32")): + def before(X: T.Tensor((256, 256), "int8"), W: T.Tensor((256, 256), "int8"), compute: T.Tensor((256, 256), "int32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): for i, j, r in T.grid(256, 256, 256): @@ -504,7 +504,7 @@ def before(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), com 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")): + def expected(X: T.Tensor((256, 256), "int8"), W: T.Tensor((256, 256), "int8"), compute: T.Tensor((256, 256), "int32")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): X_reindex_shared_dyn = Ts.sblock_alloc_buffer((1, 256, 256), "int8", scope="shared.dyn") @@ -631,7 +631,7 @@ def test_matmul_int8_tensorize_3d2d_dyn(): m = T.dynamic("m", "int32") @Ts.prim_func(private=True) - def before(A: T.Buffer((1, m, 22016), 'int8'), B: T.Buffer((4096, 22016), "int8"), matmul_1: T.Buffer((1, m, 4096), 'int32')): + def before(A: T.Tensor((1, m, 22016), 'int8'), B: T.Tensor((4096, 22016), "int8"), matmul_1: T.Tensor((1, m, 4096), 'int32')): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -667,7 +667,7 @@ def before(A: T.Buffer((1, m, 22016), 'int8'), B: T.Buffer((4096, 22016), "int8" C_3_s1 = T.dynamic("C_3_s1", "int32") @Ts.prim_func(private=True) - def expected(A: T.Buffer((1, m, 22016), 'int8'), B: T.Buffer((4096, 22016), "int8"), matmul_1: T.Buffer((1, m, 4096), 'int32')): + def expected(A: T.Tensor((1, m, 22016), 'int8'), B: T.Tensor((4096, 22016), "int8"), matmul_1: T.Tensor((1, m, 4096), 'int32')): T.func_attr({"op_pattern": 4, "tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -797,9 +797,9 @@ def test_matmul_metal(): @Ts.prim_func(private=True) def before( - A: T.Buffer((batch_size, 1, 4096), 'float16'), - B: T.Buffer((28672, 4096), "float16"), - C: T.Buffer((batch_size, 1, 28672), 'float16'), + A: T.Tensor((batch_size, 1, 4096), 'float16'), + B: T.Tensor((28672, 4096), "float16"), + C: T.Tensor((batch_size, 1, 28672), 'float16'), ): for i0, i1, i2, k in T.grid(batch_size, 1, 28672, 4096): @@ -833,7 +833,7 @@ def before( C_2_s1 = T.dynamic("C_2_s1", "int32") @Ts.prim_func(private=True) - def expected(A: T.Buffer((batch_size, 1, 4096), 'float16'), B: T.Buffer((28672, 4096), "float16"), C: T.Buffer((batch_size, 1, 28672), 'float16')): + def expected(A: T.Tensor((batch_size, 1, 4096), 'float16'), B: T.Tensor((28672, 4096), "float16"), C: T.Tensor((batch_size, 1, 28672), 'float16')): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): @@ -952,10 +952,10 @@ def test_matmul_metal_int4_quant(): @Ts.prim_func(private=True) def before( - B0: T.Buffer((28672, 512), "uint32"), - B1: T.Buffer((28672, 128), "float16"), - A: T.Buffer((batch_size, 1, 4096), 'float16'), - C: T.Buffer((batch_size, 1, 28672), 'float16') + B0: T.Tensor((28672, 512), "uint32"), + B1: T.Tensor((28672, 128), "float16"), + A: T.Tensor((batch_size, 1, 4096), 'float16'), + C: T.Tensor((batch_size, 1, 28672), 'float16') ): compute = Ts.sblock_alloc_buffer((28672, 4096), "float16") @@ -998,7 +998,7 @@ def before( 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"), A: T.Buffer((batch_size, 1, 4096), 'float16'), C: T.Buffer((batch_size, 1, 28672), 'float16')): + def expected(B0: T.Tensor((28672, 512), "uint32"), B1: T.Tensor((28672, 128), "float16"), A: T.Tensor((batch_size, 1, 4096), 'float16'), C: T.Tensor((batch_size, 1, 28672), 'float16')): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): diff --git a/tests/python/s_tir/dlight/test_gpu_reduction.py b/tests/python/s_tir/dlight/test_gpu_reduction.py index 2f3be305806d..41b134f66424 100644 --- a/tests/python/s_tir/dlight/test_gpu_reduction.py +++ b/tests/python/s_tir/dlight/test_gpu_reduction.py @@ -32,7 +32,7 @@ def test_decode_gemv_1(): @I.ir_module class Before: @Ts.prim_func - def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128), "float16"), V: T.Buffer((1, 1, 4096), "float16"), C: T.Buffer((1, 1, 4096), "float16")): + def func(W: T.Tensor((4096, 512), "uint32"), S: T.Tensor((4096, 128), "float16"), V: T.Tensor((1, 1, 4096), "float16"), C: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): B = Ts.sblock_alloc_buffer((4096, 4096), "float16") @@ -54,7 +54,7 @@ def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128), "float16") @I.ir_module class After: @Ts.prim_func - def func(W: T.Buffer((4096, 512), 'uint32'), S: T.Buffer((4096, 128), 'float16'), V: T.Buffer((1, 1, 4096), 'float16'), C: T.Buffer((1, 1, 4096), 'float16')): + def func(W: T.Tensor((4096, 512), 'uint32'), S: T.Tensor((4096, 128), 'float16'), V: T.Tensor((1, 1, 4096), 'float16'), C: T.Tensor((1, 1, 4096), 'float16')): T.func_attr({"global_symbol": "main", "tirx.is_scheduled": True, "tirx.noalias": True}) with Ts.sblock("root"): @@ -103,7 +103,7 @@ def test_decode_gemv_2(): @I.ir_module class Before: @Ts.prim_func - def func(W: T.Buffer((512, 4096), "uint32"), S: T.Buffer((128, 4096), "float16"), V: T.Buffer((1, 1, 4096), "float16"), C: T.Buffer((1, 1, 4096), "float16")): + def func(W: T.Tensor((512, 4096), "uint32"), S: T.Tensor((128, 4096), "float16"), V: T.Tensor((1, 1, 4096), "float16"), C: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): B = Ts.sblock_alloc_buffer((4096, 4096), "float16") @@ -125,7 +125,7 @@ def func(W: T.Buffer((512, 4096), "uint32"), S: T.Buffer((128, 4096), "float16") @I.ir_module class After: @Ts.prim_func - def func(W: T.Buffer((512, 4096), "uint32"), S: T.Buffer((128, 4096), "float16"), V: T.Buffer((1, 1, 4096), "float16"), C: T.Buffer((1, 1, 4096), "float16")): + def func(W: T.Tensor((512, 4096), "uint32"), S: T.Tensor((128, 4096), "float16"), V: T.Tensor((1, 1, 4096), "float16"), C: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"global_symbol": "main", "tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): C_rf_local = Ts.sblock_alloc_buffer((16, 1, 1, 4096), "float16", scope="local") @@ -165,7 +165,7 @@ def test_decode_gemv_3(): @I.ir_module class Before: @Ts.prim_func - def func(W: T.Buffer((512, 4096), "uint32"), S: T.Buffer((128, 4096), "float16"), V: T.Buffer((1, 1, 4096), "float16"), C: T.Buffer((1, 1, 4096), "float16")): + def func(W: T.Tensor((512, 4096), "uint32"), S: T.Tensor((128, 4096), "float16"), V: T.Tensor((1, 1, 4096), "float16"), C: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): B = Ts.sblock_alloc_buffer((4096, 4096), "float16") @@ -187,7 +187,7 @@ def func(W: T.Buffer((512, 4096), "uint32"), S: T.Buffer((128, 4096), "float16") @I.ir_module class After: @Ts.prim_func - def func(W: T.Buffer((512, 4096), 'uint32'), S: T.Buffer((128, 4096), 'float16'), V: T.Buffer((1, 1, 4096), 'float16'), C: T.Buffer((1, 1, 4096), 'float16')): + def func(W: T.Tensor((512, 4096), 'uint32'), S: T.Tensor((128, 4096), 'float16'), V: T.Tensor((1, 1, 4096), 'float16'), C: T.Tensor((1, 1, 4096), 'float16')): T.func_attr({"global_symbol": "main", "tirx.is_scheduled": True, "tirx.noalias": True}) with Ts.sblock("root"): @@ -238,7 +238,7 @@ def test_decode_gemv_4(): @I.ir_module class Before: @Ts.prim_func - def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128), "float16"), V: T.Buffer((1, 1, 4096), "float16"), C: T.Buffer((1, 1, 4096), "float16")): + def func(W: T.Tensor((4096, 512), "uint32"), S: T.Tensor((4096, 128), "float16"), V: T.Tensor((1, 1, 4096), "float16"), C: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): B = Ts.sblock_alloc_buffer((4096, 4096), "float16") @@ -260,7 +260,7 @@ def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128), "float16") @I.ir_module class After: @Ts.prim_func - def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128), "float16"), V: T.Buffer((1, 1, 4096), "float16"), C: T.Buffer((1, 1, 4096), "float16")): + def func(W: T.Tensor((4096, 512), "uint32"), S: T.Tensor((4096, 128), "float16"), V: T.Tensor((1, 1, 4096), "float16"), C: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"global_symbol": "main", "tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): C_rf_local = Ts.sblock_alloc_buffer((16, 1, 1, 4096), "float16", scope="local") @@ -302,7 +302,7 @@ def test_decode_gemv_sigmoid(): @I.ir_module class Before: @Ts.prim_func - def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128), "float16"), V: T.Buffer((1, 1, 4096), "float16"), D: T.Buffer((1, 1, 4096), "float16")): + def func(W: T.Tensor((4096, 512), "uint32"), S: T.Tensor((4096, 128), "float16"), V: T.Tensor((1, 1, 4096), "float16"), D: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): B = Ts.sblock_alloc_buffer((4096, 4096), "float16") @@ -331,7 +331,7 @@ def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128), "float16") @I.ir_module class After: @Ts.prim_func - def func(W: T.Buffer((4096, 512), 'uint32'), S: T.Buffer((4096, 128), 'float16'), V: T.Buffer((1, 1, 4096), 'float16'), D: T.Buffer((1, 1, 4096), 'float16')): + def func(W: T.Tensor((4096, 512), 'uint32'), S: T.Tensor((4096, 128), 'float16'), V: T.Tensor((1, 1, 4096), 'float16'), D: T.Tensor((1, 1, 4096), 'float16')): T.func_attr({"global_symbol": "main", "tirx.is_scheduled": True, "tirx.noalias": True}) with Ts.sblock("root"): @@ -388,7 +388,7 @@ def test_decode_gemv_1_fp32(): @I.ir_module class Before: @Ts.prim_func - def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128), "float16"), V: T.Buffer((1, 1, 4096), "float16"), C: T.Buffer((1, 1, 4096), "float16")): + def func(W: T.Tensor((4096, 512), "uint32"), S: T.Tensor((4096, 128), "float16"), V: T.Tensor((1, 1, 4096), "float16"), C: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): B = Ts.sblock_alloc_buffer((4096, 4096), "float16") @@ -417,7 +417,7 @@ def func(W: T.Buffer((4096, 512), "uint32"), S: T.Buffer((4096, 128), "float16") @I.ir_module class After: @Ts.prim_func - def func(W: T.Buffer((4096, 512), 'uint32'), S: T.Buffer((4096, 128), 'float16'), V: T.Buffer((1, 1, 4096), 'float16'), C: T.Buffer((1, 1, 4096), 'float16')): + def func(W: T.Tensor((4096, 512), 'uint32'), S: T.Tensor((4096, 128), 'float16'), V: T.Tensor((1, 1, 4096), 'float16'), C: T.Tensor((1, 1, 4096), 'float16')): T.func_attr({"global_symbol": "main", "tirx.is_scheduled": True, "tirx.noalias": True}) with Ts.sblock("root"): @@ -473,7 +473,7 @@ def test_reduction_no_spatial(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((1, 1, 4096), "float16"), B: T.Buffer((4096,), "float16"), rms_norm: T.Buffer((1, 4096), "float16")): + def main(A: T.Tensor((1, 1, 4096), "float16"), B: T.Tensor((4096,), "float16"), rms_norm: T.Tensor((1, 4096), "float16")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) Ared_temp = Ts.sblock_alloc_buffer((1, 1)) for ax0 in range(4096): @@ -490,7 +490,7 @@ def main(A: T.Buffer((1, 1, 4096), "float16"), B: T.Buffer((4096,), "float16"), @I.ir_module class After: @Ts.prim_func - def main(A: T.Buffer((1, 1, 4096), 'float16'), B: T.Buffer((4096,), 'float16'), rms_norm: T.Buffer((1, 4096), 'float16')): + def main(A: T.Tensor((1, 1, 4096), 'float16'), B: T.Tensor((4096,), 'float16'), rms_norm: T.Tensor((1, 4096), 'float16')): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) with Ts.sblock("root"): @@ -544,7 +544,7 @@ def test_spatial_inner_no_broadcasting(): @I.ir_module class Module: @Ts.prim_func - def main(lv575: T.Buffer((1376, 4096), "uint32"), lv576: T.Buffer((344, 4096), "float16"), lv574: T.Buffer((1, 1, 11008), "float16"), lv570: T.Buffer((1, 1, 4096), "float16"), p_output0_intermediate: T.Buffer((1, 1, 4096), "float16")): + def main(lv575: T.Tensor((1376, 4096), "uint32"), lv576: T.Tensor((344, 4096), "float16"), lv574: T.Tensor((1, 1, 11008), "float16"), lv570: T.Tensor((1, 1, 4096), "float16"), p_output0_intermediate: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"tirx.noalias": True}) p_output0_intermediate_1 = Ts.sblock_alloc_buffer((11008, 4096), "float16") var_matmul_intermediate = Ts.sblock_alloc_buffer((1, 1, 4096), "float16") @@ -572,7 +572,7 @@ def main(lv575: T.Buffer((1376, 4096), "uint32"), lv576: T.Buffer((344, 4096), " @I.ir_module class Expected: @Ts.prim_func - def main(lv575: T.Buffer((1376, 4096), "uint32"), lv576: T.Buffer((344, 4096), "float16"), lv574: T.Buffer((1, 1, 11008), "float16"), lv570: T.Buffer((1, 1, 4096), "float16"), p_output0_intermediate: T.Buffer((1, 1, 4096), "float16")): + def main(lv575: T.Tensor((1376, 4096), "uint32"), lv576: T.Tensor((344, 4096), "float16"), lv574: T.Tensor((1, 1, 11008), "float16"), lv570: T.Tensor((1, 1, 4096), "float16"), p_output0_intermediate: T.Tensor((1, 1, 4096), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) var_matmul_intermediate_local = Ts.sblock_alloc_buffer((1, 1, 4096), "float16", scope="local") var_matmul_intermediate_rf_local = Ts.sblock_alloc_buffer((16, 1, 1, 4096), "float16", scope="local") @@ -623,7 +623,7 @@ def test_spatial_inner_broadcasting(): @I.ir_module class Module: @Ts.prim_func - def main(A: T.Buffer((256, 256), "float32"), B: T.Buffer((256, 256), "float32")): + def main(A: T.Tensor((256, 256), "float32"), B: T.Tensor((256, 256), "float32")): T.func_attr({"tirx.noalias": True}) temp_local = Ts.sblock_alloc_buffer((256,)) for j in T.serial(256): @@ -645,7 +645,7 @@ def main(A: T.Buffer((256, 256), "float32"), B: T.Buffer((256, 256), "float32")) @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((256, 256), "float32"), B: T.Buffer((256, 256), "float32")): + def main(A: T.Tensor((256, 256), "float32"), B: T.Tensor((256, 256), "float32")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) temp_local_shared = Ts.sblock_alloc_buffer((256,), scope="shared") temp_local_rf_local = Ts.sblock_alloc_buffer((16, 256), scope="local") @@ -698,7 +698,7 @@ def test_reduction_inner_no_broadcasting(): @I.ir_module class Module: @Ts.prim_func - def main(A: T.Buffer((256, 256), "float32"), B: T.Buffer((256,), "float32")): + def main(A: T.Tensor((256, 256), "float32"), B: T.Tensor((256,), "float32")): T.func_attr({"tirx.noalias": True}) temp_local = Ts.sblock_alloc_buffer((256,)) for i in T.serial(256): @@ -720,7 +720,7 @@ def main(A: T.Buffer((256, 256), "float32"), B: T.Buffer((256,), "float32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((256, 256), "float32"), B: T.Buffer((256,), "float32")): + def main(A: T.Tensor((256, 256), "float32"), B: T.Tensor((256,), "float32")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): temp_local_local = Ts.sblock_alloc_buffer((256,), scope="local") @@ -766,7 +766,7 @@ def test_reduction_inner_no_broadcasting2(): @I.ir_module class Module: @Ts.prim_func - def main(lv9: T.Buffer((2560, 320), "uint32"), lv10: T.Buffer((2560, 80), "float16"), lv1: T.Buffer((1, 2560), "float16"), p_output0_intermediate: T.Buffer((1, 2560), "float32")): + def main(lv9: T.Tensor((2560, 320), "uint32"), lv10: T.Tensor((2560, 80), "float16"), lv1: T.Tensor((1, 2560), "float16"), p_output0_intermediate: T.Tensor((1, 2560), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): p_output0_intermediate_1 = Ts.sblock_alloc_buffer((2560, 2560), "float16") @@ -795,7 +795,7 @@ def main(lv9: T.Buffer((2560, 320), "uint32"), lv10: T.Buffer((2560, 80), "float @I.ir_module class Expected: @Ts.prim_func - def main(lv9: T.Buffer((2560, 320), "uint32"), lv10: T.Buffer((2560, 80), "float16"), lv1: T.Buffer((1, 2560), "float16"), p_output0_intermediate: T.Buffer((1, 2560), "float32")): + def main(lv9: T.Tensor((2560, 320), "uint32"), lv10: T.Tensor((2560, 80), "float16"), lv1: T.Tensor((1, 2560), "float16"), p_output0_intermediate: T.Tensor((1, 2560), "float32")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): var_matmul_intermediate_local = Ts.sblock_alloc_buffer((1, 2560), "float16", scope="local") @@ -850,7 +850,7 @@ def test_reduction_inner_spatial_choose_perfect_factor(): @I.ir_module class Module: @Ts.prim_func - def main(A: T.Buffer((T.int64(1), T.int64(32), T.int64(1), n), 'float16'), B: T.Buffer((T.int64(1), T.int64(32), n, T.int64(100)), 'float16'), matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(100)), "float16")): + def main(A: T.Tensor((T.int64(1), T.int64(32), T.int64(1), n), 'float16'), B: T.Tensor((T.int64(1), T.int64(32), n, T.int64(100)), 'float16'), matmul: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(100)), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -867,7 +867,7 @@ def main(A: T.Buffer((T.int64(1), T.int64(32), T.int64(1), n), 'float16'), B: T. @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((T.int64(1), T.int64(32), T.int64(1), n), 'float16'), B: T.Buffer((T.int64(1), T.int64(32), n, T.int64(100)), 'float16'), matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(100)), "float16")): + def main(A: T.Tensor((T.int64(1), T.int64(32), T.int64(1), n), 'float16'), B: T.Tensor((T.int64(1), T.int64(32), n, T.int64(100)), 'float16'), matmul: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(100)), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -916,9 +916,9 @@ def test_reduction_inner_spatial_non_affine_write_back_falls_back(): class Before: @Ts.prim_func def main( - A: T.Buffer((2, 4, 20), "float32"), - W: T.Buffer((3,), "float32"), - C: T.Buffer((2, 2, 20), "float32"), + A: T.Tensor((2, 4, 20), "float32"), + W: T.Tensor((3,), "float32"), + C: T.Tensor((2, 2, 20), "float32"), ): for n, y, x, k in T.grid(2, 2, 20, 3): with Ts.sblock("conv"): @@ -933,9 +933,9 @@ def main( class Expected: @Ts.prim_func def main( - A: T.Buffer((2, 4, 20), "float32"), - W: T.Buffer((3,), "float32"), - C: T.Buffer((2, 2, 20), "float32"), + A: T.Tensor((2, 4, 20), "float32"), + W: T.Tensor((3,), "float32"), + C: T.Tensor((2, 2, 20), "float32"), ): T.func_attr({"tirx.is_scheduled": True}) for ax0_ax1_ax2_fused_0 in T.thread_binding(1, thread="blockIdx.x"): @@ -988,8 +988,8 @@ def test_reduction_inner_spatial_non_affine_without_mixed_access(): class Before: @Ts.prim_func def main( - A: T.Buffer((2, 2, 3, 20), "float32"), - C: T.Buffer((2, 2, 20), "float32"), + A: T.Tensor((2, 2, 3, 20), "float32"), + C: T.Tensor((2, 2, 20), "float32"), ): for n, y, x, k in T.grid(2, 2, 20, 3): with Ts.sblock("sum"): @@ -1016,8 +1016,8 @@ def test_reduction_inner_spatial_reordered_access_declines(): class Before: @Ts.prim_func def main( - A: T.Buffer((2, 16, 3, 20), "float32"), - C: T.Buffer((2, 20, 16), "float32"), + A: T.Tensor((2, 16, 3, 20), "float32"), + C: T.Tensor((2, 20, 16), "float32"), ): for n, y, x, k in T.grid(2, 20, 16, 3): with Ts.sblock("sum"): @@ -1038,9 +1038,9 @@ def test_reduction_inner_spatial_affine_write_back_still_applies(): class Before: @Ts.prim_func def main( - A: T.Buffer((2, 4, 16), "float32"), - W: T.Buffer((3,), "float32"), - C: T.Buffer((2, 2, 16), "float32"), + A: T.Tensor((2, 4, 16), "float32"), + W: T.Tensor((3,), "float32"), + C: T.Tensor((2, 2, 16), "float32"), ): for n, y, x, k in T.grid(2, 2, 16, 3): with Ts.sblock("conv"): @@ -1062,8 +1062,8 @@ def test_reduction_inner_spatial_uses_normalized_split_extent(): class Before: @Ts.prim_func def main( - A: T.Buffer((2, 3, 3, 4), "float32"), - C: T.Buffer((2, 3, 8), "float32"), + A: T.Tensor((2, 3, 3, 4), "float32"), + C: T.Tensor((2, 3, 8), "float32"), ): for n, y, x, k in T.grid(2, 3, 8, 3): with Ts.sblock("sum"): @@ -1088,7 +1088,7 @@ def test_repeat_transpose_gemv(): @I.ir_module class Before: @Ts.prim_func(private=True) - def fused_relax_repeat_relax_permute_dims_relax_matmul1(lv716: T.Buffer((T.int64(1), kv_seq_len, T.int64(8), T.int64(128)), 'float16'), astype66: T.Buffer((T.int64(1), T.int64(32), T.int64(1), kv_seq_len), 'float16'), var_matmul_intermediate: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16")): + def fused_relax_repeat_relax_permute_dims_relax_matmul1(lv716: T.Tensor((T.int64(1), kv_seq_len, T.int64(8), T.int64(128)), 'float16'), astype66: T.Tensor((T.int64(1), T.int64(32), T.int64(1), kv_seq_len), 'float16'), var_matmul_intermediate: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1119,7 +1119,7 @@ def fused_relax_repeat_relax_permute_dims_relax_matmul1(lv716: T.Buffer((T.int64 @I.ir_module class Expected: @Ts.prim_func(private=True) - def fused_relax_repeat_relax_permute_dims_relax_matmul1(lv716: T.Buffer((T.int64(1), kv_seq_len, T.int64(8), T.int64(128)), 'float16'), astype66: T.Buffer((T.int64(1), T.int64(32), T.int64(1), kv_seq_len), 'float16'), var_matmul_intermediate: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16")): + def fused_relax_repeat_relax_permute_dims_relax_matmul1(lv716: T.Tensor((T.int64(1), kv_seq_len, T.int64(8), T.int64(128)), 'float16'), astype66: T.Tensor((T.int64(1), T.int64(32), T.int64(1), kv_seq_len), 'float16'), var_matmul_intermediate: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -1171,9 +1171,9 @@ def test_gemv_dyn_shape_epilogue(): class Module: @Ts.prim_func(private=True) def main( - A: T.Buffer((T.int64(4096), vocab_size), "float16"), - B: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), - C: T.Buffer((T.int64(1), T.int64(1), vocab_size)), + A: T.Tensor((T.int64(4096), vocab_size), "float16"), + B: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), + C: T.Tensor((T.int64(1), T.int64(1), vocab_size)), ): T.func_attr({"tirx.noalias": True}) @@ -1201,7 +1201,7 @@ def main( @I.ir_module class Expected: @Ts.prim_func(private=True) - def main(A: T.Buffer((T.int64(4096), vocab_size), 'float16'), B: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float16"), C: T.Buffer((T.int64(1), T.int64(1), vocab_size))): + def main(A: T.Tensor((T.int64(4096), vocab_size), 'float16'), B: T.Tensor((T.int64(1), T.int64(1), T.int64(4096)), "float16"), C: T.Tensor((T.int64(1), T.int64(1), vocab_size))): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -1252,7 +1252,7 @@ def test_gemv_output_one_element(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((T.int64(1), T.int64(2048)), "float16"), weight: T.Buffer((T.int64(1), T.int64(2048)), "float16"), out: T.Buffer((T.int64(1), T.int64(1)), "float16")): + def main(A: T.Tensor((T.int64(1), T.int64(2048)), "float16"), weight: T.Tensor((T.int64(1), T.int64(2048)), "float16"), out: T.Tensor((T.int64(1), T.int64(1)), "float16")): T.func_attr({"tirx.noalias": True}) NT_matmul_intermediate = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1)), "float16") for i0, i1, k in T.grid(T.int64(1), T.int64(1), T.int64(2048)): @@ -1269,7 +1269,7 @@ def main(A: T.Buffer((T.int64(1), T.int64(2048)), "float16"), weight: T.Buffer(( @I.ir_module class Expected: @Ts.prim_func(private=True) - def main(A: T.Buffer((T.int64(1), T.int64(2048)), "float16"), weight: T.Buffer((T.int64(1), T.int64(2048)), "float16"), out: T.Buffer((T.int64(1), T.int64(1)), "float16")): + def main(A: T.Tensor((T.int64(1), T.int64(2048)), "float16"), weight: T.Tensor((T.int64(1), T.int64(2048)), "float16"), out: T.Tensor((T.int64(1), T.int64(1)), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) NT_matmul_intermediate_shared = Ts.sblock_alloc_buffer((T.int64(1), T.int64(1)), "float16", scope="shared") NT_matmul_intermediate_rf_local = Ts.sblock_alloc_buffer((T.int64(1024), T.int64(1), T.int64(1)), "float16", scope="local") @@ -1313,7 +1313,7 @@ def test_no_reduction_loop_check(): @I.ir_module class Before: @Ts.prim_func(private=True) - def matmul(lv43: T.Buffer((T.int64(1), T.int64(32), T.int64(1)), "float16"), lv44: T.Buffer((T.int64(1), T.int64(1), T.int64(1)), "float16"), matmul: T.Buffer((T.int64(1), T.int64(32), T.int64(1)), "float16")): + def matmul(lv43: T.Tensor((T.int64(1), T.int64(32), T.int64(1)), "float16"), lv44: T.Tensor((T.int64(1), T.int64(1), T.int64(1)), "float16"), matmul: T.Tensor((T.int64(1), T.int64(32), T.int64(1)), "float16")): T.func_attr({"op_pattern": 4, "tirx.noalias": True}) # with Ts.sblock("root"): for i0, i1, i2, k in T.grid(T.int64(1), T.int64(32), T.int64(1), T.int64(1)): diff --git a/tests/python/s_tir/dlight/test_gpu_rmsnorm.py b/tests/python/s_tir/dlight/test_gpu_rmsnorm.py index 1f37069004ba..9f94fc100b91 100644 --- a/tests/python/s_tir/dlight/test_gpu_rmsnorm.py +++ b/tests/python/s_tir/dlight/test_gpu_rmsnorm.py @@ -41,7 +41,7 @@ def test_rms_norm_with_casting(): @I.ir_module class Before: @Ts.prim_func - def main(data: T.Buffer((1, n, 4096), 'float16'), weight: T.Buffer((4096,), "float16"), T_cast: T.Buffer((1, n, 4096), 'float16')): + def main(data: T.Tensor((1, n, 4096), 'float16'), weight: T.Tensor((4096,), "float16"), T_cast: T.Tensor((1, n, 4096), 'float16')): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -101,7 +101,7 @@ def main(data: T.Buffer((1, n, 4096), 'float16'), weight: T.Buffer((4096,), "flo @I.ir_module class After: @Ts.prim_func - def main(data: T.Buffer((1, n, 4096), 'float16'), weight: T.Buffer((4096,), "float16"), T_cast: T.Buffer((1, n, 4096), 'float16')): + def main(data: T.Tensor((1, n, 4096), 'float16'), weight: T.Tensor((4096,), "float16"), T_cast: T.Tensor((1, n, 4096), 'float16')): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -173,7 +173,7 @@ def test_rms_norm_without_casting(): @I.ir_module class Before: @Ts.prim_func - def main(data: T.Buffer((1, n, 4096)), weight: T.Buffer((4096,), "float32"), T_cast: T.Buffer((1, n, 4096))): + def main(data: T.Tensor((1, n, 4096)), weight: T.Tensor((4096,), "float32"), T_cast: T.Tensor((1, n, 4096))): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -219,7 +219,7 @@ def main(data: T.Buffer((1, n, 4096)), weight: T.Buffer((4096,), "float32"), T_c @I.ir_module class After: @Ts.prim_func - def main(data: T.Buffer((1, n, 4096)), weight: T.Buffer((4096,), "float32"), T_cast: T.Buffer((1, n, 4096))): + def main(data: T.Tensor((1, n, 4096)), weight: T.Tensor((4096,), "float32"), T_cast: T.Tensor((1, n, 4096))): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): diff --git a/tests/python/s_tir/dlight/test_gpu_transpose.py b/tests/python/s_tir/dlight/test_gpu_transpose.py index 36b47bfdcee4..7ce9a60f8064 100644 --- a/tests/python/s_tir/dlight/test_gpu_transpose.py +++ b/tests/python/s_tir/dlight/test_gpu_transpose.py @@ -39,7 +39,7 @@ def test_transpose(): @I.ir_module class Before: @Ts.prim_func - def main(rxplaceholder: T.Buffer((T.int64(512), T.int64(4096)), "float32"), T_transpose: T.Buffer((T.int64(4096), T.int64(512)), "float32")): + def main(rxplaceholder: T.Tensor((T.int64(512), T.int64(4096)), "float32"), T_transpose: T.Tensor((T.int64(4096), T.int64(512)), "float32")): T.func_attr({"tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(4096), T.int64(512)): with Ts.sblock("T_transpose"): @@ -49,7 +49,7 @@ def main(rxplaceholder: T.Buffer((T.int64(512), T.int64(4096)), "float32"), T_tr @I.ir_module class After: @Ts.prim_func - def main(rxplaceholder: T.Buffer((T.int64(512), T.int64(4096)), "float32"), T_transpose: T.Buffer((T.int64(4096), T.int64(512)), "float32")): + def main(rxplaceholder: T.Tensor((T.int64(512), T.int64(4096)), "float32"), T_transpose: T.Tensor((T.int64(4096), T.int64(512)), "float32")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): rxplaceholder_shared = Ts.sblock_alloc_buffer((T.int64(512), T.int64(4096)), scope="shared") @@ -85,7 +85,7 @@ def test_decode_transpose(): @I.ir_module class Before: @Ts.prim_func - def main(rxplaceholder: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), rxplaceholder_1: T.Buffer((T.int64(128), T.int64(4096)), "uint32"), T_transpose: T.Buffer((T.int64(4096), T.int64(4096)), "float32")): + def main(rxplaceholder: T.Tensor((T.int64(512), T.int64(4096)), "uint32"), rxplaceholder_1: T.Tensor((T.int64(128), T.int64(4096)), "uint32"), T_transpose: T.Tensor((T.int64(4096), T.int64(4096)), "float32")): T.func_attr({"tirx.noalias": True}) decode = Ts.sblock_alloc_buffer((T.int64(4096), T.int64(4096))) for i, j in T.grid(T.int64(4096), T.int64(4096)): @@ -104,7 +104,7 @@ def main(rxplaceholder: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), rxpla @I.ir_module class After: @Ts.prim_func - def main(rxplaceholder: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), rxplaceholder_1: T.Buffer((T.int64(128), T.int64(4096)), "uint32"), T_transpose: T.Buffer((T.int64(4096), T.int64(4096)), "float32")): + def main(rxplaceholder: T.Tensor((T.int64(512), T.int64(4096)), "uint32"), rxplaceholder_1: T.Tensor((T.int64(128), T.int64(4096)), "uint32"), T_transpose: T.Tensor((T.int64(4096), T.int64(4096)), "float32")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) decode_shared = Ts.sblock_alloc_buffer((T.int64(4096), T.int64(4096)), scope="shared") for ax0_0_0 in T.thread_binding(T.int64(64), thread="blockIdx.y", annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}): @@ -139,7 +139,7 @@ def test_decode_int3_transpose(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((T.int64(412), T.int64(4096)), "uint32"), B: T.Buffer((T.int64(103), T.int64(4096)), "float16"), T_transpose: T.Buffer((T.int64(4096), T.int64(4096)), "float16")): + def main(A: T.Tensor((T.int64(412), T.int64(4096)), "uint32"), B: T.Tensor((T.int64(103), T.int64(4096)), "float16"), T_transpose: T.Tensor((T.int64(4096), T.int64(4096)), "float16")): T.func_attr({"tirx.noalias": True}) decode_1 = Ts.sblock_alloc_buffer((T.int64(4096), T.int64(4096)), "float16") for i, j in T.grid(T.int64(4096), T.int64(4096)): @@ -158,7 +158,7 @@ def main(A: T.Buffer((T.int64(412), T.int64(4096)), "uint32"), B: T.Buffer((T.in @I.ir_module class After: @Ts.prim_func - def main(A: T.Buffer((T.int64(412), T.int64(4096)), "uint32"), B: T.Buffer((T.int64(103), T.int64(4096)), "float16"), T_transpose: T.Buffer((T.int64(4096), T.int64(4096)), "float16")): + def main(A: T.Tensor((T.int64(412), T.int64(4096)), "uint32"), B: T.Tensor((T.int64(103), T.int64(4096)), "float16"), T_transpose: T.Tensor((T.int64(4096), T.int64(4096)), "float16")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): decode_1_shared = Ts.sblock_alloc_buffer((T.int64(4096), T.int64(4096)), "float16", scope="shared") diff --git a/tests/python/s_tir/dlight/test_primitives.py b/tests/python/s_tir/dlight/test_primitives.py index f7f4f8f62e9e..2dcdaa1c18b8 100644 --- a/tests/python/s_tir/dlight/test_primitives.py +++ b/tests/python/s_tir/dlight/test_primitives.py @@ -27,7 +27,7 @@ @Ts.prim_func -def main(p0: T.Buffer((), "int32"), T_stack: T.Buffer((T.int64(3),), "int32")): +def main(p0: T.Tensor((), "int32"), T_stack: T.Tensor((T.int64(3),), "int32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): compile_engine_const = Ts.sblock_alloc_buffer((), "int32") diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_arg_info.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_arg_info.py index b1d12cf96b24..f5d2889fd021 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_arg_info.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_arg_info.py @@ -25,9 +25,9 @@ @Ts.prim_func def Matmul( - A: T.Buffer((128, 256), "float32"), - B: T.Buffer((256, 512), "float32"), - C: T.Buffer((128, 512), "float32"), + A: T.Tensor((128, 256), "float32"), + B: T.Tensor((256, 512), "float32"), + C: T.Tensor((128, 512), "float32"), ) -> None: T.func_attr({"global_symbol": "main"}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_builder.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_builder.py index c299fbe0d38b..dde07b86dc94 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_builder.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_builder.py @@ -44,9 +44,9 @@ class MatmulModule: @Ts.prim_func def matmul( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "matmul", "tirx.noalias": True}) @@ -62,9 +62,9 @@ def matmul( class MatmulReluModule: @Ts.prim_func def matmul_relu( # pylint: disable=no-self-argument - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - D: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + D: T.Tensor((1024, 1024), "float32"), ) -> None: T.func_attr({"global_symbol": "matmul_relu", "tirx.noalias": True}) @@ -85,7 +85,7 @@ def matmul_relu( # pylint: disable=no-self-argument class BatchMatmulModule: @Ts.prim_func def batch_matmul( # pylint: disable=no-self-argument - A: T.Buffer([16, 128, 128]), B: T.Buffer([16, 128, 128]), C: T.Buffer([16, 128, 128]) + A: T.Tensor([16, 128, 128]), B: T.Tensor([16, 128, 128]), C: T.Tensor([16, 128, 128]) ) -> None: T.func_attr({"global_symbol": "batch_matmul", "tirx.noalias": True}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_cost_model.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_cost_model.py index 3a6c14013c0e..636e28c4f7d7 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_cost_model.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_cost_model.py @@ -49,9 +49,9 @@ class Matmul: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -66,7 +66,7 @@ def main( @tvm.script.ir_module class FullModule: @Ts.prim_func - def main(T_full: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def main(T_full: T.Tensor((T.int64(2), T.int64(3)), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for ax0, ax1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_database.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_database.py index 375c0b7d1e9d..e485b42416c7 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_database.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_database.py @@ -45,9 +45,9 @@ class Matmul: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: T.func_attr({"global_symbol": "main"}) @@ -62,9 +62,9 @@ def main( class MatmulRelu: @Ts.prim_func def main( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_feature_extractor_per_store_feature.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_feature_extractor_per_store_feature.py index 919e11e1e7c6..4348f8ba3da7 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_feature_extractor_per_store_feature.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_feature_extractor_per_store_feature.py @@ -34,9 +34,9 @@ @Ts.prim_func def matmul( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - C: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + C: T.Tensor((512, 512), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -59,7 +59,7 @@ def matmul( @tvm.script.ir_module class LayoutTransform: @Ts.prim_func - def main(placeholder: T.Buffer((1, 16, 7, 7, 32), "float32"), placeholder_1: T.Buffer((25088,), "float32"), T_layout_trans: T.Buffer((1, 1, 7, 7, 512), "float32")) -> None: + def main(placeholder: T.Tensor((1, 16, 7, 7, 32), "float32"), placeholder_1: T.Tensor((25088,), "float32"), T_layout_trans: T.Tensor((1, 1, 7, 7, 512), "float32")) -> None: # function attr dict T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) # body @@ -419,9 +419,9 @@ def test_cpu_fusion(): # pylint: disable=all @Ts.prim_func def func( - A: T.Buffer([64, 32], dtype="float32"), - B: T.Buffer([64, 32], dtype="float32"), - C: T.Buffer([64, 32], dtype="float32"), + A: T.Tensor([64, 32], dtype="float32"), + B: T.Tensor([64, 32], dtype="float32"), + C: T.Tensor([64, 32], dtype="float32"), ) -> None: for i, j in T.grid(64, 32): # type: ignore with Ts.sblock(): @@ -716,7 +716,7 @@ def _create_schedule(): def test_empty_feature(): @Ts.prim_func - def full(T_full: T.Buffer((T.int64(2), T.int64(3)), "float32")): + def full(T_full: T.Tensor((T.int64(2), T.int64(3)), "float32")): for ax0, ax1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): v_ax0, v_ax1 = Ts.axis.remap("SS", [ax0, ax1]) @@ -1627,7 +1627,7 @@ def test_cpu_layout_transform(): @Ts.prim_func -def negative_extent(A: T.Buffer((1,), "float32")): +def negative_extent(A: T.Tensor((1,), "float32")): for j in range(0, -1): A[j] = A[j] + 1.0 diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_measure_callback.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_measure_callback.py index 830d6c1cadff..7126d3e76ece 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_measure_callback.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_measure_callback.py @@ -33,9 +33,9 @@ class Matmul: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: T.func_attr({"global_symbol": "main"}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_mma_tensorize.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_mma_tensorize.py index ce6165737311..5ed900b0be4f 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_mma_tensorize.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_mma_tensorize.py @@ -38,9 +38,9 @@ class Gemm_F16F16F16: # fmt: off @Ts.prim_func def main( - A: T.Buffer((M, K), "float16"), # type: ignore - B: T.Buffer((K, N), "float16"), # type: ignore - C: T.Buffer((M, N), "float16"), # type: ignore + A: T.Tensor((M, K), "float16"), # type: ignore + B: T.Tensor((K, N), "float16"), # type: ignore + C: T.Tensor((M, N), "float16"), # type: ignore ): for i, j, k in T.grid(M, N, K): with Ts.sblock("C"): @@ -55,9 +55,9 @@ class Gemm_F16F16F32: # fmt: off @Ts.prim_func def main( - A: T.Buffer((M, K), "float16"), # type: ignore - B: T.Buffer((K, N), "float16"), # type: ignore - C: T.Buffer((M, N), "float32"), # type: ignore + A: T.Tensor((M, K), "float16"), # type: ignore + B: T.Tensor((K, N), "float16"), # type: ignore + C: T.Tensor((M, N), "float32"), # type: ignore ): for i, j, k in T.grid(M, N, K): with Ts.sblock("C"): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_compute_location.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_compute_location.py index dfa9c9c94a8e..89c5d6713c29 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_compute_location.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_compute_location.py @@ -26,8 +26,8 @@ @Ts.prim_func def add( - A: T.Buffer([2048, 2048, 2048], dtype="float32"), - B: T.Buffer([2048, 2048, 2048], dtype="float32"), + A: T.Tensor([2048, 2048, 2048], dtype="float32"), + B: T.Tensor([2048, 2048, 2048], dtype="float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main"}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_parallel.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_parallel.py index e72613f16bf2..6fa084b96420 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_parallel.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_parallel.py @@ -26,7 +26,7 @@ @Ts.prim_func -def matmul(A: T.Buffer([512, 512]), B: T.Buffer([512, 512]), C: T.Buffer([512, 512])) -> None: +def matmul(A: T.Tensor([512, 512]), B: T.Tensor([512, 512]), C: T.Tensor([512, 512])) -> None: for i, j, k in T.grid(512, 512, 512): # type: ignore with Ts.sblock("C"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) # type: ignore diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_thread_binding.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_thread_binding.py index 6c92a0c7a6bb..4ba6dd714ed7 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_thread_binding.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_thread_binding.py @@ -26,7 +26,7 @@ @Ts.prim_func def element_wise( - A: T.Buffer([512, 512], dtype="float32"), B: T.Buffer([512, 512], dtype="float32") + A: T.Tensor([512, 512], dtype="float32"), B: T.Tensor([512, 512], dtype="float32") ) -> None: for i, j in T.grid(512, 512): with Ts.sblock("C"): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_tile_size.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_tile_size.py index 14321d4eb81b..74da2bbb05f5 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_tile_size.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_tile_size.py @@ -28,7 +28,7 @@ @Ts.prim_func -def matmul(A: T.Buffer([512, 512]), B: T.Buffer([512, 512]), C: T.Buffer([512, 512])) -> None: +def matmul(A: T.Tensor([512, 512]), B: T.Tensor([512, 512]), C: T.Tensor([512, 512])) -> None: for i, j, k in T.grid(512, 512, 512): # type: ignore with Ts.sblock("C"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) # type: ignore diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_unroll.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_unroll.py index 92d162305964..3642ada63033 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_unroll.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_mutator_mutate_unroll.py @@ -26,7 +26,7 @@ @Ts.prim_func -def matmul(A: T.Buffer([512, 512]), B: T.Buffer([512, 512]), C: T.Buffer([512, 512])) -> None: +def matmul(A: T.Tensor([512, 512]), B: T.Tensor([512, 512]), C: T.Tensor([512, 512])) -> None: for i, j, k in T.grid(512, 512, 512): # type: ignore with Ts.sblock("C"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) # type: ignore diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_post_order_apply.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_post_order_apply.py index 1092722900f9..718f7f1dca5b 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_post_order_apply.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_post_order_apply.py @@ -59,9 +59,9 @@ def get_matmul_packed(m, n, k, lhs_type="int8", rhs_dtype="int8", acc_dtype="int class Matmul: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: T.func_attr({"global_symbol": "main"}) @@ -76,9 +76,9 @@ def main( class DuplicateMatmul: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: T.func_attr({"global_symbol": "main"}) @@ -96,7 +96,7 @@ def main( @tvm.script.ir_module class TrinityMatmul: @Ts.prim_func - def main(A: T.Buffer((1024, 1024), 'float32'), D: T.Buffer((1024, 1024), 'float32')) -> None: + def main(A: T.Tensor((1024, 1024), 'float32'), D: T.Tensor((1024, 1024), 'float32')) -> None: T.func_attr({"global_symbol": "main"}) B = Ts.sblock_alloc_buffer((1024, 1024), "float32") @@ -119,7 +119,7 @@ def main(A: T.Buffer((1024, 1024), 'float32'), D: T.Buffer((1024, 1024), 'float3 class TrinityMatmulProcessedForReference: @Ts.prim_func def main( - A: T.Buffer([1024, 1024], dtype="float32"), D: T.Buffer([1024, 1024], dtype="float32") + A: T.Tensor([1024, 1024], dtype="float32"), D: T.Tensor([1024, 1024], dtype="float32") ) -> None: # function attr dict T.func_attr({"global_symbol": "main"}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_disallow_async_strided_mem_copy.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_disallow_async_strided_mem_copy.py index 3b454ec02651..82b132ab3f74 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_disallow_async_strided_mem_copy.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_disallow_async_strided_mem_copy.py @@ -52,9 +52,9 @@ def _create_context(mod, target) -> ms.TuneContext: class Matmul: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: T.func_attr({"global_symbol": "main"}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_disallow_dynamic_loop.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_disallow_dynamic_loop.py index 4762a9ee4479..b0fd90069b6f 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_disallow_dynamic_loop.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_disallow_dynamic_loop.py @@ -52,9 +52,9 @@ def _create_context(mod, target) -> ms.TuneContext: class Matmul: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: T.func_attr({"global_symbol": "main"}) @@ -69,9 +69,9 @@ def main( class DynamicLoop: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: T.func_attr({"global_symbol": "main"}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_cooperative_fetch.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_cooperative_fetch.py index f0ffedbd9267..a819e056c628 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_cooperative_fetch.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_cooperative_fetch.py @@ -54,7 +54,7 @@ def _create_context(mod, target) -> ms.TuneContext: @tvm.script.ir_module class AfterRewrite0: @Ts.prim_func - def main(A: T.Buffer([512, 512], dtype='float32'), B: T.Buffer([512, 512], dtype='float32'), C: T.Buffer([512, 512], dtype='float32')) -> None: + def main(A: T.Tensor([512, 512], dtype='float32'), B: T.Tensor([512, 512], dtype='float32'), C: T.Tensor([512, 512], dtype='float32')) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -106,9 +106,9 @@ def main(A: T.Buffer([512, 512], dtype='float32'), B: T.Buffer([512, 512], dtype class WarpExecutionAfterRewrite: @Ts.prim_func def main( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - C: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + C: T.Tensor((512, 512), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_layout.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_layout.py index 9d03d103e92a..63ca669dbf39 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_layout.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_layout.py @@ -78,9 +78,9 @@ def test_tir_matmul(): @Ts.prim_func(private=True) def before( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: T.func_attr({"layout_free_buffers": [1]}) for i0, j, k0, i1, k1 in T.grid(4, 16, 4, 4, 4): @@ -94,9 +94,9 @@ def before( @Ts.prim_func(private=True) def expected( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: T.func_attr({"layout_free_buffers": [1]}) B_reindex = Ts.sblock_alloc_buffer([16, 4, 4], dtype="float32") @@ -124,7 +124,7 @@ def test_rewritten_buffers_must_occur_within_block(): @Ts.prim_func(private=True) def before( - A: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), ) -> None: T.func_attr({"layout_free_buffers": [0]}) for i, j in T.grid(16, 16): @@ -144,7 +144,7 @@ def test_extent_one(): @Ts.prim_func(private=True) def before( - A: T.Buffer((16, 1), "float32"), + A: T.Tensor((16, 1), "float32"), ) -> None: T.func_attr({"layout_free_buffers": [0]}) for i, j in T.grid(16, 1): @@ -153,7 +153,7 @@ def before( T.evaluate(A[vi, vj]) @Ts.prim_func(private=True) - def expected(A: T.Buffer((16, 1), "float32")): + def expected(A: T.Tensor((16, 1), "float32")): T.func_attr({"layout_free_buffers": [0]}) A_global = Ts.sblock_alloc_buffer([16], dtype="float32") @@ -175,9 +175,9 @@ def expected(A: T.Buffer((16, 1), "float32")): @Ts.prim_func def tir_matmul( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: T.func_attr({"layout_free_buffers": [1]}) for i0, j, k0, i1, k1 in T.grid(4, 16, 4, 4, 4): @@ -192,9 +192,9 @@ def tir_matmul( @Ts.prim_func def rewritten_tir_matmul( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: T.func_attr({"layout_free_buffers": [1]}) B_reindex = Ts.sblock_alloc_buffer([16, 4, 4], dtype="float32") @@ -226,7 +226,7 @@ def test_layout_rewrite(): @tvm.script.ir_module class Conv2dCacheRead: @Ts.prim_func - def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((3, 3, 64, 64), "float32"), conv2d_nhwc: T.Buffer((1, 56, 56, 64), "float32")): + def main(p0: T.Tensor((1, 56, 56, 64), "float32"), p1: T.Tensor((3, 3, 64, 64), "float32"), conv2d_nhwc: T.Tensor((1, 56, 56, 64), "float32")): T.func_attr({"layout_free_buffers": [1], "tirx.noalias": True, "global_symbol": "main"}) pad_temp = Ts.sblock_alloc_buffer([1, 58, 58, 64], dtype="float32") conv2d_nhwc_global = Ts.sblock_alloc_buffer([1, 56, 56, 64], dtype="float32") @@ -303,7 +303,7 @@ def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((3, 3, 64, 64), @tvm.script.ir_module class Conv2dCacheReadRewritten: @Ts.prim_func - def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((3, 3, 64, 64), "float32"), conv2d_nhwc: T.Buffer((1, 56, 56, 64), "float32")): + def main(p0: T.Tensor((1, 56, 56, 64), "float32"), p1: T.Tensor((3, 3, 64, 64), "float32"), conv2d_nhwc: T.Tensor((1, 56, 56, 64), "float32")): T.func_attr({"layout_free_buffers": [1], "tirx.noalias": True, "global_symbol": "main"}) pad_temp = Ts.sblock_alloc_buffer([1, 58, 58, 64], dtype="float32") conv2d_nhwc_global = Ts.sblock_alloc_buffer([1, 56, 56, 64], dtype="float32") @@ -388,7 +388,7 @@ def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((3, 3, 64, 64), @tvm.script.ir_module class Conv2dCacheReadMultipleRewritten: @Ts.prim_func - def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((3, 3, 64, 64), "float32"), conv2d_nhwc: T.Buffer((1, 56, 56, 64), "float32")): + def main(p0: T.Tensor((1, 56, 56, 64), "float32"), p1: T.Tensor((3, 3, 64, 64), "float32"), conv2d_nhwc: T.Tensor((1, 56, 56, 64), "float32")): T.func_attr({"layout_free_buffers": [1], "tirx.noalias": True, "global_symbol": "main"}) pad_temp = Ts.sblock_alloc_buffer([1, 58, 58, 64], dtype="float32") conv2d_nhwc_global = Ts.sblock_alloc_buffer([1, 56, 56, 64], dtype="float32") @@ -501,9 +501,9 @@ def test_layout_rewrite_cache_read_multiple(): def test_layout_rewrite_int64_index(): @Ts.prim_func(private=True) def before( - p0: T.Buffer((T.int64(12), T.int64(197), T.int64(64)), "int8"), - p1: T.Buffer((T.int64(12), T.int64(197), T.int64(64)), "int8"), - T_batch_matmul_NT: T.Buffer((T.int64(12), T.int64(197), T.int64(197)), "int32"), + p0: T.Tensor((T.int64(12), T.int64(197), T.int64(64)), "int8"), + p1: T.Tensor((T.int64(12), T.int64(197), T.int64(64)), "int8"), + T_batch_matmul_NT: T.Tensor((T.int64(12), T.int64(197), T.int64(197)), "int32"), ): T.func_attr({"layout_free_buffers": [1], "tirx.noalias": True}) for b_0_i_0_fused in T.parallel(T.int64(394)): @@ -562,9 +562,9 @@ def before( @Ts.prim_func(private=True) def expected( - p0: T.Buffer((T.int64(12), T.int64(197), T.int64(64)), "int8"), - p1: T.Buffer((T.int64(12), T.int64(197), T.int64(64)), "int8"), - T_batch_matmul_NT: T.Buffer((T.int64(12), T.int64(197), T.int64(197)), "int32"), + p0: T.Tensor((T.int64(12), T.int64(197), T.int64(64)), "int8"), + p1: T.Tensor((T.int64(12), T.int64(197), T.int64(64)), "int8"), + T_batch_matmul_NT: T.Tensor((T.int64(12), T.int64(197), T.int64(197)), "int32"), ): T.func_attr({"tirx.noalias": True, "layout_free_buffers": [1]}) p1_global = Ts.sblock_alloc_buffer( 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 9f984a945c5c..d5b2f0a0b84a 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 @@ -30,7 +30,7 @@ @tvm.script.ir_module class Move_PUV: @Ts.prim_func - def main(A: T.Buffer([1024, 1024, 1024], dtype='float32'), B: T.Buffer([1024, 1024, 1024], dtype='float32')) -> None: + def main(A: T.Tensor([1024, 1024, 1024], dtype='float32'), B: T.Tensor([1024, 1024, 1024], dtype='float32')) -> None: # function attr dict T.func_attr({"global_symbol": "main"}) @@ -48,7 +48,7 @@ def main(A: T.Buffer([1024, 1024, 1024], dtype='float32'), B: T.Buffer([1024, 10 B[vi, vj, vk] = A[vi, vj, vk] @Ts.prim_func -def Move_PUV0(A: T.Buffer([1024, 1024, 1024], dtype='float32'), B: T.Buffer([1024, 1024, 1024], dtype='float32')) -> None: +def Move_PUV0(A: T.Tensor([1024, 1024, 1024], dtype='float32'), B: T.Tensor([1024, 1024, 1024], dtype='float32')) -> None: # function attr dict T.func_attr({"global_symbol": "main"}) @@ -73,7 +73,7 @@ def Move_PUV0(A: T.Buffer([1024, 1024, 1024], dtype='float32'), B: T.Buffer([102 @tvm.script.ir_module class Fused_NN_Dense: @Ts.prim_func - def main(placeholder: T.Buffer((64, 768), "float32"), placeholder_1: T.Buffer((768, 768), "float32"), T_matmul_NT: T.Buffer((64, 768), "float32")) -> None: + def main(placeholder: T.Tensor((64, 768), "float32"), placeholder_1: T.Tensor((768, 768), "float32"), T_matmul_NT: T.Tensor((64, 768), "float32")) -> None: for i0, i1, i2 in T.grid(64, 768, 768): with Ts.sblock("T_matmul_NT"): i, j, k = Ts.axis.remap("SSR", [i0, i1, i2]) @@ -85,9 +85,9 @@ def main(placeholder: T.Buffer((64, 768), "float32"), placeholder_1: T.Buffer((7 @Ts.prim_func def before_matmul_vectorize( - placeholder: T.Buffer((64, 768), "float32"), - placeholder_1: T.Buffer((768, 768), "float32"), - T_matmul_NT: T.Buffer((64, 768), "float32"), + placeholder: T.Tensor((64, 768), "float32"), + placeholder_1: T.Tensor((768, 768), "float32"), + T_matmul_NT: T.Tensor((64, 768), "float32"), ) -> None: with Ts.sblock("root"): Ts.reads() @@ -115,9 +115,9 @@ def before_matmul_vectorize( @Ts.prim_func def after_matmul_vectorize( - placeholder: T.Buffer((64, 768), "float32"), - placeholder_1: T.Buffer((768, 768), "float32"), - T_matmul_NT: T.Buffer((64, 768), "float32"), + placeholder: T.Tensor((64, 768), "float32"), + placeholder_1: T.Tensor((768, 768), "float32"), + T_matmul_NT: T.Tensor((64, 768), "float32"), ) -> None: T_matmul_NT_global = Ts.sblock_alloc_buffer([64, 768], dtype="float32") for i0_0, i1_0, i0_1, i1_1 in T.grid(1, 16, 1, 3): @@ -143,9 +143,9 @@ def after_matmul_vectorize( @Ts.prim_func def before_postproc_add( - lhs: T.Buffer((1, 8, 56, 56, 32), "uint8"), - rhs: T.Buffer((1, 8, 56, 56, 32), "uint8"), - add_compute: T.Buffer((1, 8, 56, 56, 32), "uint8"), + lhs: T.Tensor((1, 8, 56, 56, 32), "uint8"), + rhs: T.Tensor((1, 8, 56, 56, 32), "uint8"), + add_compute: T.Tensor((1, 8, 56, 56, 32), "uint8"), ) -> None: with Ts.sblock("root"): Ts.sblock_attr({"meta_schedule.parallel":64, "meta_schedule.vectorize":128}) @@ -158,9 +158,9 @@ def before_postproc_add( @Ts.prim_func def after_postproc_add( - lhs: T.Buffer((1, 8, 56, 56, 32), "uint8"), - rhs: T.Buffer((1, 8, 56, 56, 32), "uint8"), - add_compute: T.Buffer((1, 8, 56, 56, 32), "uint8"), + lhs: T.Tensor((1, 8, 56, 56, 32), "uint8"), + rhs: T.Tensor((1, 8, 56, 56, 32), "uint8"), + add_compute: T.Tensor((1, 8, 56, 56, 32), "uint8"), ) -> None: with Ts.sblock("root"): for n_c0_h_w_c1_fused_0 in T.parallel(0, 6272): @@ -179,8 +179,8 @@ def after_postproc_add( @Ts.prim_func def before_postproc_dynamic_shape_vectorize( - A: T.Buffer((n,), dtype='float32'), - B: T.Buffer((n,), dtype='float32'), + A: T.Tensor((n,), dtype='float32'), + B: T.Tensor((n,), dtype='float32'), ) -> None: with Ts.sblock("root"): @@ -221,7 +221,7 @@ def test_parallel_vectorize_add(): def test_no_unroll_for_spatial_block(): # fmt: off @Ts.prim_func - def layer_norm(A: T.Buffer((1, 4, 4, 32), "float32"), B: T.Buffer((4, 4, 32), "float32"), C: T.Buffer((4, 4, 32), "float32"), T_layer_norm: T.Buffer((1, 4, 4, 32), "float32")): + def layer_norm(A: T.Tensor((1, 4, 4, 32), "float32"), B: T.Tensor((4, 4, 32), "float32"), C: T.Tensor((4, 4, 32), "float32"), T_layer_norm: T.Tensor((1, 4, 4, 32), "float32")): with Ts.sblock("root"): Ts.sblock_attr({"meta_schedule.unroll_explicit": 512}) A_red_temp_v0 = Ts.sblock_alloc_buffer((1,)) @@ -246,7 +246,7 @@ def layer_norm(A: T.Buffer((1, 4, 4, 32), "float32"), B: T.Buffer((4, 4, 32), "f T_layer_norm[v_ax0, v_ax1, v_ax2, v_ax3] = (A[v_ax0, v_ax1, v_ax2, v_ax3] - A_red_temp_v0[v_ax0] * T.float32(0.001953125)) * T.rsqrt(A_red_temp_v1[v_ax0] * T.float32(0.001953125) - A_red_temp_v0[v_ax0] * T.float32(0.001953125) * (A_red_temp_v0[v_ax0] * T.float32(0.001953125)) + T.float32(1.0000000000000001e-05)) * B[v_ax1, v_ax2, v_ax3] + C[v_ax1, v_ax2, v_ax3] @Ts.prim_func - def expected(A: T.Buffer((1, 4, 4, 32), "float32"), B: T.Buffer((4, 4, 32), "float32"), C: T.Buffer((4, 4, 32), "float32"), T_layer_norm: T.Buffer((1, 4, 4, 32), "float32")): + def expected(A: T.Tensor((1, 4, 4, 32), "float32"), B: T.Tensor((4, 4, 32), "float32"), C: T.Tensor((4, 4, 32), "float32"), T_layer_norm: T.Tensor((1, 4, 4, 32), "float32")): with Ts.sblock("root"): A_red_temp_v0 = Ts.sblock_alloc_buffer((1,)) A_red_temp_v1 = Ts.sblock_alloc_buffer((1,)) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_reduction_block.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_reduction_block.py index b515cb06c6bd..0d3def88a67f 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_reduction_block.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_reduction_block.py @@ -51,7 +51,7 @@ def _create_context(mod, target) -> ms.TuneContext: @tvm.script.ir_module class Matmul_before_rewrite: @Ts.prim_func - def main(A: T.Buffer([512, 512], dtype='float32'), B: T.Buffer([512, 512], dtype='float32'), C: T.Buffer([512, 512], dtype='float32')) -> None: + def main(A: T.Tensor([512, 512], dtype='float32'), B: T.Tensor([512, 512], dtype='float32'), C: T.Tensor([512, 512], dtype='float32')) -> None: C_local = Ts.sblock_alloc_buffer([512, 512], dtype="float32", scope="local") A_shared = Ts.sblock_alloc_buffer([512, 512], dtype="float32", scope="shared") @@ -100,7 +100,7 @@ def main(A: T.Buffer([512, 512], dtype='float32'), B: T.Buffer([512, 512], dtype @tvm.script.ir_module class Matmul_after_rewrite: @Ts.prim_func - def main(A: T.Buffer([512, 512], dtype='float32'), B: T.Buffer([512, 512], dtype='float32'), C: T.Buffer([512, 512], dtype='float32')) -> None: + def main(A: T.Tensor([512, 512], dtype='float32'), B: T.Tensor([512, 512], dtype='float32'), C: T.Tensor([512, 512], dtype='float32')) -> None: C_local = Ts.sblock_alloc_buffer([512, 512], dtype="float32", scope="local") A_shared = Ts.sblock_alloc_buffer([512, 512], dtype="float32", scope="shared") @@ -154,7 +154,7 @@ def main(A: T.Buffer([512, 512], dtype='float32'), B: T.Buffer([512, 512], dtype @tvm.script.ir_module class Softmax_cross_thread_reduction: @Ts.prim_func - def main(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def main(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T_softmax_maxelem_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") T_softmax_expsum_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") for i0 in T.serial(256): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_tensorize.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_tensorize.py index 43fef33d1c23..2d6b81033822 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_tensorize.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_tensorize.py @@ -27,9 +27,9 @@ class Conv2dNCHWcVNNIModuleTiled: @Ts.prim_func def main( - placeholder: T.Buffer((1, 4, 56, 56, 16), "uint8"), - placeholder_1: T.Buffer((16, 4, 1, 1, 4, 16, 4), "int8"), - conv2d_NCHWc_int8: T.Buffer((1, 16, 56, 56, 16), "int32"), + placeholder: T.Tensor((1, 4, 56, 56, 16), "uint8"), + placeholder_1: T.Tensor((16, 4, 1, 1, 4, 16, 4), "int8"), + conv2d_NCHWc_int8: T.Tensor((1, 16, 56, 56, 16), "int32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for ( @@ -147,9 +147,9 @@ def main( class Conv2dNCHWcVNNIModuleTensorized: @Ts.prim_func def main( - placeholder: T.Buffer((1, 4, 56, 56, 16), "uint8"), - placeholder_1: T.Buffer((16, 4, 1, 1, 4, 16, 4), "int8"), - conv2d_NCHWc_int8: T.Buffer((1, 16, 56, 56, 16), "int32"), + placeholder: T.Tensor((1, 4, 56, 56, 16), "uint8"), + placeholder_1: T.Tensor((16, 4, 1, 1, 4, 16, 4), "int8"), + conv2d_NCHWc_int8: T.Tensor((1, 16, 56, 56, 16), "int32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -249,9 +249,9 @@ def main( class DenseDP4ATiled: @Ts.prim_func def main( - X: T.Buffer((128, 128), "int8"), - W: T.Buffer((128, 128), "int8"), - compute: T.Buffer((128, 128), "int32"), + X: T.Tensor((128, 128), "int8"), + W: T.Tensor((128, 128), "int8"), + compute: T.Tensor((128, 128), "int32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) compute_local = Ts.sblock_alloc_buffer([128, 128], dtype="int32", scope="local") @@ -339,9 +339,9 @@ def main( class DenseDP4ATensorized: @Ts.prim_func def main( - X: T.Buffer((128, 128), "int8"), - W: T.Buffer((128, 128), "int8"), - compute: T.Buffer((128, 128), "int32"), + X: T.Tensor((128, 128), "int8"), + W: T.Tensor((128, 128), "int8"), + compute: T.Tensor((128, 128), "int32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_unbound_block.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_unbound_block.py index 100087058788..14a58be27b95 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_unbound_block.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_rewrite_unbound_block.py @@ -50,7 +50,7 @@ def _create_context(mod, target) -> ms.TuneContext: class Before_cooperative_fetch: @Ts.prim_func def main( - A: T.Buffer([512, 512], dtype="float32"), B: T.Buffer([512, 512], dtype="float32") + A: T.Tensor([512, 512], dtype="float32"), B: T.Tensor([512, 512], dtype="float32") ) -> None: for i, j in T.grid(512, 512): with Ts.sblock("C"): @@ -62,7 +62,7 @@ def main( class After_cooperative_fetch: @Ts.prim_func def main( - A: T.Buffer([512, 512], dtype="float32"), B: T.Buffer([512, 512], dtype="float32") + A: T.Tensor([512, 512], dtype="float32"), B: T.Tensor([512, 512], dtype="float32") ) -> None: for i_j_fused_0 in T.thread_binding(256, thread="blockIdx.x"): for i_j_fused_1 in T.thread_binding(1024, thread="threadIdx.x"): @@ -75,7 +75,7 @@ def main( @tvm.script.ir_module class Before_norm_bmn: @Ts.prim_func - def main(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer((1,), "float32")) -> None: + def main(A: T.Tensor((1, 256, 256), "float32"), D: T.Tensor((1,), "float32")) -> None: C = Ts.sblock_alloc_buffer([1], dtype="float32") for i0, i1, i2 in T.grid(1, 256, 256): with Ts.sblock("C"): @@ -92,7 +92,7 @@ def main(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer((1,), "float32")) -> @tvm.script.ir_module class After_norm_bmn: @Ts.prim_func - def main(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer((1,), "float32")) -> None: + def main(A: T.Tensor((1, 256, 256), "float32"), D: T.Tensor((1,), "float32")) -> None: C = Ts.sblock_alloc_buffer([1], dtype="float32") for i0_fused_0 in T.thread_binding(1, thread="blockIdx.x"): for i0_fused_1 in T.thread_binding(1, thread="threadIdx.x"): @@ -114,7 +114,7 @@ def main(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer((1,), "float32")) -> class Bert_fused_reshape_transpose_reshape: @Ts.prim_func def main( - placeholder: T.Buffer((12, 64, 64), "float32"), T_reshape: T.Buffer((64, 768), "float32") + placeholder: T.Tensor((12, 64, 64), "float32"), T_reshape: T.Tensor((64, 768), "float32") ) -> None: for i0_i1_fused_0, i0_i1_fused_1 in T.grid(1536, 32): with Ts.sblock("T_reshape_1"): @@ -133,7 +133,7 @@ def main( class Bert_fused_reshape_transpose_reshape_large: @Ts.prim_func def main( - placeholder: T.Buffer((12, 64, 64), "float32"), T_reshape: T.Buffer((64, 768), "float32") + placeholder: T.Tensor((12, 64, 64), "float32"), T_reshape: T.Tensor((64, 768), "float32") ) -> None: for i0_i1_fused_0, i0_i1_fused_1 in T.grid(1536000, 32): with Ts.sblock("T_reshape_1"): @@ -152,7 +152,7 @@ def main( class Bert_fused_reshape_transpose_reshape_after_rub: @Ts.prim_func def main( - placeholder: T.Buffer((12, 64, 64), "float32"), T_reshape: T.Buffer((64, 768), "float32") + placeholder: T.Tensor((12, 64, 64), "float32"), T_reshape: T.Tensor((64, 768), "float32") ) -> None: for i0_i1_fused_0_i0_i1_fused_1_fused_0 in T.thread_binding(48, thread="blockIdx.x"): for i0_i1_fused_0_i0_i1_fused_1_fused_1 in T.thread_binding(1024, thread="threadIdx.x"): @@ -186,7 +186,7 @@ def main( class Bert_fused_reshape_transpose_reshape_after_rub_large: @Ts.prim_func def main( - placeholder: T.Buffer((12, 64, 64), "float32"), T_reshape: T.Buffer((64, 768), "float32") + placeholder: T.Tensor((12, 64, 64), "float32"), T_reshape: T.Tensor((64, 768), "float32") ) -> None: # body # with Ts.sblock("root") @@ -233,7 +233,7 @@ def main( @Ts.prim_func def before_unrolled_loop( - placeholder: T.Buffer((1, 56, 56, 64), "float32"), + placeholder: T.Tensor((1, 56, 56, 64), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -258,7 +258,7 @@ def before_unrolled_loop( @Ts.prim_func def after_unrolled_loop( - placeholder: T.Buffer((1, 56, 56, 64), "float32"), + placeholder: T.Tensor((1, 56, 56, 64), "float32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body 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 b90100292193..8b0ff8300fda 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 @@ -50,7 +50,7 @@ def _create_context(mod, target) -> ms.TuneContext: @tvm.script.ir_module class Conv2dCuda0: @Ts.prim_func - def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * 14 * 512 * 256], dtype='float32')) -> None: + def main(A: T.Tensor([14 * 14 * 256 * 256], dtype='float32'), B: T.Tensor([14 * 14 * 512 * 256], dtype='float32')) -> None: # function attr dict T.func_attr({"global_symbol": "main", "T.noalias": True}) # var definition @@ -62,9 +62,9 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * # body T.launch_thread(blockIdx_z, 196) - B_local = T.decl_buffer([64], "float32", scope="local") - Apad_shared = T.decl_buffer([512], "float32", scope="shared") - Apad_shared_local = T.decl_buffer([8], "float32", scope="local") + B_local = T.decl_tensor([64], "float32", scope="local") + Apad_shared = T.decl_tensor([512], "float32", scope="shared") + Apad_shared_local = T.decl_tensor([8], "float32", scope="local") T.launch_thread(blockIdx_y, 8) T.launch_thread(blockIdx_x, 4) T.launch_thread(threadIdx_y, 8) @@ -90,7 +90,7 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * @tvm.script.ir_module class Conv2dCuda1: @Ts.prim_func - def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * 14 * 512 * 256], dtype='float32')) -> None: + def main(A: T.Tensor([14 * 14 * 256 * 256], dtype='float32'), B: T.Tensor([14 * 14 * 512 * 256], dtype='float32')) -> None: # function attr dict T.func_attr({"global_symbol": "main", "T.noalias": True}) # var definition @@ -102,9 +102,9 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * # body T.launch_thread(blockIdx_z, 196) - B_local = T.decl_buffer([6400000], "float32", scope="local") - Apad_shared = T.decl_buffer([512], "float32", scope="shared") - Apad_shared_local = T.decl_buffer([8], "float32", scope="local") + B_local = T.decl_tensor([6400000], "float32", scope="local") + Apad_shared = T.decl_tensor([512], "float32", scope="shared") + Apad_shared_local = T.decl_tensor([8], "float32", scope="local") T.launch_thread(blockIdx_y, 8) T.launch_thread(blockIdx_x, 4) T.launch_thread(threadIdx_y, 8) @@ -134,7 +134,7 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * @tvm.script.ir_module class Conv2dCuda2: @Ts.prim_func - def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * 14 * 512 * 256], dtype='float32')) -> None: + def main(A: T.Tensor([14 * 14 * 256 * 256], dtype='float32'), B: T.Tensor([14 * 14 * 512 * 256], dtype='float32')) -> None: # function attr dict T.func_attr({"global_symbol": "main", "T.noalias": True}) # var definition @@ -146,9 +146,9 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * # body T.launch_thread(blockIdx_z, 196) - B_local = T.decl_buffer([64], "float32", scope="local") - Apad_shared = T.decl_buffer([512000], "float32", scope="shared") - Apad_shared_local = T.decl_buffer([8], "float32", scope="local") + B_local = T.decl_tensor([64], "float32", scope="local") + Apad_shared = T.decl_tensor([512000], "float32", scope="shared") + Apad_shared_local = T.decl_tensor([8], "float32", scope="local") T.launch_thread(blockIdx_y, 8) T.launch_thread(blockIdx_x, 4) T.launch_thread(threadIdx_y, 8) @@ -178,7 +178,7 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * @tvm.script.ir_module class Conv2dCuda3: @Ts.prim_func - def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * 14 * 512 * 256], dtype='float32')) -> None: + def main(A: T.Tensor([14 * 14 * 256 * 256], dtype='float32'), B: T.Tensor([14 * 14 * 512 * 256], dtype='float32')) -> None: # function attr dict T.func_attr({"global_symbol": "main", "T.noalias": True}) # var definition @@ -190,9 +190,9 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * # body T.launch_thread(blockIdx_z, 196) - B_local = T.decl_buffer([64], "float32", scope="local") - Apad_shared = T.decl_buffer([512], "float32", scope="shared") - Apad_shared_local = T.decl_buffer([8], "float32", scope="local") + B_local = T.decl_tensor([64], "float32", scope="local") + Apad_shared = T.decl_tensor([512], "float32", scope="shared") + Apad_shared_local = T.decl_tensor([8], "float32", scope="local") T.launch_thread(blockIdx_y, 8) T.launch_thread(blockIdx_x, 4) T.launch_thread(threadIdx_y, 8) @@ -216,7 +216,7 @@ def main(A: T.Buffer([14 * 14 * 256 * 256], dtype='float32'), B: T.Buffer([14 * B[blockIdx_z * 131072 + blockIdx_y * 16384 + threadIdx_y * 2048 + ff_inner_inner_inner * 256 + blockIdx_x * 64 + threadIdx_x * 8 + nn_inner_inner_inner] = B_local[ff_inner_inner_inner * 8 + nn_inner_inner_inner] @Ts.prim_func -def GmmCuda0(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), "float32"), Z: T.Buffer((1, 128, 128), "float32")) -> None: +def GmmCuda0(X: T.Tensor((1, 128, 128), "float32"), Y: T.Tensor((1, 128, 128), "float32"), Z: T.Tensor((1, 128, 128), "float32")) -> None: Z_local = Ts.sblock_alloc_buffer([1, 128, 128], dtype="float32", scope="local") X_shared = Ts.sblock_alloc_buffer([1, 128, 128], dtype="float32", scope="shared") Y_shared = Ts.sblock_alloc_buffer([1, 128, 128], dtype="float32", scope="shared") @@ -270,7 +270,7 @@ def GmmCuda0(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), " Z[v0, v1, v2] = Z_local[v0, v1, v2] @Ts.prim_func -def GmmCuda1(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), "float32"), Z: T.Buffer((1, 128, 128), "float32")) -> None: +def GmmCuda1(X: T.Tensor((1, 128, 128), "float32"), Y: T.Tensor((1, 128, 128), "float32"), Z: T.Tensor((1, 128, 128), "float32")) -> None: Z_local = Ts.sblock_alloc_buffer([1, 128, 128], dtype="float32", scope="local") X_shared = Ts.sblock_alloc_buffer([1, 128, 128], dtype="float32", scope="shared") Y_shared = Ts.sblock_alloc_buffer([1, 128, 128], dtype="float32", scope="shared") @@ -328,7 +328,7 @@ def GmmCuda1(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), " Z[v0, v1, v2] = Z_local[v0, v1, v2] @Ts.prim_func -def GmmCuda2(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), "float32"), Z: T.Buffer((1, 128, 128), "float32")) -> None: +def GmmCuda2(X: T.Tensor((1, 128, 128), "float32"), Y: T.Tensor((1, 128, 128), "float32"), Z: T.Tensor((1, 128, 128), "float32")) -> None: Z_local = Ts.sblock_alloc_buffer([1, 128, 128], dtype="float32", scope="local") X_shared = Ts.sblock_alloc_buffer([1, 128, 128], dtype="float32", scope="shared") Y_shared = Ts.sblock_alloc_buffer([1, 128, 128], dtype="float32", scope="shared") @@ -394,9 +394,9 @@ def GmmCuda2(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), " @Ts.prim_func def GMMCUDATensorCore( - X: T.Buffer((1024, 1024), "float16"), - Y: T.Buffer((1024, 1024), "float16"), - Z: T.Buffer((1024, 1024), "float32"), + X: T.Tensor((1024, 1024), "float16"), + Y: T.Tensor((1024, 1024), "float16"), + Z: T.Tensor((1024, 1024), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_vtcm_limit.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_vtcm_limit.py index 76dd45f1eed0..5d14da347ddb 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_vtcm_limit.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_postproc_verify_vtcm_limit.py @@ -44,7 +44,7 @@ def _create_context(mod, target) -> ms.TuneContext: @tvm.script.ir_module class Conv2dNCHWcVTCM: @Ts.prim_func - def main(p0: T.Buffer((T.int64(1), T.int64(2), T.int64(56), T.int64(56), T.int64(32)), "uint8"), p1: T.Buffer((T.int64(2), T.int64(2), T.int64(3), T.int64(3), T.int64(8), T.int64(32), T.int64(4)), "uint8"), conv2d_NCHWc_int8: T.Buffer((T.int64(1), T.int64(2), T.int64(54), T.int64(54), T.int64(32)), "int32")): + def main(p0: T.Tensor((T.int64(1), T.int64(2), T.int64(56), T.int64(56), T.int64(32)), "uint8"), p1: T.Tensor((T.int64(2), T.int64(2), T.int64(3), T.int64(3), T.int64(8), T.int64(32), T.int64(4)), "uint8"), conv2d_NCHWc_int8: T.Tensor((T.int64(1), T.int64(2), T.int64(54), T.int64(54), T.int64(32)), "int32")): T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) p0_global_vtcm = Ts.sblock_alloc_buffer([T.int64(1), T.int64(2), T.int64(56), T.int64(56), T.int64(32)], dtype="uint8", scope="global.vtcm") p1_global_vtcm = Ts.sblock_alloc_buffer([T.int64(2), T.int64(2), T.int64(3), T.int64(3), T.int64(8), T.int64(32), T.int64(4)], dtype="uint8", scope="global.vtcm") diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_runner.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_runner.py index 775c3bc238d5..715bb1f2250c 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_runner.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_runner.py @@ -73,9 +73,9 @@ class MatmulModule: @Ts.prim_func def main( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -91,9 +91,9 @@ def main( class MatmulReluModule: @Ts.prim_func def main( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -114,7 +114,7 @@ def main( class BatchMatmulModule: @Ts.prim_func def main( - A: T.Buffer([16, 32, 32]), B: T.Buffer([16, 32, 32]), C: T.Buffer([16, 32, 32]) + A: T.Tensor([16, 32, 32]), B: T.Tensor([16, 32, 32]), C: T.Tensor([16, 32, 32]) ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -130,7 +130,7 @@ def main( class AddModule: @Ts.prim_func def main( - A: T.Buffer([32], "float32"), B: T.Buffer([32], "float32"), C: T.Buffer([32], "float32") + A: T.Tensor([32], "float32"), B: T.Tensor([32], "float32"), C: T.Tensor([32], "float32") ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -145,9 +145,9 @@ def main( class MatmulHugeModule: @Ts.prim_func def main( - A: T.Buffer((4096, 4096), "float32"), - B: T.Buffer((4096, 4096), "float32"), - C: T.Buffer((4096, 4096), "float32"), + A: T.Tensor((4096, 4096), "float32"), + B: T.Tensor((4096, 4096), "float32"), + C: T.Tensor((4096, 4096), "float32"), ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_add_rfactor.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_add_rfactor.py index e18d870d0945..35416b9ca3ca 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_add_rfactor.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_add_rfactor.py @@ -30,9 +30,9 @@ def test_cpu_matmul(): @Ts.prim_func def cpu_matmul_0( - A: T.Buffer((4, 512), "float32"), - B: T.Buffer((512, 4), "float32"), - C: T.Buffer((4, 4), "float32"), + A: T.Tensor((4, 512), "float32"), + B: T.Tensor((512, 4), "float32"), + C: T.Tensor((4, 4), "float32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0, i1, i2 in T.grid(4, 4, 512): @@ -46,9 +46,9 @@ def cpu_matmul_0( @Ts.prim_func def cpu_matmul_1( - A: T.Buffer((4, 512), "float32"), - B: T.Buffer((512, 4), "float32"), - C: T.Buffer((4, 4), "float32"), + A: T.Tensor((4, 512), "float32"), + B: T.Tensor((512, 4), "float32"), + C: T.Tensor((4, 4), "float32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) C_rf = Ts.sblock_alloc_buffer([4, 4, 128], dtype="float32") @@ -74,9 +74,9 @@ def cpu_matmul_1( @Ts.prim_func def cpu_matmul_2( - A: T.Buffer((4, 512), "float32"), - B: T.Buffer((512, 4), "float32"), - C: T.Buffer((4, 4), "float32"), + A: T.Tensor((4, 512), "float32"), + B: T.Tensor((512, 4), "float32"), + C: T.Tensor((4, 4), "float32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) C_rf = Ts.sblock_alloc_buffer([4, 4, 4], dtype="float32") @@ -125,10 +125,10 @@ def cpu_matmul_2( def test_cpu_argmax(): @Ts.prim_func def argmax( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1 in T.grid(128, 128): with Ts.sblock("argmax"): @@ -150,10 +150,10 @@ def argmax( @Ts.prim_func def argmax_0( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer(128, "int32"), - argmax_v1: T.Buffer(128, "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor(128, "int32"), + argmax_v1: T.Tensor(128, "float32"), ) -> None: for i0, i1 in T.grid(128, 128): with Ts.sblock("argmax"): @@ -174,10 +174,10 @@ def argmax_0( @Ts.prim_func def argmax_1( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer(128, "int32"), - argmax_v1: T.Buffer(128, "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor(128, "int32"), + argmax_v1: T.Tensor(128, "float32"), ) -> None: argmax_v0_rf = Ts.sblock_alloc_buffer([128, 16], dtype="int32") argmax_v1_rf = Ts.sblock_alloc_buffer([128, 16], dtype="float32") @@ -221,10 +221,10 @@ def argmax_1( @Ts.prim_func def argmax_2( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer(128, "int32"), - argmax_v1: T.Buffer(128, "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor(128, "int32"), + argmax_v1: T.Tensor(128, "float32"), ) -> None: # body # with Ts.sblock("root") diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_apply_custom_rule.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_apply_custom_rule.py index d1275d5cfea7..6e60df009abb 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_apply_custom_rule.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_apply_custom_rule.py @@ -30,9 +30,9 @@ class Matmul: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: T.func_attr({"global_symbol": "main"}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_bind.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_bind.py index f1e858caac65..410139078148 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_bind.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_bind.py @@ -28,7 +28,7 @@ @Ts.prim_func def element_wise( - A: T.Buffer([512, 512], dtype="float32"), B: T.Buffer([512, 512], dtype="float32") + A: T.Tensor([512, 512], dtype="float32"), B: T.Tensor([512, 512], dtype="float32") ) -> None: for i, j in T.grid(512, 512): with Ts.sblock("C"): @@ -38,9 +38,9 @@ def element_wise( @Ts.prim_func def reduction_loop_only( - A: T.Buffer(2, "float32"), - B: T.Buffer(2, "float32"), - C: T.Buffer((), "float32"), + A: T.Tensor(2, "float32"), + B: T.Tensor(2, "float32"), + C: T.Tensor((), "float32"), ) -> None: for i0 in T.serial(2): with Ts.sblock("C"): @@ -54,9 +54,9 @@ def reduction_loop_only( @Ts.prim_func def zero_dim_add( - A: T.Buffer((), "float32"), - B: T.Buffer((), "float32"), - C: T.Buffer((), "float32"), + A: T.Tensor((), "float32"), + B: T.Tensor((), "float32"), + C: T.Tensor((), "float32"), ) -> None: with Ts.sblock("C"): vi = Ts.axis.spatial(1, 0) @@ -66,8 +66,8 @@ def zero_dim_add( def test_cuda_element_wise(): @Ts.prim_func def elementwise_0( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), ) -> None: # body # with Ts.sblock("root") @@ -101,9 +101,9 @@ def elementwise_0( def test_cuda_reduction_loop_only(): @Ts.prim_func def reduction_loop_only_0( - A: T.Buffer(2, "float32"), - B: T.Buffer(2, "float32"), - C: T.Buffer((), "float32"), + A: T.Tensor(2, "float32"), + B: T.Tensor(2, "float32"), + C: T.Tensor((), "float32"), ) -> None: for u_fused_0 in T.thread_binding(1, thread="blockIdx.x"): for u_fused_1 in T.thread_binding(1, thread="threadIdx.x"): @@ -134,9 +134,9 @@ def reduction_loop_only_0( def test_cuda_zero_dim_add(): @Ts.prim_func def zero_dim_add_0( - A: T.Buffer((), "float32"), - B: T.Buffer((), "float32"), - C: T.Buffer((), "float32"), + A: T.Tensor((), "float32"), + B: T.Tensor((), "float32"), + C: T.Tensor((), "float32"), ) -> None: for u_fused_0 in T.thread_binding(1, thread="blockIdx.x"): for u_fused_1 in T.thread_binding(1, thread="threadIdx.x"): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_inline.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_inline.py index cc4189b71bef..c45a8b78907c 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_inline.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_auto_inline.py @@ -34,7 +34,7 @@ @tvm.script.ir_module class Conv2DBiasBnReLU: @Ts.prim_func - def main(X: T.Buffer([1, 512, 56, 56], dtype='float32'), W: T.Buffer([512, 512, 3, 3], dtype='float32'), B: T.Buffer([512, 1, 1], dtype='float32'), bn_scale: T.Buffer([512, 1, 1], dtype='float32'), bn_offset: T.Buffer([512, 1, 1], dtype='float32'), compute: T.Buffer([1, 512, 56, 56], dtype='float32')) -> None: + def main(X: T.Tensor([1, 512, 56, 56], dtype='float32'), W: T.Tensor([512, 512, 3, 3], dtype='float32'), B: T.Tensor([512, 1, 1], dtype='float32'), bn_scale: T.Tensor([512, 1, 1], dtype='float32'), bn_offset: T.Tensor([512, 1, 1], dtype='float32'), compute: T.Tensor([1, 512, 56, 56], dtype='float32')) -> None: pad_temp = Ts.sblock_alloc_buffer([1, 512, 58, 58], dtype="float32") compute_1 = Ts.sblock_alloc_buffer([1, 512, 56, 56], dtype="float32") @@ -71,7 +71,7 @@ def main(X: T.Buffer([1, 512, 56, 56], dtype='float32'), W: T.Buffer([512, 512, @tvm.script.ir_module class Conv2DBiasBnReLUInlined: @Ts.prim_func - def main(X: T.Buffer([1, 512, 56, 56], dtype='float32'), W: T.Buffer([512, 512, 3, 3], dtype='float32'), B: T.Buffer([512, 1, 1], dtype='float32'), bn_scale: T.Buffer([512, 1, 1], dtype='float32'), bn_offset: T.Buffer([512, 1, 1], dtype='float32'), compute: T.Buffer([1, 512, 56, 56], dtype='float32')) -> None: + def main(X: T.Tensor([1, 512, 56, 56], dtype='float32'), W: T.Tensor([512, 512, 3, 3], dtype='float32'), B: T.Tensor([512, 1, 1], dtype='float32'), bn_scale: T.Tensor([512, 1, 1], dtype='float32'), bn_offset: T.Tensor([512, 1, 1], dtype='float32'), compute: T.Tensor([1, 512, 56, 56], dtype='float32')) -> None: pad_temp = Ts.sblock_alloc_buffer([1, 512, 58, 58], dtype="float32") compute_1 = Ts.sblock_alloc_buffer([1, 512, 56, 56], dtype="float32") @@ -93,7 +93,7 @@ def main(X: T.Buffer([1, 512, 56, 56], dtype='float32'), W: T.Buffer([512, 512, @tvm.script.ir_module class MultiLevelTiledConv2D: @Ts.prim_func - def main(X: T.Buffer([1, 512, 56, 56], dtype='float32'), W: T.Buffer([512, 512, 3, 3], dtype='float32'), B: T.Buffer([512, 1, 1], dtype='float32'), bn_scale: T.Buffer([512, 1, 1], dtype='float32'), bn_offset: T.Buffer([512, 1, 1], dtype='float32'), compute: T.Buffer([1, 512, 56, 56], dtype='float32')) -> None: + def main(X: T.Tensor([1, 512, 56, 56], dtype='float32'), W: T.Tensor([512, 512, 3, 3], dtype='float32'), B: T.Tensor([512, 1, 1], dtype='float32'), bn_scale: T.Tensor([512, 1, 1], dtype='float32'), bn_offset: T.Tensor([512, 1, 1], dtype='float32'), compute: T.Tensor([1, 512, 56, 56], dtype='float32')) -> None: pad_temp = Ts.sblock_alloc_buffer([1, 512, 58, 58], dtype="float32") compute_1 = Ts.sblock_alloc_buffer([1, 512, 56, 56], dtype="float32") @@ -150,7 +150,7 @@ def main(X: T.Buffer([1, 512, 56, 56], dtype='float32'), W: T.Buffer([512, 512, @tvm.script.ir_module class MultiLevelTiledConv2DAfterInline: @Ts.prim_func - def main(X: T.Buffer((1, 512, 56, 56), "float32"), W: T.Buffer((512, 512, 3, 3), "float32"), B: T.Buffer((512, 1, 1), "float32"), bn_scale: T.Buffer((512, 1, 1), "float32"), bn_offset: T.Buffer((512, 1, 1), "float32"), compute: T.Buffer((1, 512, 56, 56), "float32")) -> None: + def main(X: T.Tensor((1, 512, 56, 56), "float32"), W: T.Tensor((512, 512, 3, 3), "float32"), B: T.Tensor((512, 1, 1), "float32"), bn_scale: T.Tensor((512, 1, 1), "float32"), bn_offset: T.Tensor((512, 1, 1), "float32"), compute: T.Tensor((1, 512, 56, 56), "float32")) -> None: compute_local = Ts.sblock_alloc_buffer([1, 512, 56, 56], dtype="float32", scope="local") for i0_0_i1_0_i2_0_i3_0_fused in T.thread_binding(224, thread="blockIdx.x"): for i0_1_i1_1_i2_1_i3_1_fused in T.thread_binding(2, thread="vthread.x"): @@ -177,7 +177,7 @@ def main(X: T.Buffer((1, 512, 56, 56), "float32"), W: T.Buffer((512, 512, 3, 3), @tvm.script.ir_module class SoftmaxBeforeInline: @Ts.prim_func - def main(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def main(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T_softmax_maxelem = Ts.sblock_alloc_buffer([256], dtype="float32") T_softmax_exp = Ts.sblock_alloc_buffer([256, 256], dtype="float32") T_softmax_expsum = Ts.sblock_alloc_buffer([256], dtype="float32") @@ -205,7 +205,7 @@ def main(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256) @tvm.script.ir_module class SoftmaxAfterInline: @Ts.prim_func - def main(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def main(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T_softmax_maxelem = Ts.sblock_alloc_buffer([256], dtype="float32") T_softmax_expsum = Ts.sblock_alloc_buffer([256], dtype="float32") for i0, i1 in T.grid(256, 256): @@ -229,10 +229,10 @@ def main(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256) class BeforePureSpatial: @Ts.prim_func def main( - placeholder: T.Buffer((1, 384), "int64"), - placeholder_1: T.Buffer((30522, 768), "float32"), - placeholder_2: T.Buffer((1, 384, 768), "float32"), - T_add: T.Buffer((1, 384, 768), "float32"), + placeholder: T.Tensor((1, 384), "int64"), + placeholder_1: T.Tensor((30522, 768), "float32"), + placeholder_2: T.Tensor((1, 384, 768), "float32"), + T_add: T.Tensor((1, 384, 768), "float32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) compile_engine_const = Ts.sblock_alloc_buffer([], dtype="int64") @@ -292,7 +292,7 @@ def main( @tvm.script.ir_module class AfterPureSpatial: @Ts.prim_func - def main(placeholder: T.Buffer((1, 384), "int64"), placeholder_1: T.Buffer((30522, 768), "float32"), placeholder_2: T.Buffer((1, 384, 768), "float32"), T_add: T.Buffer((1, 384, 768), "float32")) -> None: + def main(placeholder: T.Tensor((1, 384), "int64"), placeholder_1: T.Tensor((30522, 768), "float32"), placeholder_2: T.Tensor((1, 384, 768), "float32"), T_add: T.Tensor((1, 384, 768), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -307,7 +307,7 @@ def main(placeholder: T.Buffer((1, 384), "int64"), placeholder_1: T.Buffer((3052 @tvm.script.ir_module class ConstConsumer: @Ts.prim_func - def main(T_full: T.Buffer((1, 12, 4096), "int64")) -> None: + def main(T_full: T.Tensor((1, 12, 4096), "int64")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -322,7 +322,7 @@ def main(T_full: T.Buffer((1, 12, 4096), "int64")) -> None: @tvm.script.ir_module class Conv2dInt8: @Ts.prim_func - def main(p0: T.Buffer((16, 14, 14, 256), "int8"), p1: T.Buffer((1024, 1, 1, 256), "int8"), p2: T.Buffer((1, 1, 1, 1024), "int32"), p3: T.Buffer((1, 1, 1, 1024), "int32"), p4: T.Buffer(1024, "int32"), p5: T.Buffer(1024, "int32"), p6: T.Buffer(1024, "int32"), p7: T.Buffer(1, "int32"), p8: T.Buffer((16, 14, 14, 1024), "int32"), compute: T.Buffer((16, 14, 14, 1024), "int32")) -> None: + def main(p0: T.Tensor((16, 14, 14, 256), "int8"), p1: T.Tensor((1024, 1, 1, 256), "int8"), p2: T.Tensor((1, 1, 1, 1024), "int32"), p3: T.Tensor((1, 1, 1, 1024), "int32"), p4: T.Tensor(1024, "int32"), p5: T.Tensor(1024, "int32"), p6: T.Tensor(1024, "int32"), p7: T.Tensor(1, "int32"), p8: T.Tensor((16, 14, 14, 1024), "int32"), compute: T.Tensor((16, 14, 14, 1024), "int32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -498,7 +498,7 @@ def test_inline_constant_scalars_skip_output_block(): @tvm.script.ir_module class Full: @Ts.prim_func - def main(T_full: T.Buffer((), "float32")): + def main(T_full: T.Tensor((), "float32")): with Ts.sblock("T_full"): vi = Ts.axis.spatial(1, 0) Ts.reads() @@ -515,8 +515,8 @@ def test_no_inline_root_block(): class MaxReduction: @Ts.prim_func def main( - data: T.Buffer((8, 8), "float32"), - data_red: T.Buffer((), "float32"), + data: T.Tensor((8, 8), "float32"), + data_red: T.Tensor((), "float32"), ): T.func_attr({"tir.noalias": T.bool(True)}) with Ts.sblock("data_red"): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_cross_thread_reduction.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_cross_thread_reduction.py index dc0c298ac58d..5fa5fb54d2f7 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_cross_thread_reduction.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_cross_thread_reduction.py @@ -33,7 +33,7 @@ class Softmax_mn_after_inline: @Ts.prim_func def main( - A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32") + A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32") ) -> None: T_softmax_maxelem = Ts.sblock_alloc_buffer([256], dtype="float32") T_softmax_expsum = Ts.sblock_alloc_buffer([256], dtype="float32") @@ -63,8 +63,8 @@ def main( def test_gpu_softmax_mn(): @Ts.prim_func def softmax_mn_0( - A: T.Buffer((256, 256), "float32"), - T_softmax_norm: T.Buffer((256, 256), "float32"), + A: T.Tensor((256, 256), "float32"), + T_softmax_norm: T.Tensor((256, 256), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -105,7 +105,7 @@ def softmax_mn_0( @Ts.prim_func def softmax_mn_1( - A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32") + A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32") ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -157,7 +157,7 @@ def softmax_mn_1( @Ts.prim_func def softmax_mn_2( - A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32") + A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32") ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -207,7 +207,7 @@ def softmax_mn_2( @Ts.prim_func def softmax_mn_3( - A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32") + A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32") ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -295,7 +295,7 @@ def softmax_mn_3( def test_gpu_softmax_mn_after_inline(): @Ts.prim_func def softmax_mn_after_inline_0( - A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32") + A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32") ) -> None: T_softmax_maxelem = Ts.sblock_alloc_buffer([256], dtype="float32") T_softmax_expsum = Ts.sblock_alloc_buffer([256], dtype="float32") @@ -329,7 +329,7 @@ def softmax_mn_after_inline_0( @Ts.prim_func def softmax_mn_after_inline_1( - A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32") + A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32") ) -> None: T_softmax_maxelem = Ts.sblock_alloc_buffer([256], dtype="float32") T_softmax_expsum = Ts.sblock_alloc_buffer([256], dtype="float32") @@ -365,7 +365,7 @@ def softmax_mn_after_inline_1( @Ts.prim_func def softmax_mn_after_inline_2( - A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32") + A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32") ) -> None: T_softmax_maxelem = Ts.sblock_alloc_buffer([256], dtype="float32") T_softmax_expsum_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") @@ -409,7 +409,7 @@ def softmax_mn_after_inline_2( @Ts.prim_func def softmax_mn_after_inline_3( - A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32") + A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32") ) -> None: T_softmax_maxelem_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") T_softmax_expsum_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") @@ -492,7 +492,7 @@ def softmax_mn_after_inline_3( def test_gpu_batch_norm_bmn(): @Ts.prim_func - def batch_norm_bmn_0(A: T.Buffer((1, 512, 512), "float32"), D: T.Buffer(1, "float32")) -> None: + def batch_norm_bmn_0(A: T.Tensor((1, 512, 512), "float32"), D: T.Tensor(1, "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -514,7 +514,7 @@ def batch_norm_bmn_0(A: T.Buffer((1, 512, 512), "float32"), D: T.Buffer(1, "floa D[b] = T.sqrt(C[b]) @Ts.prim_func - def batch_norm_bmn_1(A: T.Buffer((1, 512, 512), "float32"), D: T.Buffer(1, "float32")) -> None: + def batch_norm_bmn_1(A: T.Tensor((1, 512, 512), "float32"), D: T.Tensor(1, "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -562,10 +562,10 @@ def batch_norm_bmn_1(A: T.Buffer((1, 512, 512), "float32"), D: T.Buffer(1, "floa @Ts.prim_func def argmax( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1 in T.grid(128, 128): with Ts.sblock("argmax"): @@ -588,10 +588,10 @@ def argmax( @Ts.prim_func def argmax_32( - idx: T.Buffer((1, 32), "int32"), - val: T.Buffer((1, 32), "float32"), - argmax_v0: T.Buffer((1,), "int32"), - argmax_v1: T.Buffer((1,), "float32"), + idx: T.Tensor((1, 32), "int32"), + val: T.Tensor((1, 32), "float32"), + argmax_v0: T.Tensor((1,), "int32"), + argmax_v1: T.Tensor((1,), "float32"), ) -> None: for i0, i1 in T.grid(1, 32): with Ts.sblock("argmax"): @@ -615,10 +615,10 @@ def argmax_32( def test_gpu_argmax(): @Ts.prim_func def argmax_0( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer(128, "int32"), - argmax_v1: T.Buffer(128, "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor(128, "int32"), + argmax_v1: T.Tensor(128, "float32"), ) -> None: # body # with Ts.sblock("root") @@ -641,10 +641,10 @@ def argmax_0( @Ts.prim_func def argmax_1( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer(128, "int32"), - argmax_v1: T.Buffer(128, "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor(128, "int32"), + argmax_v1: T.Tensor(128, "float32"), ) -> None: # body # with Ts.sblock("root") @@ -690,10 +690,10 @@ def argmax_1( def test_gpu_argmax_32(): @Ts.prim_func def argmax_0( - idx: T.Buffer((1, 32), "int32"), - val: T.Buffer((1, 32), "float32"), - argmax_v0: T.Buffer((1,), "int32"), - argmax_v1: T.Buffer((1,), "float32"), + idx: T.Tensor((1, 32), "int32"), + val: T.Tensor((1, 32), "float32"), + argmax_v0: T.Tensor((1,), "int32"), + argmax_v1: T.Tensor((1,), "float32"), ) -> None: # body # with Ts.sblock("root") @@ -716,10 +716,10 @@ def argmax_0( @Ts.prim_func def argmax_1( - idx: T.Buffer((1, 32), "int32"), - val: T.Buffer((1, 32), "float32"), - argmax_v0: T.Buffer((1,), "int32"), - argmax_v1: T.Buffer((1,), "float32"), + idx: T.Tensor((1, 32), "int32"), + val: T.Tensor((1, 32), "float32"), + argmax_v0: T.Tensor((1,), "int32"), + argmax_v1: T.Tensor((1,), "float32"), ) -> None: # body # with Ts.sblock("root") diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt.py index 4f1aa0e13811..3deeec7ab971 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt.py @@ -33,9 +33,9 @@ def test_cpu_matmul(): @Ts.prim_func def cpu_matmul_0( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - C: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + C: T.Tensor((512, 512), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -64,9 +64,9 @@ def cpu_matmul_0( @Ts.prim_func def cpu_matmul_1( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - C: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + C: T.Tensor((512, 512), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -95,9 +95,9 @@ def cpu_matmul_1( @Ts.prim_func def cpu_matmul_2( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - C: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + C: T.Tensor((512, 512), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -151,9 +151,9 @@ def cpu_matmul_2( def test_cpu_matmul_relu(): @Ts.prim_func def cpu_matmul_relu_0( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - compute: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + compute: T.Tensor((512, 512), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -182,9 +182,9 @@ def cpu_matmul_relu_0( @Ts.prim_func def cpu_matmul_relu_1( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - compute: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + compute: T.Tensor((512, 512), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -213,9 +213,9 @@ def cpu_matmul_relu_1( @Ts.prim_func def cpu_matmul_relu_2( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - compute: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + compute: T.Tensor((512, 512), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -275,9 +275,9 @@ def cpu_matmul_relu_2( def test_cuda_matmul(): @Ts.prim_func def cuda_matmul_0( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - C: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + C: T.Tensor((512, 512), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -381,9 +381,9 @@ def cuda_matmul_0( def test_cuda_matmul_relu(): @Ts.prim_func def cuda_matmul_relu_0( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - compute: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + compute: T.Tensor((512, 512), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -499,8 +499,8 @@ def cuda_matmul_relu_0( def test_cuda_sum_with_trivial_block_iter(): @Ts.prim_func def sum_with_trivial_block_iter( - A: T.Buffer((1, 64, 768), "float32"), - B: T.Buffer((1, 64, 1), "float32"), + A: T.Tensor((1, 64, 768), "float32"), + B: T.Tensor((1, 64, 1), "float32"), ) -> None: for i0, i1, i2, i3 in T.grid(1, 64, 1, 768): with Ts.sblock("sum"): @@ -525,9 +525,9 @@ def sum_with_trivial_block_iter( def test_multi_level_tiling_hexagon(): @Ts.prim_func def cpu_conv2d_nhwc( - inputs: T.Buffer((1, 56, 56, 64), "float16"), - weight: T.Buffer((3, 3, 64, 64), "float16"), - conv2d_nhwc: T.Buffer((1, 56, 56, 64), "float16"), + inputs: T.Tensor((1, 56, 56, 64), "float16"), + weight: T.Tensor((3, 3, 64, 64), "float16"), + conv2d_nhwc: T.Tensor((1, 56, 56, 64), "float16"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) PadInput = Ts.sblock_alloc_buffer((1, 58, 58, 64), "float16") @@ -630,9 +630,9 @@ def cpu_conv2d_nhwc( def test_cache_read_specify_consumer(): @Ts.prim_func def cache_read_specify_consumer_0( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - T_add: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + T_add: T.Tensor((512, 512), "float32"), ): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) C = Ts.sblock_alloc_buffer((512, 512)) @@ -740,8 +740,8 @@ def test_max_pool_blocked(): # fmt off @Ts.prim_func def pool_blocked_cache_read_write( - X: T.Buffer((1, 2, 8, 8, 8, 8, 32), "uint8"), - pool: T.Buffer((1, 2, 4, 4, 8, 8, 32), "uint8"), + X: T.Tensor((1, 2, 8, 8, 8, 8, 32), "uint8"), + pool: T.Tensor((1, 2, 4, 4, 8, 8, 32), "uint8"), ): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) pool_global = Ts.sblock_alloc_buffer((1, 2, 4, 4, 8, 8, 32), "uint8") diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt_intrin.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt_intrin.py index bb3c036ed24d..fe2478ae5676 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt_intrin.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt_intrin.py @@ -37,9 +37,9 @@ def test_x86_conv2d_nchwc( ): @Ts.prim_func def conv2d_nchwc( - placeholder: T.Buffer((1, 4, 56, 56, 16), "uint8"), - placeholder_1: T.Buffer((16, 4, 1, 1, 4, 16, 4), "int8"), - conv2d_NCHWc_int8: T.Buffer((1, 16, 56, 56, 16), "int32"), + placeholder: T.Tensor((1, 4, 56, 56, 16), "uint8"), + placeholder_1: T.Tensor((16, 4, 1, 1, 4, 16, 4), "int8"), + conv2d_NCHWc_int8: T.Tensor((1, 16, 56, 56, 16), "int32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0, i1, i2, i3, i4, i5, i6, i7, i8, i9 in T.grid(1, 16, 56, 56, 16, 1, 1, 4, 4, 4): @@ -74,7 +74,7 @@ def conv2d_nchwc( # fmt: off @Ts.prim_func - def x86_conv2d_nchwc_0(placeholder: T.Buffer((1, 4, 56, 56, 16), "uint8"), placeholder_1: T.Buffer((16, 4, 1, 1, 4, 16, 4), "int8"), conv2d_NCHWc_int8: T.Buffer((1, 16, 56, 56, 16), "int32")) -> None: + def x86_conv2d_nchwc_0(placeholder: T.Tensor((1, 4, 56, 56, 16), "uint8"), placeholder_1: T.Tensor((16, 4, 1, 1, 4, 16, 4), "int8"), conv2d_NCHWc_int8: T.Tensor((1, 16, 56, 56, 16), "int32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): conv2d_NCHWc_int8_global = Ts.sblock_alloc_buffer((1, 16, 56, 56, 16), "int32") @@ -120,7 +120,7 @@ def x86_conv2d_nchwc_0(placeholder: T.Buffer((1, 4, 56, 56, 16), "uint8"), place conv2d_NCHWc_int8[v0, v1, v2, v3, v4] = conv2d_NCHWc_int8_global[v0, v1, v2, v3, v4] @Ts.prim_func - def x86_conv2d_nchwc_1(placeholder: T.Buffer((1, 4, 56, 56, 16), "uint8"), placeholder_1: T.Buffer((16, 4, 1, 1, 4, 16, 4), "int8"), conv2d_NCHWc_int8: T.Buffer((1, 16, 56, 56, 16), "int32")) -> None: + def x86_conv2d_nchwc_1(placeholder: T.Tensor((1, 4, 56, 56, 16), "uint8"), placeholder_1: T.Tensor((16, 4, 1, 1, 4, 16, 4), "int8"), conv2d_NCHWc_int8: T.Tensor((1, 16, 56, 56, 16), "int32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): conv2d_NCHWc_int8_global = Ts.sblock_alloc_buffer((1, 16, 56, 56, 16), "int32") @@ -166,7 +166,7 @@ def x86_conv2d_nchwc_1(placeholder: T.Buffer((1, 4, 56, 56, 16), "uint8"), place conv2d_NCHWc_int8[v0, v1, v2, v3, v4] = conv2d_NCHWc_int8_global[v0, v1, v2, v3, v4] @Ts.prim_func - def x86_conv2d_nchwc_2(placeholder: T.Buffer((1, 4, 56, 56, 16), "uint8"), placeholder_1: T.Buffer((16, 4, 1, 1, 4, 16, 4), "int8"), conv2d_NCHWc_int8: T.Buffer((1, 16, 56, 56, 16), "int32")) -> None: + def x86_conv2d_nchwc_2(placeholder: T.Tensor((1, 4, 56, 56, 16), "uint8"), placeholder_1: T.Tensor((16, 4, 1, 1, 4, 16, 4), "int8"), conv2d_NCHWc_int8: T.Tensor((1, 16, 56, 56, 16), "int32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): for i0_0, i1_0, i2_0, i3_0, i4_0_0, i0_1, i1_1, i2_1, i3_1, i4_0_1, i5_0, i6_0, i7_0, i8_0, i9_0_0, i0_2, i1_2, i2_2, i3_2, i4_0_2, i5_1, i6_1, i7_1, i8_1, i9_0_1, i0_3, i1_3, i2_3, i3_3, i4_0_3 in T.grid(1, 8, 28, 56, 1, 1, 2, 1, 1, 1, 1, 1, 1, 4, 1, 1, 1, 2, 1, 1, 1, 1, 4, 1, 1, 1, 1, 1, 1, 1): @@ -306,9 +306,9 @@ def _dense(m, n, k, in_dtype, out_dtype): def test_dp4a_dense(): @Ts.prim_func def dp4a_dense_0( - X: T.Buffer((128, 128), "int8"), - W: T.Buffer((128, 128), "int8"), - compute: T.Buffer((128, 128), "int32"), + X: T.Tensor((128, 128), "int8"), + W: T.Tensor((128, 128), "int8"), + compute: T.Tensor((128, 128), "int32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt_tc.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt_tc.py index 21ff2080c04d..e441d8193917 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt_tc.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_mlt_tc.py @@ -85,7 +85,7 @@ def test_matmul_relu(shared_scope): # fmt: off @Ts.prim_func - def matmul_relu_0(A: T.Buffer((128, 128), "float16"), B: T.Buffer((128, 128), "float16"), compute: T.Buffer((128, 128), "float32")) -> None: + def matmul_relu_0(A: T.Tensor((128, 128), "float16"), B: T.Tensor((128, 128), "float16"), compute: T.Tensor((128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): C_reindex_shared = Ts.sblock_alloc_buffer((4, 8, 2, 1, 16, 16), scope=shared_scope) @@ -236,7 +236,7 @@ def matmul_relu_0(A: T.Buffer((128, 128), "float16"), B: T.Buffer((128, 128), "f def test_matmul_relu_with_fallback(): # fmt: off @Ts.prim_func - def matmul_relu_fallback_0(A: T.Buffer((128, 128), "float16"), B: T.Buffer((128, 128), "float16"), compute: T.Buffer((128, 128), "float32")) -> None: + def matmul_relu_fallback_0(A: T.Tensor((128, 128), "float16"), B: T.Tensor((128, 128), "float16"), compute: T.Tensor((128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): C_reindex_shared = Ts.sblock_alloc_buffer((4, 2, 2, 4, 16, 16), scope="shared") @@ -394,7 +394,7 @@ def test_conv2d(shared_scope): # fmt: off @Ts.prim_func - def conv2d_0(inputs: T.Buffer((1, 16, 16, 32), "float16"), weight: T.Buffer((3, 3, 32, 32), "float16"), conv2d_nhwc: T.Buffer((1, 16, 16, 32), "float32")): + def conv2d_0(inputs: T.Tensor((1, 16, 16, 32), "float16"), weight: T.Tensor((3, 3, 32, 32), "float16"), conv2d_nhwc: T.Tensor((1, 16, 16, 32), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): PadInput = Ts.sblock_alloc_buffer((1, 18, 18, 32), "float16") @@ -576,7 +576,7 @@ def test_matmul_relu_pipeline(shared_scope): # fmt: off @Ts.prim_func - def matmul_relu_pipeline_0(A: T.Buffer((128, 128), "float16"), B: T.Buffer((128, 128), "float16"), compute: T.Buffer((128, 128), "float32")) -> None: + def matmul_relu_pipeline_0(A: T.Tensor((128, 128), "float16"), B: T.Tensor((128, 128), "float16"), compute: T.Tensor((128, 128), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -757,7 +757,7 @@ def test_matmul_relu_non_tensorizable(): def test_padded_matmul_relu(): # fmt: off @Ts.prim_func - def padded_matmul_relu_0(A: T.Buffer((127, 127), "float16"), B: T.Buffer((127, 127), "float16"), compute: T.Buffer((127, 127), "float32")) -> None: + def padded_matmul_relu_0(A: T.Tensor((127, 127), "float16"), B: T.Tensor((127, 127), "float16"), compute: T.Tensor((127, 127), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) C_reindex_shared = Ts.sblock_alloc_buffer((4, 8, 2, 1, 16, 16), scope="shared") C_reindex_shared_wmma_accumulator = Ts.sblock_alloc_buffer((4, 8, 2, 1, 16, 16), scope="wmma.accumulator") @@ -905,7 +905,7 @@ def padded_matmul_relu_0(A: T.Buffer((127, 127), "float16"), B: T.Buffer((127, 1 def test_conv_1x1(): # fmt: off @Ts.prim_func - def conv2d_1x1_0(inputs: T.Buffer((1, 16, 16, 64), "float16"), weight: T.Buffer((1, 1, 64, 64), "float16"), conv2d_nhwc: T.Buffer((1, 16, 16, 64), "float32")): + def conv2d_1x1_0(inputs: T.Tensor((1, 16, 16, 64), "float16"), weight: T.Tensor((1, 1, 64, 64), "float16"), conv2d_nhwc: T.Tensor((1, 16, 16, 64), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): conv2d_nhwc_reindex_shared = Ts.sblock_alloc_buffer((2, 1, 8, 4, 16, 16), scope="shared") @@ -1063,7 +1063,7 @@ def conv2d_1x1_0(inputs: T.Buffer((1, 16, 16, 64), "float16"), weight: T.Buffer( def test_padded_conv(): # fmt: off @Ts.prim_func - def padded_conv2d_0(inputs: T.Buffer((1, 224, 224, 3), "float16"), weight: T.Buffer((7, 7, 3, 64), "float16"), conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32")): + def padded_conv2d_0(inputs: T.Tensor((1, 224, 224, 3), "float16"), weight: T.Tensor((7, 7, 3, 64), "float16"), conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): conv2d_nhwc_reindex_shared = Ts.sblock_alloc_buffer((56, 2, 14, 2, 16, 16), scope="shared") @@ -1215,7 +1215,7 @@ def padded_conv2d_0(inputs: T.Buffer((1, 224, 224, 3), "float16"), weight: T.Buf def test_padded_matmul_single_padded_input(): # fmt: off @Ts.prim_func - def padded_matmul_single_padded_input_0(A: T.Buffer((1023, 4096), "float16"), B: T.Buffer((4096, 1024), "float16"), C: T.Buffer((1023, 1024), "float32")): + def padded_matmul_single_padded_input_0(A: T.Tensor((1023, 4096), "float16"), B: T.Tensor((4096, 1024), "float16"), C: T.Tensor((1023, 1024), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): C_reindex_pad_shared = Ts.sblock_alloc_buffer((8, 32, 8, 2, 16, 16), scope="shared") @@ -1363,7 +1363,7 @@ def padded_matmul_single_padded_input_0(A: T.Buffer((1023, 4096), "float16"), B: def test_padded_matmul_no_padded_output(): # fmt: off @Ts.prim_func - def padded_matmul_no_padded_output_0(A: T.Buffer((1024, 4095), "float16"), B: T.Buffer((4095, 1024), "float16"), C: T.Buffer((1024, 1024), "float32")): + def padded_matmul_no_padded_output_0(A: T.Tensor((1024, 4095), "float16"), B: T.Tensor((4095, 1024), "float16"), C: T.Tensor((1024, 1024), "float32")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): C_reindex_shared = Ts.sblock_alloc_buffer((32, 16, 2, 4, 16, 16), scope="shared") diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_parallel_vectorize_unroll.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_parallel_vectorize_unroll.py index 3d40ecf663fe..71e3f9d85823 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_parallel_vectorize_unroll.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_parallel_vectorize_unroll.py @@ -32,7 +32,7 @@ @tvm.script.ir_module class Matmul: @Ts.prim_func - def main(A: T.Buffer((1024, 1024), 'float32'), B: T.Buffer((1024, 1024), 'float32'), C: T.Buffer((1024, 1024), 'float32')) -> None: + def main(A: T.Tensor((1024, 1024), 'float32'), B: T.Tensor((1024, 1024), 'float32'), C: T.Tensor((1024, 1024), 'float32')) -> None: T.func_attr({"global_symbol": "main"}) for i, j, k in T.grid(1024, 1024, 1024): @@ -45,7 +45,7 @@ def main(A: T.Buffer((1024, 1024), 'float32'), B: T.Buffer((1024, 1024), 'float3 @tvm.script.ir_module class ParallelizeVectorizeUnroll: @Ts.prim_func - def main(A: T.Buffer((1024, 1024), 'float32'), B: T.Buffer((1024, 1024), 'float32'), C: T.Buffer((1024, 1024), 'float32')) -> None: + def main(A: T.Tensor((1024, 1024), 'float32'), B: T.Tensor((1024, 1024), 'float32'), C: T.Tensor((1024, 1024), 'float32')) -> None: T.func_attr({"global_symbol": "main"}) with Ts.sblock("root"): @@ -63,7 +63,7 @@ def main(A: T.Buffer((1024, 1024), 'float32'), B: T.Buffer((1024, 1024), 'float3 @tvm.script.ir_module class PureSpatial: @Ts.prim_func - def main(placeholder: T.Buffer((1, 13, 13, 3, 85), "float32"), placeholder_1: T.Buffer((1, 26, 26, 3, 85), "float32"), placeholder_2: T.Buffer((1, 52, 52, 3, 85), "float32"), T_expand_dims: T.Buffer((1, 80, 10647), "float32")) -> None: + def main(placeholder: T.Tensor((1, 13, 13, 3, 85), "float32"), placeholder_1: T.Tensor((1, 26, 26, 3, 85), "float32"), placeholder_2: T.Tensor((1, 52, 52, 3, 85), "float32"), T_expand_dims: T.Tensor((1, 80, 10647), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) T_strided_slice_with_axes = Ts.sblock_alloc_buffer([1, 52, 52, 3, 1], dtype="float32") T_sigmoid = Ts.sblock_alloc_buffer([1, 52, 52, 3, 1], dtype="float32") @@ -219,9 +219,9 @@ def main(placeholder: T.Buffer((1, 13, 13, 3, 85), "float32"), placeholder_1: T. def test_parallel_vectorize_unroll(): @Ts.prim_func def Matmul_0( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main"}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_random_compute_location.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_random_compute_location.py index 1df8f7442c2f..21a3ced79ead 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_random_compute_location.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_schedule_rule_random_compute_location.py @@ -32,8 +32,8 @@ class Add: @Ts.prim_func def main( - A: T.Buffer([2048, 2048, 2048], dtype="float32"), - B: T.Buffer([2048, 2048, 2048], dtype="float32"), + A: T.Tensor([2048, 2048, 2048], dtype="float32"), + B: T.Tensor([2048, 2048, 2048], dtype="float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main"}) @@ -62,8 +62,8 @@ def main( def test_random_compute_location(): @Ts.prim_func def add_0( - A: T.Buffer((2048, 2048, 2048), "float32"), - B: T.Buffer((2048, 2048, 2048), "float32"), + A: T.Tensor((2048, 2048, 2048), "float32"), + B: T.Tensor((2048, 2048, 2048), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main"}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_search_strategy.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_search_strategy.py index 450c838dfc7c..1dfa624a47e0 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_search_strategy.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_search_strategy.py @@ -38,7 +38,7 @@ @tvm.script.ir_module class Matmul: @Ts.prim_func - def main(A: T.Buffer((32, 32), 'float32'), B: T.Buffer((32, 32), 'float32'), C: T.Buffer((32, 32), 'float32')) -> None: # type: ignore + def main(A: T.Tensor((32, 32), 'float32'), B: T.Tensor((32, 32), 'float32'), C: T.Tensor((32, 32), 'float32')) -> None: # type: ignore T.func_attr({"global_symbol": "main"}) for i, j, k in T.grid(32, 32, 32): @@ -51,7 +51,7 @@ def main(A: T.Buffer((32, 32), 'float32'), B: T.Buffer((32, 32), 'float32'), C: @tvm.script.ir_module class OtherBlock: @Ts.prim_func - def main(A: T.Buffer((32, 32), 'float32'), B: T.Buffer((32, 32), 'float32'), C: T.Buffer((32, 32), 'float32')) -> None: # type: ignore + def main(A: T.Tensor((32, 32), 'float32'), B: T.Tensor((32, 32), 'float32'), C: T.Tensor((32, 32), 'float32')) -> None: # type: ignore T.func_attr({"global_symbol": "main"}) for i, j, k in T.grid(32, 32, 32): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cpu.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cpu.py index 6fa9fa05a263..45c3f07beb11 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cpu.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cpu.py @@ -45,7 +45,7 @@ def _design_space(mod): def test_cpu_c1d(): # fmt: off @Ts.prim_func - def c1d_0(inputs: T.Buffer((1, 256, 64), "float32"), weight: T.Buffer((3, 64, 128), "float32"), conv1d_nlc: T.Buffer((1, 128, 128), "float32")): + def c1d_0(inputs: T.Tensor((1, 256, 64), "float32"), weight: T.Tensor((3, 64, 128), "float32"), conv1d_nlc: T.Tensor((1, 128, 128), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -81,7 +81,7 @@ def c1d_0(inputs: T.Buffer((1, 256, 64), "float32"), weight: T.Buffer((3, 64, 12 Ts.writes(conv1d_nlc[v0, v1, v2]) conv1d_nlc[v0, v1, v2] = conv1d_nlc_global[v0, v1, v2] @Ts.prim_func - def c1d_1(inputs: T.Buffer((1, 256, 64), "float32"), weight: T.Buffer((3, 64, 128), "float32"), conv1d_nlc: T.Buffer((1, 128, 128), "float32")) -> None: + def c1d_1(inputs: T.Tensor((1, 256, 64), "float32"), weight: T.Tensor((3, 64, 128), "float32"), conv1d_nlc: T.Tensor((1, 128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -121,7 +121,7 @@ def c1d_1(inputs: T.Buffer((1, 256, 64), "float32"), weight: T.Buffer((3, 64, 12 conv1d_nlc[v0, v1, v2] = conv1d_nlc_global[v0, v1, v2] @Ts.prim_func - def c1d_2(inputs: T.Buffer((1, 256, 64), "float32"), weight: T.Buffer((3, 64, 128), "float32"), conv1d_nlc: T.Buffer((1, 128, 128), "float32")) -> None: + def c1d_2(inputs: T.Tensor((1, 256, 64), "float32"), weight: T.Tensor((3, 64, 128), "float32"), conv1d_nlc: T.Tensor((1, 128, 128), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): @@ -184,7 +184,7 @@ def c1d_2(inputs: T.Buffer((1, 256, 64), "float32"), weight: T.Buffer((3, 64, 12 def test_cpu_c2d(): # fmt: off @Ts.prim_func - def c2d_0(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32")) -> None: + def c2d_0(inputs: T.Tensor((1, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -228,7 +228,7 @@ def c2d_0(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, Ts.writes(conv2d_nhwc[v0, v1, v2, v3]) conv2d_nhwc[v0, v1, v2, v3] = conv2d_nhwc_global[v0, v1, v2, v3] @Ts.prim_func - def c2d_1(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32")) -> None: + def c2d_1(inputs: T.Tensor((1, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -268,7 +268,7 @@ def c2d_1(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, Ts.writes(conv2d_nhwc[v0, v1, v2, v3]) conv2d_nhwc[v0, v1, v2, v3] = conv2d_nhwc_global[v0, v1, v2, v3] @Ts.prim_func - def c2d_2(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32")) -> None: + def c2d_2(inputs: T.Tensor((1, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -349,7 +349,7 @@ def c2d_2(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, def test_cpu_c3d(): # fmt: off @Ts.prim_func - def c3d_0(inputs: T.Buffer((1, 16, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 7, 3, 64), "float32"), conv3d_ndhwc: T.Buffer((1, 8, 112, 112, 64), "float32")) -> None: + def c3d_0(inputs: T.Tensor((1, 16, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 7, 3, 64), "float32"), conv3d_ndhwc: T.Tensor((1, 8, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -397,7 +397,7 @@ def c3d_0(inputs: T.Buffer((1, 16, 224, 224, 3), "float32"), weight: T.Buffer((7 Ts.writes(conv3d_ndhwc[v0, v1, v2, v3, v4]) conv3d_ndhwc[v0, v1, v2, v3, v4] = conv3d_ndhwc_global[v0, v1, v2, v3, v4] @Ts.prim_func - def c3d_1(inputs: T.Buffer((1, 16, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 7, 3, 64), "float32"), conv3d_ndhwc: T.Buffer((1, 8, 112, 112, 64), "float32")) -> None: + def c3d_1(inputs: T.Tensor((1, 16, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 7, 3, 64), "float32"), conv3d_ndhwc: T.Tensor((1, 8, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -445,7 +445,7 @@ def c3d_1(inputs: T.Buffer((1, 16, 224, 224, 3), "float32"), weight: T.Buffer((7 Ts.writes(conv3d_ndhwc[v0, v1, v2, v3, v4]) conv3d_ndhwc[v0, v1, v2, v3, v4] = conv3d_ndhwc_global[v0, v1, v2, v3, v4] @Ts.prim_func - def c3d_2(inputs: T.Buffer((1, 16, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 7, 3, 64), "float32"), conv3d_ndhwc: T.Buffer((1, 8, 112, 112, 64), "float32")) -> None: + def c3d_2(inputs: T.Tensor((1, 16, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 7, 3, 64), "float32"), conv3d_ndhwc: T.Tensor((1, 8, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -535,7 +535,7 @@ def c3d_2(inputs: T.Buffer((1, 16, 224, 224, 3), "float32"), weight: T.Buffer((7 def test_cpu_cap(): # fmt: off @Ts.prim_func - def cap_0(inputs: T.Buffer((1, 16, 16, 4, 4, 32), "float32"), weight: T.Buffer((3, 3, 4, 4, 32, 32), "float32"), conv2d_capsule_nhwijc: T.Buffer((1, 8, 8, 4, 4, 32), "float32")) -> None: + def cap_0(inputs: T.Tensor((1, 16, 16, 4, 4, 32), "float32"), weight: T.Tensor((3, 3, 4, 4, 32, 32), "float32"), conv2d_capsule_nhwijc: T.Tensor((1, 8, 8, 4, 4, 32), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -584,7 +584,7 @@ def cap_0(inputs: T.Buffer((1, 16, 16, 4, 4, 32), "float32"), weight: T.Buffer(( Ts.writes(conv2d_capsule_nhwijc[v0, v1, v2, v3, v4, v5]) conv2d_capsule_nhwijc[v0, v1, v2, v3, v4, v5] = conv2d_capsule_nhwijc_global[v0, v1, v2, v3, v4, v5] @Ts.prim_func - def cap_1(inputs: T.Buffer((1, 16, 16, 4, 4, 32), "float32"), weight: T.Buffer((3, 3, 4, 4, 32, 32), "float32"), conv2d_capsule_nhwijc: T.Buffer((1, 8, 8, 4, 4, 32), "float32")) -> None: + def cap_1(inputs: T.Tensor((1, 16, 16, 4, 4, 32), "float32"), weight: T.Tensor((3, 3, 4, 4, 32, 32), "float32"), conv2d_capsule_nhwijc: T.Tensor((1, 8, 8, 4, 4, 32), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -630,7 +630,7 @@ def cap_1(inputs: T.Buffer((1, 16, 16, 4, 4, 32), "float32"), weight: T.Buffer(( Ts.writes(conv2d_capsule_nhwijc[v0, v1, v2, v3, v4, v5]) conv2d_capsule_nhwijc[v0, v1, v2, v3, v4, v5] = conv2d_capsule_nhwijc_global[v0, v1, v2, v3, v4, v5] @Ts.prim_func - def cap_2(inputs: T.Buffer((1, 16, 16, 4, 4, 32), "float32"), weight: T.Buffer((3, 3, 4, 4, 32, 32), "float32"), conv2d_capsule_nhwijc: T.Buffer((1, 8, 8, 4, 4, 32), "float32")) -> None: + def cap_2(inputs: T.Tensor((1, 16, 16, 4, 4, 32), "float32"), weight: T.Tensor((3, 3, 4, 4, 32, 32), "float32"), conv2d_capsule_nhwijc: T.Tensor((1, 8, 8, 4, 4, 32), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -717,7 +717,7 @@ def cap_2(inputs: T.Buffer((1, 16, 16, 4, 4, 32), "float32"), weight: T.Buffer(( def test_cpu_dep(): # fmt: off @Ts.prim_func - def dep_0(placeholder: T.Buffer((1, 112, 112, 32), "float32"), placeholder_1: T.Buffer((1, 3, 3, 32), "float32"), depth_conv2d_nhwc: T.Buffer((1, 112, 112, 32), "float32")) -> None: + def dep_0(placeholder: T.Tensor((1, 112, 112, 32), "float32"), placeholder_1: T.Tensor((1, 3, 3, 32), "float32"), depth_conv2d_nhwc: T.Tensor((1, 112, 112, 32), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -756,7 +756,7 @@ def dep_0(placeholder: T.Buffer((1, 112, 112, 32), "float32"), placeholder_1: T. Ts.writes(depth_conv2d_nhwc[v0, v1, v2, v3]) depth_conv2d_nhwc[v0, v1, v2, v3] = depth_conv2d_nhwc_global[v0, v1, v2, v3] @Ts.prim_func - def dep_1(placeholder: T.Buffer((1, 112, 112, 32), "float32"), placeholder_1: T.Buffer((1, 3, 3, 32), "float32"), depth_conv2d_nhwc: T.Buffer((1, 112, 112, 32), "float32")) -> None: + def dep_1(placeholder: T.Tensor((1, 112, 112, 32), "float32"), placeholder_1: T.Tensor((1, 3, 3, 32), "float32"), depth_conv2d_nhwc: T.Tensor((1, 112, 112, 32), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -792,7 +792,7 @@ def dep_1(placeholder: T.Buffer((1, 112, 112, 32), "float32"), placeholder_1: T. Ts.writes(depth_conv2d_nhwc[v0, v1, v2, v3]) depth_conv2d_nhwc[v0, v1, v2, v3] = depth_conv2d_nhwc_global[v0, v1, v2, v3] @Ts.prim_func - def dep_2(placeholder: T.Buffer((1, 112, 112, 32), "float32"), placeholder_1: T.Buffer((1, 3, 3, 32), "float32"), depth_conv2d_nhwc: T.Buffer((1, 112, 112, 32), "float32")) -> None: + def dep_2(placeholder: T.Tensor((1, 112, 112, 32), "float32"), placeholder_1: T.Tensor((1, 3, 3, 32), "float32"), depth_conv2d_nhwc: T.Tensor((1, 112, 112, 32), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -866,7 +866,7 @@ def dep_2(placeholder: T.Buffer((1, 112, 112, 32), "float32"), placeholder_1: T. def test_cpu_dil(): # fmt: off @Ts.prim_func - def dil_0(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Buffer((1, 109, 109, 64), "float32")) -> None: + def dil_0(inputs: T.Tensor((1, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Tensor((1, 109, 109, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -909,7 +909,7 @@ def dil_0(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, Ts.writes(conv2d_nhwc[v0, v1, v2, v3]) conv2d_nhwc[v0, v1, v2, v3] = conv2d_nhwc_global[v0, v1, v2, v3] @Ts.prim_func - def dil_1(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Buffer((1, 109, 109, 64), "float32")) -> None: + def dil_1(inputs: T.Tensor((1, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Tensor((1, 109, 109, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -953,7 +953,7 @@ def dil_1(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, Ts.writes(conv2d_nhwc[v0, v1, v2, v3]) conv2d_nhwc[v0, v1, v2, v3] = conv2d_nhwc_global[v0, v1, v2, v3] @Ts.prim_func - def dil_2(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Buffer((1, 109, 109, 64), "float32")) -> None: + def dil_2(inputs: T.Tensor((1, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Tensor((1, 109, 109, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1032,7 +1032,7 @@ def dil_2(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, def test_cpu_gmm(): # fmt: off @Ts.prim_func - def gmm_0(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), "float32"), Z: T.Buffer((1, 128, 128), "float32")) -> None: + def gmm_0(X: T.Tensor((1, 128, 128), "float32"), Y: T.Tensor((1, 128, 128), "float32"), Z: T.Tensor((1, 128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1061,7 +1061,7 @@ def gmm_0(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), "flo Ts.writes(Z[v0, v1, v2]) Z[v0, v1, v2] = Z_global[v0, v1, v2] @Ts.prim_func - def gmm_1(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), "float32"), Z: T.Buffer((1, 128, 128), "float32")) -> None: + def gmm_1(X: T.Tensor((1, 128, 128), "float32"), Y: T.Tensor((1, 128, 128), "float32"), Z: T.Tensor((1, 128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1090,7 +1090,7 @@ def gmm_1(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), "flo Ts.writes(Z[v0, v1, v2]) Z[v0, v1, v2] = Z_global[v0, v1, v2] @Ts.prim_func - def gmm_2(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), "float32"), Z: T.Buffer((1, 128, 128), "float32")) -> None: + def gmm_2(X: T.Tensor((1, 128, 128), "float32"), Y: T.Tensor((1, 128, 128), "float32"), Z: T.Tensor((1, 128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1143,7 +1143,7 @@ def gmm_2(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), "flo def test_cpu_grp(): # fmt: off @Ts.prim_func - def grp_0(inputs: T.Buffer((1, 56, 56, 64), "float32"), weight: T.Buffer((3, 3, 16, 128), "float32"), conv2d_nhwc: T.Buffer((1, 28, 28, 128), "float32")) -> None: + def grp_0(inputs: T.Tensor((1, 56, 56, 64), "float32"), weight: T.Tensor((3, 3, 16, 128), "float32"), conv2d_nhwc: T.Tensor((1, 28, 28, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1187,7 +1187,7 @@ def grp_0(inputs: T.Buffer((1, 56, 56, 64), "float32"), weight: T.Buffer((3, 3, Ts.writes(conv2d_nhwc[v0, v1, v2, v3]) conv2d_nhwc[v0, v1, v2, v3] = conv2d_nhwc_global[v0, v1, v2, v3] @Ts.prim_func - def grp_1(inputs: T.Buffer((1, 56, 56, 64), "float32"), weight: T.Buffer((3, 3, 16, 128), "float32"), conv2d_nhwc: T.Buffer((1, 28, 28, 128), "float32")) -> None: + def grp_1(inputs: T.Tensor((1, 56, 56, 64), "float32"), weight: T.Tensor((3, 3, 16, 128), "float32"), conv2d_nhwc: T.Tensor((1, 28, 28, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1227,7 +1227,7 @@ def grp_1(inputs: T.Buffer((1, 56, 56, 64), "float32"), weight: T.Buffer((3, 3, Ts.writes(conv2d_nhwc[v0, v1, v2, v3]) conv2d_nhwc[v0, v1, v2, v3] = conv2d_nhwc_global[v0, v1, v2, v3] @Ts.prim_func - def grp_2(inputs: T.Buffer((1, 56, 56, 64), "float32"), weight: T.Buffer((3, 3, 16, 128), "float32"), conv2d_nhwc: T.Buffer((1, 28, 28, 128), "float32")) -> None: + def grp_2(inputs: T.Tensor((1, 56, 56, 64), "float32"), weight: T.Tensor((3, 3, 16, 128), "float32"), conv2d_nhwc: T.Tensor((1, 28, 28, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1306,7 +1306,7 @@ def grp_2(inputs: T.Buffer((1, 56, 56, 64), "float32"), weight: T.Buffer((3, 3, def test_cpu_t2d(): # fmt: off @Ts.prim_func - def t2d_0(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Buffer((1, 8, 8, 256), "float32")) -> None: + def t2d_0(inputs: T.Tensor((1, 4, 4, 512), "float32"), weight: T.Tensor((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Tensor((1, 8, 8, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1346,7 +1346,7 @@ def t2d_0(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 5 Ts.writes(conv2d_transpose_nhwc[v0, v1, v2, v3]) conv2d_transpose_nhwc[v0, v1, v2, v3] = conv2d_transpose_nhwc_global[v0, v1, v2, v3] @Ts.prim_func - def t2d_1(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Buffer((1, 8, 8, 256), "float32")) -> None: + def t2d_1(inputs: T.Tensor((1, 4, 4, 512), "float32"), weight: T.Tensor((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Tensor((1, 8, 8, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1387,7 +1387,7 @@ def t2d_1(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 5 Ts.writes(conv2d_transpose_nhwc[v0, v1, v2, v3]) conv2d_transpose_nhwc[v0, v1, v2, v3] = conv2d_transpose_nhwc_global[v0, v1, v2, v3] @Ts.prim_func - def t2d_2(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Buffer((1, 8, 8, 256), "float32")) -> None: + def t2d_2(inputs: T.Tensor((1, 4, 4, 512), "float32"), weight: T.Tensor((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Tensor((1, 8, 8, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1456,7 +1456,7 @@ def t2d_2(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 5 def test_cpu_nrm(): # fmt: off @Ts.prim_func - def nrm_0(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer(1, "float32")) -> None: + def nrm_0(A: T.Tensor((1, 256, 256), "float32"), D: T.Tensor(1, "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1487,7 +1487,7 @@ def nrm_0(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer(1, "float32")) -> N Ts.writes(D[v_b]) D[v_b] = T.sqrt(C[v_b]) @Ts.prim_func - def nrm_1(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer(1, "float32")) -> None: + def nrm_1(A: T.Tensor((1, 256, 256), "float32"), D: T.Tensor(1, "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1518,7 +1518,7 @@ def nrm_1(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer(1, "float32")) -> N Ts.writes(D[v_b]) D[v_b] = T.sqrt(C[v_b]) @Ts.prim_func - def nrm_2(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer(1, "float32")) -> None: + def nrm_2(A: T.Tensor((1, 256, 256), "float32"), D: T.Tensor(1, "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1569,7 +1569,7 @@ def nrm_2(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer(1, "float32")) -> N def test_cpu_sfm(): # fmt: off @Ts.prim_func - def sfm_0(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_0(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1620,7 +1620,7 @@ def sfm_0(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 Ts.sblock_attr({"axis": 1}) T_softmax_norm[v_i0, v_i1] = T.exp(A[v_i0, v_i1] - T_softmax_maxelem[v_i0]) / T_softmax_expsum[v_i0] @Ts.prim_func - def sfm_1(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_1(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1681,7 +1681,7 @@ def sfm_1(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 Ts.sblock_attr({"axis": 1}) T_softmax_norm[v_i0, v_i1] = T_softmax_exp[v_i0, v_i1] / T_softmax_expsum[v_i0] @Ts.prim_func - def sfm_2(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_2(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1722,7 +1722,7 @@ def sfm_2(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 Ts.sblock_attr({"axis": 1}) T_softmax_norm[v_i0, v_i1] = T.exp(A[v_i0, v_i1] - T_softmax_maxelem[v_i0]) / T_softmax_expsum[v_i0] @Ts.prim_func - def sfm_3(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_3(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1787,7 +1787,7 @@ def sfm_3(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 Ts.sblock_attr({"axis": 1}) T_softmax_norm[v_i0, v_i1] = T_softmax_exp[v_i0, v_i1] / T_softmax_expsum[v_i0] @Ts.prim_func - def sfm_4(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_4(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1847,7 +1847,7 @@ def sfm_4(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 Ts.sblock_attr({"axis": 1}) T_softmax_norm[v_i0, v_i1] = T_softmax_exp[v_i0, v_i1] / T_softmax_expsum[v_i0] @Ts.prim_func - def sfm_5(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_5(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1902,7 +1902,7 @@ def sfm_5(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 Ts.sblock_attr({"axis": 1}) T_softmax_norm[v_i0, v_i1] = T_softmax_exp[v_i0, v_i1] / T_softmax_expsum[v_i0] @Ts.prim_func - def sfm_6(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_6(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1946,7 +1946,7 @@ def sfm_6(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 Ts.sblock_attr({"axis": 1}) T_softmax_norm[v_i0, v_i1] = T.exp(A[v_i0, v_i1] - T_softmax_maxelem[v_i0]) / T_softmax_expsum[v_i0] @Ts.prim_func - def sfm_7(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_7(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1988,7 +1988,7 @@ def sfm_7(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 Ts.sblock_attr({"axis": 1}) T_softmax_norm[v_i0, v_i1] = T.exp(A[v_i0, v_i1] - T_softmax_maxelem[v_i0]) / T_softmax_expsum[v_i0] @Ts.prim_func - def sfm_8(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_8(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -2130,7 +2130,7 @@ def sfm_8(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 def test_cpu_cbr(): # fmt: off @Ts.prim_func - def cbr_0(data: T.Buffer((1, 224, 224, 3), "float32"), kernel: T.Buffer((7, 7, 3, 64), "float32"), bias: T.Buffer(64, "float32"), bn_offset: T.Buffer(64, "float32"), bn_scale: T.Buffer(64, "float32"), compute: T.Buffer((1, 112, 112, 64), "float32")) -> None: + def cbr_0(data: T.Tensor((1, 224, 224, 3), "float32"), kernel: T.Tensor((7, 7, 3, 64), "float32"), bias: T.Tensor(64, "float32"), bn_offset: T.Tensor(64, "float32"), bn_scale: T.Tensor(64, "float32"), compute: T.Tensor((1, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -2159,7 +2159,7 @@ def cbr_0(data: T.Buffer((1, 224, 224, 3), "float32"), kernel: T.Buffer((7, 7, 3 Ts.writes(compute[v_i0, v_i1, v_i2, v_i3]) compute[v_i0, v_i1, v_i2, v_i3] = T.max((Conv2dOutput[v_i0, v_i1, v_i2, v_i3] + bias[v_i3]) * bn_scale[v_i3] + bn_offset[v_i3], T.float32(0)) @Ts.prim_func - def cbr_1(data: T.Buffer((1, 224, 224, 3), "float32"), kernel: T.Buffer((7, 7, 3, 64), "float32"), bias: T.Buffer(64, "float32"), bn_offset: T.Buffer(64, "float32"), bn_scale: T.Buffer(64, "float32"), compute: T.Buffer((1, 112, 112, 64), "float32")) -> None: + def cbr_1(data: T.Tensor((1, 224, 224, 3), "float32"), kernel: T.Tensor((7, 7, 3, 64), "float32"), bias: T.Tensor(64, "float32"), bn_offset: T.Tensor(64, "float32"), bn_scale: T.Tensor(64, "float32"), compute: T.Tensor((1, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -2203,7 +2203,7 @@ def cbr_1(data: T.Buffer((1, 224, 224, 3), "float32"), kernel: T.Buffer((7, 7, 3 Ts.writes(compute[v_i0, v_i1, v_i2, v_i3]) compute[v_i0, v_i1, v_i2, v_i3] = T.max((Conv2dOutput[v_i0, v_i1, v_i2, v_i3] + bias[v_i3]) * bn_scale[v_i3] + bn_offset[v_i3], T.float32(0)) @Ts.prim_func - def cbr_2(data: T.Buffer((1, 224, 224, 3), "float32"), kernel: T.Buffer((7, 7, 3, 64), "float32"), bias: T.Buffer(64, "float32"), bn_offset: T.Buffer(64, "float32"), bn_scale: T.Buffer(64, "float32"), compute: T.Buffer((1, 112, 112, 64), "float32")) -> None: + def cbr_2(data: T.Tensor((1, 224, 224, 3), "float32"), kernel: T.Tensor((7, 7, 3, 64), "float32"), bias: T.Tensor(64, "float32"), bn_offset: T.Tensor(64, "float32"), bn_scale: T.Tensor(64, "float32"), compute: T.Tensor((1, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -2293,7 +2293,7 @@ def cbr_2(data: T.Buffer((1, 224, 224, 3), "float32"), kernel: T.Buffer((7, 7, 3 def test_cpu_tbg(): # fmt: off @Ts.prim_func - def tbg_0(query: T.Buffer((1, 128, 12, 64), "float32"), value: T.Buffer((1, 128, 12, 64), "float32"), C: T.Buffer((1, 12, 128, 128), "float32")) -> None: + def tbg_0(query: T.Tensor((1, 128, 12, 64), "float32"), value: T.Tensor((1, 128, 12, 64), "float32"), C: T.Tensor((1, 12, 128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -2345,7 +2345,7 @@ def tbg_0(query: T.Buffer((1, 128, 12, 64), "float32"), value: T.Buffer((1, 128, Ts.writes(C[v0, v1, v2, v3]) C[v0, v1, v2, v3] = C_global[v0, v1, v2, v3] @Ts.prim_func - def tbg_1(query: T.Buffer((1, 128, 12, 64), "float32"), value: T.Buffer((1, 128, 12, 64), "float32"), C: T.Buffer((1, 12, 128, 128), "float32")) -> None: + def tbg_1(query: T.Tensor((1, 128, 12, 64), "float32"), value: T.Tensor((1, 128, 12, 64), "float32"), C: T.Tensor((1, 12, 128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -2392,7 +2392,7 @@ def tbg_1(query: T.Buffer((1, 128, 12, 64), "float32"), value: T.Buffer((1, 128, Ts.writes(C[v0, v1, v2, v3]) C[v0, v1, v2, v3] = C_global[v0, v1, v2, v3] @Ts.prim_func - def tbg_2(query: T.Buffer((1, 128, 12, 64), "float32"), value: T.Buffer((1, 128, 12, 64), "float32"), C: T.Buffer((1, 12, 128, 128), "float32")) -> None: + def tbg_2(query: T.Tensor((1, 128, 12, 64), "float32"), value: T.Tensor((1, 128, 12, 64), "float32"), C: T.Tensor((1, 12, 128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cuda.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cuda.py index 9f51a0bcaa8f..b32004952c05 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cuda.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cuda.py @@ -45,7 +45,7 @@ def _design_space(mod): def test_cuda_c1d(): # fmt: off @Ts.prim_func - def c1d_0(inputs: T.Buffer((1, 256, 64), "float32"), weight: T.Buffer((3, 64, 128), "float32"), conv1d_nlc: T.Buffer((1, 128, 128), "float32")) -> None: + def c1d_0(inputs: T.Tensor((1, 256, 64), "float32"), weight: T.Tensor((3, 64, 128), "float32"), conv1d_nlc: T.Tensor((1, 128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -123,7 +123,7 @@ def c1d_0(inputs: T.Buffer((1, 256, 64), "float32"), weight: T.Buffer((3, 64, 12 def test_cuda_c2d(): # fmt: off @Ts.prim_func - def c2d_0(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32")) -> None: + def c2d_0(inputs: T.Tensor((1, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -207,7 +207,7 @@ def c2d_0(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, def test_cuda_c3d(): # fmt: off @Ts.prim_func - def c3d_0(inputs: T.Buffer((1, 16, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 7, 3, 64), "float32"), conv3d_ndhwc: T.Buffer((1, 8, 112, 112, 64), "float32")) -> None: + def c3d_0(inputs: T.Tensor((1, 16, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 7, 3, 64), "float32"), conv3d_ndhwc: T.Tensor((1, 8, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -297,7 +297,7 @@ def c3d_0(inputs: T.Buffer((1, 16, 224, 224, 3), "float32"), weight: T.Buffer((7 def test_cuda_cap(): # fmt: off @Ts.prim_func - def cap_0(inputs: T.Buffer((1, 16, 16, 4, 4, 32), "float32"), weight: T.Buffer((3, 3, 4, 4, 32, 32), "float32"), conv2d_capsule_nhwijc: T.Buffer((1, 8, 8, 4, 4, 32), "float32")) -> None: + def cap_0(inputs: T.Tensor((1, 16, 16, 4, 4, 32), "float32"), weight: T.Tensor((3, 3, 4, 4, 32, 32), "float32"), conv2d_capsule_nhwijc: T.Tensor((1, 8, 8, 4, 4, 32), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -391,7 +391,7 @@ def cap_0(inputs: T.Buffer((1, 16, 16, 4, 4, 32), "float32"), weight: T.Buffer(( def test_cuda_dep(): # fmt: off @Ts.prim_func - def dep_0(placeholder: T.Buffer((1, 112, 112, 32), "float32"), placeholder_1: T.Buffer((1, 3, 3, 32), "float32"), depth_conv2d_nhwc: T.Buffer((1, 112, 112, 32), "float32")) -> None: + def dep_0(placeholder: T.Tensor((1, 112, 112, 32), "float32"), placeholder_1: T.Tensor((1, 3, 3, 32), "float32"), depth_conv2d_nhwc: T.Tensor((1, 112, 112, 32), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -472,7 +472,7 @@ def dep_0(placeholder: T.Buffer((1, 112, 112, 32), "float32"), placeholder_1: T. def test_cuda_dil(): # fmt: off @Ts.prim_func - def dil_0(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Buffer((1, 109, 109, 64), "float32")) -> None: + def dil_0(inputs: T.Tensor((1, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Tensor((1, 109, 109, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -553,7 +553,7 @@ def dil_0(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, def test_cuda_gmm(): # fmt: off @Ts.prim_func - def gmm_0(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), "float32"), Z: T.Buffer((1, 128, 128), "float32")) -> None: + def gmm_0(X: T.Tensor((1, 128, 128), "float32"), Y: T.Tensor((1, 128, 128), "float32"), Z: T.Tensor((1, 128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -627,7 +627,7 @@ def gmm_0(X: T.Buffer((1, 128, 128), "float32"), Y: T.Buffer((1, 128, 128), "flo def test_cuda_grp(): # fmt: off @Ts.prim_func - def grp_0(inputs: T.Buffer((1, 56, 56, 64), "float32"), weight: T.Buffer((3, 3, 16, 128), "float32"), conv2d_nhwc: T.Buffer((1, 28, 28, 128), "float32")) -> None: + def grp_0(inputs: T.Tensor((1, 56, 56, 64), "float32"), weight: T.Tensor((3, 3, 16, 128), "float32"), conv2d_nhwc: T.Tensor((1, 28, 28, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -709,7 +709,7 @@ def grp_0(inputs: T.Buffer((1, 56, 56, 64), "float32"), weight: T.Buffer((3, 3, def test_cuda_t2d(): # fmt: off @Ts.prim_func - def t2d_0(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Buffer((1, 8, 8, 256), "float32")) -> None: + def t2d_0(inputs: T.Tensor((1, 4, 4, 512), "float32"), weight: T.Tensor((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Tensor((1, 8, 8, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -793,7 +793,7 @@ def t2d_0(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 5 def test_cuda_nrm(): # fmt: off @Ts.prim_func - def nrm_0(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer(1, "float32")) -> None: + def nrm_0(A: T.Tensor((1, 256, 256), "float32"), D: T.Tensor(1, "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -819,7 +819,7 @@ def nrm_0(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer(1, "float32")) -> N Ts.writes(D[v_b]) D[v_b] = T.sqrt(C[v_b]) @Ts.prim_func - def nrm_1(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer(1, "float32")) -> None: + def nrm_1(A: T.Tensor((1, 256, 256), "float32"), D: T.Tensor(1, "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -866,7 +866,7 @@ def nrm_1(A: T.Buffer((1, 256, 256), "float32"), D: T.Buffer(1, "float32")) -> N def test_cuda_sfm(): # fmt: off @Ts.prim_func - def sfm_0(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_0(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -906,7 +906,7 @@ def sfm_0(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 Ts.sblock_attr({"axis": 1}) T_softmax_norm[v_i0, v_i1] = T.exp(A[v_i0, v_i1] - T_softmax_maxelem[v_i0]) / T_softmax_expsum[v_i0] @Ts.prim_func - def sfm_1(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_1(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -946,7 +946,7 @@ def sfm_1(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 Ts.sblock_attr({"axis": 1}) T_softmax_norm[v_i0, v_i1] = T.exp(A[v_i0, v_i1] - T_softmax_maxelem[v_i0]) / T_softmax_expsum[v_i0] @Ts.prim_func - def sfm_2(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_2(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -988,7 +988,7 @@ def sfm_2(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 Ts.sblock_attr({"axis": 1}) T_softmax_norm[v_i0, v_i1] = T.exp(A[v_i0, v_i1] - T_softmax_maxelem[v_i0]) / T_softmax_expsum_shared[v_i0] @Ts.prim_func - def sfm_3(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32")) -> None: + def sfm_3(A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1065,7 +1065,7 @@ def sfm_3(A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256 def test_cuda_cbr(): # fmt: off @Ts.prim_func - def cbr_0(data: T.Buffer((1, 224, 224, 3), "float32"), kernel: T.Buffer((7, 7, 3, 64), "float32"), bias: T.Buffer(64, "float32"), bn_offset: T.Buffer(64, "float32"), bn_scale: T.Buffer(64, "float32"), compute: T.Buffer((1, 112, 112, 64), "float32")) -> None: + def cbr_0(data: T.Tensor((1, 224, 224, 3), "float32"), kernel: T.Tensor((7, 7, 3, 64), "float32"), bias: T.Tensor(64, "float32"), bn_offset: T.Tensor(64, "float32"), bn_scale: T.Tensor(64, "float32"), compute: T.Tensor((1, 112, 112, 64), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1148,7 +1148,7 @@ def cbr_0(data: T.Buffer((1, 224, 224, 3), "float32"), kernel: T.Buffer((7, 7, 3 def test_cuda_tbg(): # fmt: off @Ts.prim_func - def tbg_0(query: T.Buffer((1, 128, 12, 64), "float32"), value: T.Buffer((1, 128, 12, 64), "float32"), C: T.Buffer((1, 12, 128, 128), "float32")) -> None: + def tbg_0(query: T.Tensor((1, 128, 12, 64), "float32"), value: T.Tensor((1, 128, 12, 64), "float32"), C: T.Tensor((1, 12, 128, 128), "float32")) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cuda_async.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cuda_async.py index d0961b7a1fc8..79f8b3a76d8d 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cuda_async.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_space_cuda_async.py @@ -46,7 +46,7 @@ def get_c2d_prim_func(stage: int): if stage == 0: # fmt: off @Ts.prim_func - def c2d(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32")): + def c2d(inputs: T.Tensor((1, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -107,7 +107,7 @@ def c2d(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 3 else: # fmt: off @Ts.prim_func - def c2d(inputs: T.Buffer((1, 224, 224, 3), "float32"), weight: T.Buffer((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32")): + def c2d(inputs: T.Tensor((1, 224, 224, 3), "float32"), weight: T.Tensor((7, 7, 3, 64), "float32"), conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -200,7 +200,7 @@ def get_gmm_prim_func(stage: int): if stage == 0: # fmt: off @Ts.prim_func - def gmm(X: T.Buffer((1, 1024, 1024), "float32"), Y: T.Buffer((1, 1024, 1024), "float32"), Z: T.Buffer((1, 1024, 1024), "float32")): + def gmm(X: T.Tensor((1, 1024, 1024), "float32"), Y: T.Tensor((1, 1024, 1024), "float32"), Z: T.Tensor((1, 1024, 1024), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -255,7 +255,7 @@ def gmm(X: T.Buffer((1, 1024, 1024), "float32"), Y: T.Buffer((1, 1024, 1024), "f else: # fmt: off @Ts.prim_func - def gmm(X: T.Buffer((1, 1024, 1024), "float32"), Y: T.Buffer((1, 1024, 1024), "float32"), Z: T.Buffer((1, 1024, 1024), "float32")): + def gmm(X: T.Tensor((1, 1024, 1024), "float32"), Y: T.Tensor((1, 1024, 1024), "float32"), Z: T.Tensor((1, 1024, 1024), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_space_generator.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_space_generator.py index f08caf4a750e..411108caf83c 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_space_generator.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_space_generator.py @@ -41,9 +41,9 @@ class Matmul: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: T.func_attr({"global_symbol": "main"}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_space_post_opt.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_space_post_opt.py index e72b81d1f790..ae853583b5ec 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_space_post_opt.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_space_post_opt.py @@ -36,7 +36,7 @@ @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_task_scheduler.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_task_scheduler.py index 9598bfa2646b..e4b37971e2a9 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_task_scheduler.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_task_scheduler.py @@ -38,9 +38,9 @@ class MatmulModule: @Ts.prim_func def main( # type: ignore - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -56,9 +56,9 @@ def main( # type: ignore class MatmulReluModule: @Ts.prim_func def main( # type: ignore - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - D: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + D: T.Tensor((1024, 1024), "float32"), ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -79,9 +79,9 @@ def main( # type: ignore class BatchMatmulModule: @Ts.prim_func def main( # type: ignore - A: T.Buffer([16, 128, 128]), - B: T.Buffer([16, 128, 128]), - C: T.Buffer([16, 128, 128]), + A: T.Tensor([16, 128, 128]), + B: T.Tensor([16, 128, 128]), + C: T.Tensor([16, 128, 128]), ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) 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 bcc2775ea9e1..05483bcbc165 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 @@ -35,9 +35,9 @@ class Dense: @Ts.prim_func def main( - p0: T.Buffer((128, 128), "float32"), - p1: T.Buffer((128, 128), "float32"), - T_matmul_NT: T.Buffer((128, 128), "float32"), + p0: T.Tensor((128, 128), "float32"), + p1: T.Tensor((128, 128), "float32"), + T_matmul_NT: T.Tensor((128, 128), "float32"), ) -> None: # function attr dict T.func_attr({"layout_free_buffers": [1], "tirx.noalias": True, "global_symbol": "main"}) @@ -57,9 +57,9 @@ def main( class DenseAdd: @Ts.prim_func def main( - p0: T.Buffer((128, 128), "float32"), - p1: T.Buffer((128, 128), "float32"), - T_add: T.Buffer((128, 128), "float32"), + p0: T.Tensor((128, 128), "float32"), + p1: T.Tensor((128, 128), "float32"), + T_add: T.Tensor((128, 128), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True, "layout_free_buffers": [1]}) @@ -92,9 +92,9 @@ def main( class DenseAdd_scheduled_cpu: @Ts.prim_func def main( - p0: T.Buffer((128, 128), "float32"), - p1: T.Buffer((128, 128), "float32"), - T_add: T.Buffer((128, 128), "float32"), + p0: T.Tensor((128, 128), "float32"), + p1: T.Tensor((128, 128), "float32"), + T_add: T.Tensor((128, 128), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True, "layout_free_buffers": [1]}) @@ -173,7 +173,7 @@ def main( @tvm.script.ir_module class DenseAdd_cpu_no_write_cache: @Ts.prim_func - def main(p0: T.Buffer((128, 128), "float32"), p1: T.Buffer((128, 128), "float32"), T_add: T.Buffer((128, 128), "float32")) -> None: + def main(p0: T.Tensor((128, 128), "float32"), p1: T.Tensor((128, 128), "float32"), T_add: T.Tensor((128, 128), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True, "layout_free_buffers": [1]}) # body @@ -219,9 +219,9 @@ def main(p0: T.Buffer((128, 128), "float32"), p1: T.Buffer((128, 128), "float32" class DenseAdd_scheduled_gpu: @Ts.prim_func def main( - p0: T.Buffer((128, 128), "float32"), - p1: T.Buffer((128, 128), "float32"), - T_add: T.Buffer((128, 128), "float32"), + p0: T.Tensor((128, 128), "float32"), + p1: T.Tensor((128, 128), "float32"), + T_add: T.Tensor((128, 128), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True, "layout_free_buffers": [1]}) @@ -371,7 +371,7 @@ def main( @tvm.script.ir_module class Conv2dInt8: @Ts.prim_func - def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), "int8"), p2: T.Buffer((1, 1, 1, 256), "int32"), p3: T.Buffer((1, 1, 1, 256), "int32"), p4: T.Buffer((1, 1, 1, 256), "int64"), p5: T.Buffer((1, 1, 1, 256), "int64"), p6: T.Buffer((1, 1, 1, 256), "int64"), p7: T.Buffer((), "int32"), p8: T.Buffer(1, "int32"), compute: T.Buffer((16, 56, 56, 256), "int32")) -> None: + def main(p0: T.Tensor((16, 56, 56, 64), "int8"), p1: T.Tensor((256, 1, 1, 64), "int8"), p2: T.Tensor((1, 1, 1, 256), "int32"), p3: T.Tensor((1, 1, 1, 256), "int32"), p4: T.Tensor((1, 1, 1, 256), "int64"), p5: T.Tensor((1, 1, 1, 256), "int64"), p6: T.Tensor((1, 1, 1, 256), "int64"), p7: T.Tensor((), "int32"), p8: T.Tensor(1, "int32"), compute: T.Tensor((16, 56, 56, 256), "int32")) -> None: # function attr dict T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) # body @@ -486,7 +486,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " @tvm.script.ir_module class Conv2dInt8_target: @Ts.prim_func - def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), "int8"), p2: T.Buffer((1, 1, 1, 256), "int32"), p3: T.Buffer((1, 1, 1, 256), "int32"), p4: T.Buffer((1, 1, 1, 256), "int64"), p5: T.Buffer((1, 1, 1, 256), "int64"), p6: T.Buffer((1, 1, 1, 256), "int64"), p7: T.Buffer((), "int32"), p8: T.Buffer(1, "int32"), p9: T.Buffer((16, 56, 56, 256), "int32"), compute: T.Buffer((16, 56, 56, 256), "uint8")) -> None: + def main(p0: T.Tensor((16, 56, 56, 64), "int8"), p1: T.Tensor((256, 1, 1, 64), "int8"), p2: T.Tensor((1, 1, 1, 256), "int32"), p3: T.Tensor((1, 1, 1, 256), "int32"), p4: T.Tensor((1, 1, 1, 256), "int64"), p5: T.Tensor((1, 1, 1, 256), "int64"), p6: T.Tensor((1, 1, 1, 256), "int64"), p7: T.Tensor((), "int32"), p8: T.Tensor(1, "int32"), p9: T.Tensor((16, 56, 56, 256), "int32"), compute: T.Tensor((16, 56, 56, 256), "uint8")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -650,7 +650,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " @tvm.script.ir_module class Conv2dInt8_tensorcore_scheduled: @Ts.prim_func - def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), "int8"), p2: T.Buffer((1, 1, 1, 256), "int32"), p3: T.Buffer((1, 1, 1, 256), "int32"), p4: T.Buffer((1, 1, 1, 256), "int64"), p5: T.Buffer((1, 1, 1, 256), "int64"), p6: T.Buffer((1, 1, 1, 256), "int64"), p7: T.Buffer((), "int32"), p8: T.Buffer((1,), "int32"), p9: T.Buffer((16, 56, 56, 256), "int32"), compute: T.Buffer((16, 56, 56, 256), "uint8")): + def main(p0: T.Tensor((16, 56, 56, 64), "int8"), p1: T.Tensor((256, 1, 1, 64), "int8"), p2: T.Tensor((1, 1, 1, 256), "int32"), p3: T.Tensor((1, 1, 1, 256), "int32"), p4: T.Tensor((1, 1, 1, 256), "int64"), p5: T.Tensor((1, 1, 1, 256), "int64"), p6: T.Tensor((1, 1, 1, 256), "int64"), p7: T.Tensor((), "int32"), p8: T.Tensor((1,), "int32"), p9: T.Tensor((16, 56, 56, 256), "int32"), compute: T.Tensor((16, 56, 56, 256), "uint8")): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): conv2d_nhwc_reindex_shared = Ts.sblock_alloc_buffer((50176, 256), "int32", scope="shared") @@ -751,7 +751,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " @tvm.script.ir_module class Conv2dInt8_NCHWc: @Ts.prim_func - def main(p0: T.Buffer((1, 32, 7, 7, 16), "uint8"), p1: T.Buffer((128, 32, 1, 1, 4, 16, 4), "int8"), p2: T.Buffer((1, 128, 1, 1, 16), "int32"), p3: T.Buffer((1, 128, 1, 1, 16), "float32"), p4: T.Buffer(1, "float32"), p5: T.Buffer((1, 128, 7, 7, 16), "int32"), compute: T.Buffer((1, 128, 7, 7, 16), "uint8")) -> None: + def main(p0: T.Tensor((1, 32, 7, 7, 16), "uint8"), p1: T.Tensor((128, 32, 1, 1, 4, 16, 4), "int8"), p2: T.Tensor((1, 128, 1, 1, 16), "int32"), p3: T.Tensor((1, 128, 1, 1, 16), "float32"), p4: T.Tensor(1, "float32"), p5: T.Tensor((1, 128, 7, 7, 16), "int32"), compute: T.Tensor((1, 128, 7, 7, 16), "uint8")) -> None: # function attr dict T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) # body @@ -913,7 +913,7 @@ def main(p0: T.Buffer((1, 32, 7, 7, 16), "uint8"), p1: T.Buffer((128, 32, 1, 1, @tvm.script.ir_module class Conv2dInt8_NCHWc_target: @Ts.prim_func - def main(p0: T.Buffer((1, 32, 7, 7, 16), "uint8"), p1: T.Buffer((128, 32, 1, 1, 4, 16, 4), "int8"), p2: T.Buffer((1, 128, 1, 1, 16), "int32"), p3: T.Buffer((1, 128, 1, 1, 16), "float32"), p4: T.Buffer(1, "float32"), p5: T.Buffer((1, 128, 7, 7, 16), "uint8"), T_cast: T.Buffer((1, 128, 7, 7, 16), "int32")) -> None: + def main(p0: T.Tensor((1, 32, 7, 7, 16), "uint8"), p1: T.Tensor((128, 32, 1, 1, 4, 16, 4), "int8"), p2: T.Tensor((1, 128, 1, 1, 16), "int32"), p3: T.Tensor((1, 128, 1, 1, 16), "float32"), p4: T.Tensor(1, "float32"), p5: T.Tensor((1, 128, 7, 7, 16), "uint8"), T_cast: T.Tensor((1, 128, 7, 7, 16), "int32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -1130,7 +1130,7 @@ def get_conv2d_vnni_mod(intrin_id): @tvm.script.ir_module class Conv2dInt8_NCHWc_scheduled: @Ts.prim_func - def main(p0: T.Buffer((1, 32, 7, 7, 16), "uint8"), p1: T.Buffer((128, 32, 1, 1, 4, 16, 4), "int8"), p2: T.Buffer((1, 128, 1, 1, 16), "int32"), p3: T.Buffer((1, 128, 1, 1, 16), "float32"), p4: T.Buffer(1, "float32"), p5: T.Buffer((1, 128, 7, 7, 16), "uint8"), T_cast: T.Buffer((1, 128, 7, 7, 16), "int32")) -> None: + def main(p0: T.Tensor((1, 32, 7, 7, 16), "uint8"), p1: T.Tensor((128, 32, 1, 1, 4, 16, 4), "int8"), p2: T.Tensor((1, 128, 1, 1, 16), "int32"), p3: T.Tensor((1, 128, 1, 1, 16), "float32"), p4: T.Tensor(1, "float32"), p5: T.Tensor((1, 128, 7, 7, 16), "uint8"), T_cast: T.Tensor((1, 128, 7, 7, 16), "int32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -1192,7 +1192,7 @@ def main(p0: T.Buffer((1, 32, 7, 7, 16), "uint8"), p1: T.Buffer((128, 32, 1, 1, @tvm.script.ir_module class Conv2dWinogradAddRelu: @Ts.prim_func - def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((6, 6, 64, 64), "float32"), p2: T.Buffer((1, 1, 1, 64), "float32"), T_relu: T.Buffer((1, 56, 56, 64), "float32")) -> None: + def main(p0: T.Tensor((1, 56, 56, 64), "float32"), p1: T.Tensor((6, 6, 64, 64), "float32"), p2: T.Tensor((1, 1, 1, 64), "float32"), T_relu: T.Tensor((1, 56, 56, 64), "float32")) -> None: # function attr dict T.func_attr({"layout_free_buffers": [1], "tirx.noalias": True, "global_symbol": "main"}) # body @@ -1283,7 +1283,7 @@ def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((6, 6, 64, 64), @tvm.script.ir_module class Conv2dWinogradAddResidualRelu: @Ts.prim_func - def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((6, 6, 64, 64), "float32"), p2: T.Buffer((1, 1, 1, 64), "float32"), p3: T.Buffer((1, 56, 56, 64), "float32"), T_relu: T.Buffer((1, 56, 56, 64), "float32")) -> None: + def main(p0: T.Tensor((1, 56, 56, 64), "float32"), p1: T.Tensor((6, 6, 64, 64), "float32"), p2: T.Tensor((1, 1, 1, 64), "float32"), p3: T.Tensor((1, 56, 56, 64), "float32"), T_relu: T.Tensor((1, 56, 56, 64), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True, "layout_free_buffers": [1]}) # body @@ -1381,7 +1381,7 @@ def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((6, 6, 64, 64), @tvm.script.ir_module class Conv2dWinogradAddResidualRelu_scheduled: @Ts.prim_func - def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((6, 6, 64, 64), "float32"), p2: T.Buffer((1, 1, 1, 64), "float32"), p3: T.Buffer((1, 56, 56, 64), "float32"), T_relu: T.Buffer((1, 56, 56, 64), "float32")) -> None: + def main(p0: T.Tensor((1, 56, 56, 64), "float32"), p1: T.Tensor((6, 6, 64, 64), "float32"), p2: T.Tensor((1, 1, 1, 64), "float32"), p3: T.Tensor((1, 56, 56, 64), "float32"), T_relu: T.Tensor((1, 56, 56, 64), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True, "layout_free_buffers": [1]}) # body @@ -1520,7 +1520,7 @@ def main(p0: T.Buffer((1, 56, 56, 64), "float32"), p1: T.Buffer((6, 6, 64, 64), @tvm.script.ir_module class Conv2dInt8_with_predicate: @Ts.prim_func - def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), "int8"), p2: T.Buffer((1, 1, 1, 256), "int32"), p3: T.Buffer((1, 1, 1, 256), "int32"), p4: T.Buffer(256, "int32"), p5: T.Buffer(256, "int32"), p6: T.Buffer(256, "int32"), p7: T.Buffer((), "int32"), p8: T.Buffer(1, "int32"), compute: T.Buffer((16, 56, 56, 256), "int32")) -> None: + def main(p0: T.Tensor((16, 56, 56, 64), "int8"), p1: T.Tensor((256, 1, 1, 64), "int8"), p2: T.Tensor((1, 1, 1, 256), "int32"), p3: T.Tensor((1, 1, 1, 256), "int32"), p4: T.Tensor(256, "int32"), p5: T.Tensor(256, "int32"), p6: T.Tensor(256, "int32"), p7: T.Tensor((), "int32"), p8: T.Tensor(1, "int32"), compute: T.Tensor((16, 56, 56, 256), "int32")) -> None: # function attr dict T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) # body @@ -1593,7 +1593,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " @tvm.script.ir_module class Conv2dInt8_with_predicate_target: @Ts.prim_func - def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), "int8"), p2: T.Buffer((1, 1, 1, 256), "int32"), p3: T.Buffer((1, 1, 1, 256), "int32"), p4: T.Buffer(256, "int32"), p5: T.Buffer(256, "int32"), p6: T.Buffer(256, "int32"), p7: T.Buffer((), "int32"), p8: T.Buffer(1, "int32"), p9: T.Buffer((16, 56, 56, 256), "int32"), compute: T.Buffer((16, 56, 56, 256), "int32")) -> None: + def main(p0: T.Tensor((16, 56, 56, 64), "int8"), p1: T.Tensor((256, 1, 1, 64), "int8"), p2: T.Tensor((1, 1, 1, 256), "int32"), p3: T.Tensor((1, 1, 1, 256), "int32"), p4: T.Tensor(256, "int32"), p5: T.Tensor(256, "int32"), p6: T.Tensor(256, "int32"), p7: T.Tensor((), "int32"), p8: T.Tensor(1, "int32"), p9: T.Tensor((16, 56, 56, 256), "int32"), compute: T.Tensor((16, 56, 56, 256), "int32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -1687,7 +1687,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " @tvm.script.ir_module class Conv2dInt8_with_predicate_scheduled: @Ts.prim_func - def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), "int8"), p2: T.Buffer((1, 1, 1, 256), "int32"), p3: T.Buffer((1, 1, 1, 256), "int32"), p4: T.Buffer((256,), "int32"), p5: T.Buffer((256,), "int32"), p6: T.Buffer((256,), "int32"), p7: T.Buffer((), "int32"), p8: T.Buffer((1,), "int32"), p9: T.Buffer((16, 56, 56, 256), "int32"), compute: T.Buffer((16, 56, 56, 256), "int32")): + def main(p0: T.Tensor((16, 56, 56, 64), "int8"), p1: T.Tensor((256, 1, 1, 64), "int8"), p2: T.Tensor((1, 1, 1, 256), "int32"), p3: T.Tensor((1, 1, 1, 256), "int32"), p4: T.Tensor((256,), "int32"), p5: T.Tensor((256,), "int32"), p6: T.Tensor((256,), "int32"), p7: T.Tensor((), "int32"), p8: T.Tensor((1,), "int32"), p9: T.Tensor((16, 56, 56, 256), "int32"), compute: T.Tensor((16, 56, 56, 256), "int32")): T.func_attr({"tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_tune_context.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_tune_context.py index 39aef460435e..0392f9f16c6b 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_tune_context.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_tune_context.py @@ -36,9 +36,9 @@ class Matmul: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float32"), - B: T.Buffer((1024, 1024), "float32"), - C: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float32"), + B: T.Tensor((1024, 1024), "float32"), + C: T.Tensor((1024, 1024), "float32"), ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_tune_tir.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_tune_tir.py index c48c66bd12e6..77297cea4006 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_tune_tir.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_tune_tir.py @@ -40,7 +40,7 @@ @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) @@ -50,7 +50,7 @@ def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 1 @Ts.prim_func -def two_step(A: T.Buffer((1024, 1024), "float32"), C: T.Buffer((1024, 1024), "float32")) -> None: +def two_step(A: T.Tensor((1024, 1024), "float32"), C: T.Tensor((1024, 1024), "float32")) -> None: B = Ts.sblock_alloc_buffer((1024, 1024), "float32") for i, j in T.grid(1024, 1024): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_analysis.py b/tests/python/s_tir/schedule/test_tir_schedule_analysis.py index 551bbd5c7d9a..c898f2ebcf75 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_analysis.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_analysis.py @@ -46,7 +46,7 @@ ForKind, IndexMap, Var, - decl_buffer, + decl_tensor, floordiv, floormod, ) @@ -73,7 +73,7 @@ def _make_loops(loop_vars: list[Var], extents: list[int]) -> list[For]: def test_suggest_index_map_simple(): i, j = _make_vars("i", "j") index_map = suggest_index_map( - buffer=decl_buffer(shape=[8, 256]), + buffer=decl_tensor(shape=[8, 256]), indices=[ floordiv(i, 16) * 4 + floordiv(j, 16), floormod(i, 16) * 16 + floormod(j, 16), @@ -98,7 +98,7 @@ def test_suggest_index_map_simple(): def test_suggest_index_map_bijective(): i, j = _make_vars("i", "j") index_map = suggest_index_map( - buffer=decl_buffer(shape=[8]), + buffer=decl_tensor(shape=[8]), indices=[floormod(j, 4) * 2 + i], loops=_make_loops( loop_vars=[i, j], @@ -122,7 +122,7 @@ def test_suggest_index_map_winograd(): nu = floordiv(floormod(fused_outer, 336), 112) * 2 + floordiv(floormod(fused_outer, 8), 4) co = floormod(fused_outer, 4) * 32 + i3_3_fused ci = (i4_0 * 32) + i4_1 - buffer = decl_buffer(shape=[6, 6, 128, 128]) + buffer = decl_tensor(shape=[6, 6, 128, 128]) index_map = suggest_index_map( buffer=buffer, indices=[eps, nu, co, ci], @@ -160,9 +160,9 @@ def test_suggest_index_map_winograd(): class DenseTIRModule: @Ts.prim_func def main( - placeholder: T.Buffer((1024, 1024), "uint8"), - placeholder_1: T.Buffer((64, 256, 16, 4), "int8"), - compute: T.Buffer((1024, 1024), "int32"), + placeholder: T.Tensor((1024, 1024), "uint8"), + placeholder_1: T.Tensor((64, 256, 16, 4), "int8"), + compute: T.Tensor((1024, 1024), "int32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): @@ -184,9 +184,9 @@ def main( class Conv2dNCHWcTIRModule: @Ts.prim_func def main( - placeholder: T.Buffer((1, 4, 56, 56, 16), "uint8"), - placeholder_1: T.Buffer((16, 4, 1, 1, 4, 16, 4), "int8"), - conv2d_NCHWc_int8: T.Buffer((1, 16, 56, 56, 16), "int32"), + placeholder: T.Tensor((1, 4, 56, 56, 16), "uint8"), + placeholder_1: T.Tensor((16, 4, 1, 1, 4, 16, 4), "int8"), + conv2d_NCHWc_int8: T.Tensor((1, 16, 56, 56, 16), "int32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0, i1, i2, i3, i4, i5, i6, i7, i8, i9 in T.grid(1, 16, 56, 56, 16, 1, 1, 4, 4, 4): @@ -268,9 +268,9 @@ def test_get_tensorize_loop_mapping_conv2d_nchwc_16x4(): def test_get_tensorize_loop_mapping_matmul_mma(): @Ts.prim_func def matmul_16x16x16xf16f16f16_desc( - A: T.Buffer((16, 16), "float16", align=64, offset_factor=1), - B: T.Buffer((16, 16), "float16", align=64, offset_factor=1), - C: T.Buffer((16, 16), "float16", align=64, offset_factor=1), + A: T.Tensor((16, 16), "float16", align=64, offset_factor=1), + B: T.Tensor((16, 16), "float16", align=64, offset_factor=1), + C: T.Tensor((16, 16), "float16", align=64, offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:16, 0:16], A[0:16, 0:16], B[0:16, 0:16]) @@ -404,7 +404,7 @@ def test_get_auto_tensorize_mapping_info_matmul(n, m, k, expected): def test_is_output_block(): @Ts.prim_func def two_elementwise( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -424,7 +424,7 @@ def two_elementwise( def test_empty_grid(): @Ts.prim_func - def foo(out: T.Buffer((T.int64(1), T.int64(8), T.int64(8)), "int32")): + def foo(out: T.Tensor((T.int64(1), T.int64(8), T.int64(8)), "int32")): act = Ts.sblock_alloc_buffer((1, 8, 8), "int32") for z2, y2, x2 in T.grid(1, 8, 8): with Ts.sblock("b0"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_annotate_buffer_access.py b/tests/python/s_tir/schedule/test_tir_schedule_annotate_buffer_access.py index 2998da2795c0..790c9b893bff 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_annotate_buffer_access.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_annotate_buffer_access.py @@ -29,7 +29,7 @@ def test_annotate_read_buffer_access(): @Ts.prim_func - def before(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")): + def before(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")): B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -41,7 +41,7 @@ def before(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32" C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func - def expected(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")): + def expected(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")): B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -66,7 +66,7 @@ def expected(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float3 def test_annotate_write_buffer_access(): @Ts.prim_func - def before(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")): + def before(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")): B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -78,7 +78,7 @@ def before(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32" C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func - def expected(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")): + def expected(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")): B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -102,7 +102,7 @@ def expected(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float3 def test_annotate_buffer_access_for_resize(): # fmt: off @Ts.prim_func - def resize_before(x: T.Buffer((1, 1, 32, 32), "float16"), resize: T.Buffer((1, 1, 16, 16), "float16")): + def resize_before(x: T.Tensor((1, 1, 32, 32), "float16"), resize: T.Tensor((1, 1, 16, 16), "float16")): for i0, i1, i2, i3 in T.grid(1, 1, 16, 16): with Ts.sblock("resize"): v_i0, v_i1, v_i2, v_i3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) @@ -111,7 +111,7 @@ def resize_before(x: T.Buffer((1, 1, 32, 32), "float16"), resize: T.Buffer((1, 1 resize[v_i0, v_i1, v_i2, v_i3] = T.Cast("float16", T.Cast("float32", x[v_i0, v_i1, T.max(T.min(T.Cast("int32", T.floor((T.Cast("float32", v_i2) + T.float32(0.5)) * T.float32(2) - T.float32(0.5) + T.float32(1.0000000000000001e-05))), 31), 0), T.max(T.min(T.Cast("int32", T.floor((T.Cast("float32", v_i3) + T.float32(0.5)) * T.float32(2) - T.float32(0.5) + T.float32(1.0000000000000001e-05))), 31), 0)])) @Ts.prim_func - def resize_expected(x: T.Buffer((1, 1, 32, 32), "float16"), resize: T.Buffer((1, 1, 16, 16), "float16")): + def resize_expected(x: T.Tensor((1, 1, 32, 32), "float16"), resize: T.Tensor((1, 1, 16, 16), "float16")): for i0, i1, i2, i3 in T.grid(1, 1, 16, 16): with Ts.sblock("resize"): v_i0, v_i1, v_i2, v_i3 = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) @@ -139,7 +139,7 @@ def resize_expected(x: T.Buffer((1, 1, 32, 32), "float16"), resize: T.Buffer((1, def test_annotate_buffer_access_read_and_write(): @Ts.prim_func - def before(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")): + def before(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")): B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -155,7 +155,7 @@ def before(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32" C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func - def expected(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")): + def expected(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")): B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -186,7 +186,7 @@ def expected(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float3 def test_double_annotate_buffer_access_read(): @Ts.prim_func - def before(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")): + def before(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")): B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -202,7 +202,7 @@ def before(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32" C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func - def expected(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")): + def expected(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")): B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -236,7 +236,7 @@ def expected(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float3 def test_annotate_buffer_access_with_compute_at_for_resize(): # fmt: off @Ts.prim_func - def before(x: T.Buffer((1, 3, 200, 200), "float32"), y: T.Buffer((1, 3, 100, 100), "float32")): + def before(x: T.Tensor((1, 3, 200, 200), "float32"), y: T.Tensor((1, 3, 100, 100), "float32")): x_global = Ts.sblock_alloc_buffer([1, 3, 200, 200], dtype="float32") for ax0, ax1, ax2, ax3 in T.grid(1, 3, 200, 200): with Ts.sblock("cache"): @@ -248,7 +248,7 @@ def before(x: T.Buffer((1, 3, 200, 200), "float32"), y: T.Buffer((1, 3, 100, 100 y[v_i0, v_i1, v_i2, v_i3] = x_global[v_i0, v_i1, T.Cast("int32", T.floor(v_i2 * 2 + 0.5)), T.Cast("int32", T.floor(v_i3 * 2 + 0.5))] @Ts.prim_func - def after(x: T.Buffer((1, 3, 200, 200), "float32"), y: T.Buffer((1, 3, 100, 100), "float32")): + def after(x: T.Tensor((1, 3, 200, 200), "float32"), y: T.Tensor((1, 3, 100, 100), "float32")): x_global = Ts.sblock_alloc_buffer((1, 3, 200, 200)) for i0, i1, i2_0, i3_0 in T.grid(1, 3, 10, 10): for ax0, ax1 in T.grid(24, 24): @@ -272,7 +272,7 @@ def after(x: T.Buffer((1, 3, 200, 200), "float32"), y: T.Buffer((1, 3, 100, 100) y[v_i0, v_i1, v_i2, v_i3] = x_global[v_i0, v_i1, T.Cast("int32", T.floor(T.Cast("float32", v_i2 * 2) + T.float32(0.5))), T.Cast("int32", T.floor(T.Cast("float32", v_i3 * 2) + T.float32(0.5)))] @Ts.prim_func - def after_without_annotate_buffer_access(x: T.Buffer((1, 3, 200, 200), "float32"), y: T.Buffer((1, 3, 100, 100), "float32")): + def after_without_annotate_buffer_access(x: T.Tensor((1, 3, 200, 200), "float32"), y: T.Tensor((1, 3, 100, 100), "float32")): x_global = Ts.sblock_alloc_buffer((1, 3, 200, 200)) for i0, i1, i2_0, i3_0 in T.grid(1, 3, 10, 10): for ax0, ax1 in T.grid(200, 200): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_block_scope.py b/tests/python/s_tir/schedule/test_tir_schedule_block_scope.py index ebbf7331cac9..3b32fa2e25dc 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_block_scope.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_block_scope.py @@ -32,7 +32,7 @@ @Ts.prim_func -def elementwise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")) -> None: +def elementwise(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -45,7 +45,7 @@ def elementwise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "flo @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j in T.grid(128, 128): with Ts.sblock("init"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -58,7 +58,7 @@ def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 1 @Ts.prim_func def war_dependency( - A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), B: T.Tensor((128, 128)), C: T.Tensor((128, 128)) ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("C"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_blockize.py b/tests/python/s_tir/schedule/test_tir_schedule_blockize.py index e8c137b57a7f..1637523260fb 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_blockize.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_blockize.py @@ -29,7 +29,7 @@ # pylint: disable=no-member,invalid-name,unused-variable,line-too-long,redefined-outer-name,unexpected-keyword-arg,too-many-nested-blocks @Ts.prim_func -def single_elementwise(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32")): +def single_elementwise(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32")): for i, j in T.grid(128, 128): with Ts.sblock("B"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -42,8 +42,8 @@ def single_elementwise(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128 def test_blockize_outer(): @Ts.prim_func def after_blockize_outer( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), ) -> None: with Ts.sblock("blockized_B"): vio = Ts.axis.spatial(1, 0) @@ -66,8 +66,8 @@ def after_blockize_outer( def test_blockize_inner(): @Ts.prim_func def after_blockize_inner( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), ) -> None: for i in T.serial(128): with Ts.sblock("blockized_B"): @@ -91,8 +91,8 @@ def after_blockize_inner( def test_two_elementwise_blockize_reverse_compute_at(): @Ts.prim_func def before_blockize_rca( - A: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: B = Ts.sblock_alloc_buffer([128, 128], dtype="float32") for i, j in T.grid(8, 8): @@ -116,8 +116,8 @@ def before_blockize_rca( @Ts.prim_func def after_blockize_rca( - A: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: B = Ts.sblock_alloc_buffer([128, 128], dtype="float32") for i, j in T.grid(8, 8): @@ -155,8 +155,8 @@ def after_blockize_rca( def test_two_elementwise_blockize_compute_at(): @Ts.prim_func def before_blockize_compute_at( - A: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: # body # with Ts.sblock("root") @@ -184,8 +184,8 @@ def before_blockize_compute_at( @Ts.prim_func def after_blockize_compute_at( - A: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: B = Ts.sblock_alloc_buffer([128, 128], dtype="float32") for i_0, j_0 in T.grid(8, 8): @@ -227,7 +227,7 @@ def after_blockize_compute_at( def test_blockize_init_loops(): @Ts.prim_func - def rowsum(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128,), "float32")) -> None: + def rowsum(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128,), "float32")) -> None: for k, i in T.grid(128, 128): with Ts.sblock("B"): vk, vi = Ts.axis.remap("RS", [k, i]) @@ -237,8 +237,8 @@ def rowsum(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128,), "float32")) - @Ts.prim_func def after_rowsum_blockize( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128,), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128,), "float32"), ) -> None: with Ts.sblock("blockized_B"): vko = Ts.axis.R(1, 0) @@ -266,8 +266,8 @@ def after_rowsum_blockize( def test_blockize_outer_int64_shape(preserve_unit_iters): @Ts.prim_func def single_elementwise_int64( - A: T.Buffer((T.int64(16), T.int64(128)), "float32"), - B: T.Buffer((T.int64(16), T.int64(128)), "float32"), + A: T.Tensor((T.int64(16), T.int64(128)), "float32"), + B: T.Tensor((T.int64(16), T.int64(128)), "float32"), ) -> None: for i0, j0, i1, j1 in T.grid(T.int64(1), T.int64(8), T.int64(16), T.int64(16)): with Ts.sblock("B"): @@ -277,8 +277,8 @@ def single_elementwise_int64( @Ts.prim_func def after_single_elementwise_int64_blockize( - A: T.Buffer((T.int64(16), T.int64(128)), "float32"), - B: T.Buffer((T.int64(16), T.int64(128)), "float32"), + A: T.Tensor((T.int64(16), T.int64(128)), "float32"), + B: T.Tensor((T.int64(16), T.int64(128)), "float32"), ) -> None: for i0, j0 in T.grid(T.int64(1), T.int64(8)): with Ts.sblock("B_o"): @@ -293,8 +293,8 @@ def after_single_elementwise_int64_blockize( @Ts.prim_func def after_single_elementwise_int64_blockize_preserve_unit_iters( - A: T.Buffer((T.int64(16), T.int64(128)), "float32"), - B: T.Buffer((T.int64(16), T.int64(128)), "float32"), + A: T.Tensor((T.int64(16), T.int64(128)), "float32"), + B: T.Tensor((T.int64(16), T.int64(128)), "float32"), ) -> None: for i0, j0 in T.grid(T.int64(1), T.int64(8)): with Ts.sblock("B_o"): @@ -323,7 +323,7 @@ def after_single_elementwise_int64_blockize_preserve_unit_iters( def test_blockize_blocks(): @Ts.prim_func - def blocks_func(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32")) -> None: + def blocks_func(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32")) -> None: for m in T.serial(6): for i, j in T.grid(3, 1): with Ts.sblock("B"): @@ -341,7 +341,7 @@ def blocks_func(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "flo @Ts.prim_func def after_blocks_blockize( - A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32") ) -> None: for m in range(6): with Ts.sblock("outer_B_C_"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_cache_index.py b/tests/python/s_tir/schedule/test_tir_schedule_cache_index.py index 271501bc42e3..abe674b5ad37 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_cache_index.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_cache_index.py @@ -33,7 +33,7 @@ @Ts.prim_func -def resize(A: T.Buffer((1, 3, 40, 40)), B: T.Buffer((1, 3, 80, 80))) -> None: +def resize(A: T.Tensor((1, 3, 40, 40)), B: T.Tensor((1, 3, 80, 80))) -> None: for i0, i1, i2, i3 in T.grid(1, 3, 80, 80): with Ts.sblock("A"): n, c, vi, vj = Ts.axis.remap("SSSS", [i0, i1, i2, i3]) @@ -42,7 +42,7 @@ def resize(A: T.Buffer((1, 3, 40, 40)), B: T.Buffer((1, 3, 80, 80))) -> None: @Ts.prim_func def resize_cache_index( - A: T.Buffer((1, 3, 40, 40), "float32"), B: T.Buffer((1, 3, 80, 80), "float32") + A: T.Tensor((1, 3, 40, 40), "float32"), B: T.Tensor((1, 3, 80, 80), "float32") ) -> None: index_var_0 = Ts.sblock_alloc_buffer([80, 80], dtype="int32", strides=[1]) index_var_1 = Ts.sblock_alloc_buffer([80], dtype="int32", strides=[1]) @@ -68,7 +68,7 @@ def resize_cache_index( @Ts.prim_func def bilinear_resize( - x: T.Buffer((1, 3, 40, 40), "float16"), resize: T.Buffer((1, 3, 80, 80), "float16") + x: T.Tensor((1, 3, 40, 40), "float16"), resize: T.Tensor((1, 3, 80, 80), "float16") ): for i0, i1, i2, i3 in T.grid(1, 3, 80, 80): with Ts.sblock("resize"): @@ -323,7 +323,7 @@ def bilinear_resize( @Ts.prim_func def cached_bilinear_resize( - x: T.Buffer((1, 3, 40, 40), "float16"), resize: T.Buffer((1, 3, 80, 80), "float16") + x: T.Tensor((1, 3, 40, 40), "float16"), resize: T.Tensor((1, 3, 80, 80), "float16") ): index_var_0 = Ts.sblock_alloc_buffer([80], dtype="float32", strides=[1]) index_var_1 = Ts.sblock_alloc_buffer([80], dtype="int32", strides=[1]) diff --git a/tests/python/s_tir/schedule/test_tir_schedule_cache_read_write.py b/tests/python/s_tir/schedule/test_tir_schedule_cache_read_write.py index 33ae465b9ffc..d33d0bde8e12 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_cache_read_write.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_cache_read_write.py @@ -38,7 +38,7 @@ @Ts.prim_func -def elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -53,7 +53,7 @@ def elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: @Ts.prim_func def elementwise_shape_int64( - A: T.Buffer((T.int64(128), T.int64(128))), C: T.Buffer((T.int64(128), T.int64(128))) + A: T.Tensor((T.int64(128), T.int64(128))), C: T.Tensor((T.int64(128), T.int64(128))) ) -> None: B = Ts.sblock_alloc_buffer((T.int64(128), T.int64(128))) @@ -69,7 +69,7 @@ def elementwise_shape_int64( @Ts.prim_func def elementwise_reindex_cache_read( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ): B = Ts.sblock_alloc_buffer((128, 128)) B_shared = Ts.sblock_alloc_buffer((128, 64, 2), scope="shared") @@ -95,7 +95,7 @@ def elementwise_reindex_cache_read( @Ts.prim_func def elementwise_reindex_cache_write( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ): B = Ts.sblock_alloc_buffer((128, 128)) B_shared = Ts.sblock_alloc_buffer((128, 128), scope="shared") @@ -120,7 +120,7 @@ def elementwise_reindex_cache_write( @Ts.prim_func -def reduce(A: T.Buffer((128, 128, 128, 128), "float32"), C: T.Buffer((128, 128), "float32")): +def reduce(A: T.Tensor((128, 128, 128, 128), "float32"), C: T.Tensor((128, 128), "float32")): B = Ts.sblock_alloc_buffer((128, 128, 128), dtype="float32") for i, j, k in T.grid(128, 128, 128): for l in range(128): @@ -138,7 +138,7 @@ def reduce(A: T.Buffer((128, 128, 128, 128), "float32"), C: T.Buffer((128, 128), @Ts.prim_func def reduce_reindex_cache_write_0( - A: T.Buffer((128, 128, 128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128, 128, 128), "float32"), C: T.Tensor((128, 128), "float32") ): B = Ts.sblock_alloc_buffer((128, 128, 128)) B_shared = Ts.sblock_alloc_buffer((128, 128, 128), scope="shared") @@ -167,7 +167,7 @@ def reduce_reindex_cache_write_0( @Ts.prim_func def reduce_reindex_cache_write_1( - A: T.Buffer((128, 128, 128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128, 128, 128), "float32"), C: T.Tensor((128, 128), "float32") ): B = Ts.sblock_alloc_buffer((128, 128, 128)) B_shared = Ts.sblock_alloc_buffer((128, 128, 128), scope="shared") @@ -202,7 +202,7 @@ def reduce_reindex_cache_write_1( @Ts.prim_func -def func_nested_seq(B: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def func_nested_seq(B: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: A = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -227,7 +227,7 @@ def func_nested_seq(B: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: @Ts.prim_func -def access_under_scope(B: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def access_under_scope(B: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: A = Ts.sblock_alloc_buffer((128, 128)) for i0, j0 in T.grid(8, 8): @@ -251,10 +251,10 @@ def access_under_scope(B: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None @Ts.prim_func def opaque_access( - A: T.Buffer((128, 128), dtype="float16"), - B: T.Buffer((128, 128), dtype="float16"), - C: T.Buffer((128, 128), dtype="float16"), - D: T.Buffer((128, 128), dtype="float16"), + A: T.Tensor((128, 128), dtype="float16"), + B: T.Tensor((128, 128), dtype="float16"), + C: T.Tensor((128, 128), dtype="float16"), + D: T.Tensor((128, 128), dtype="float16"), ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("load_store"): @@ -418,7 +418,7 @@ def func_with_block_predicate() -> None: @Ts.prim_func -def inplace_func(data_io: T.Buffer((64), "int32")): +def inplace_func(data_io: T.Tensor((64), "int32")): data_1d = Ts.sblock_alloc_buffer([64], dtype="int32") for i0 in T.serial(64): with Ts.sblock("copy_in"): @@ -436,7 +436,7 @@ def inplace_func(data_io: T.Buffer((64), "int32")): @Ts.prim_func -def inplace_call(data_io: T.Buffer((64), "int32")): +def inplace_call(data_io: T.Tensor((64), "int32")): for i0 in T.serial(1): with Ts.sblock("ext_call"): Ts.reads(data_io[:64]) @@ -446,7 +446,7 @@ def inplace_call(data_io: T.Buffer((64), "int32")): @Ts.prim_func def cache_read_nested_seq_target( - B: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + B: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: A = Ts.sblock_alloc_buffer([128, 128], dtype="float32") A_global = Ts.sblock_alloc_buffer([128, 128], dtype="float32") @@ -487,9 +487,9 @@ def cache_read_nested_seq_target( @Ts.prim_func def nested_buffer_access( - A: T.Buffer((T.int64(7), T.int64(512)), dtype="float32"), - B: T.Buffer(T.int64(1), dtype="int32"), - C: T.Buffer((T.int64(1), T.int64(512)), dtype="float32"), + A: T.Tensor((T.int64(7), T.int64(512)), dtype="float32"), + B: T.Tensor(T.int64(1), dtype="int32"), + C: T.Tensor((T.int64(1), T.int64(512)), dtype="float32"), ): for ax0, ax1 in T.grid(T.int64(1), T.int64(512)): with Ts.sblock("C"): @@ -503,7 +503,7 @@ def nested_buffer_access( @Ts.prim_func -def cache_read_elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def cache_read_elementwise(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) A_global = Ts.sblock_alloc_buffer((128, 128)) B_local = Ts.sblock_alloc_buffer((128, 128), scope="local") @@ -526,7 +526,7 @@ def cache_read_elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> @Ts.prim_func -def cache_read_under_scope(B: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def cache_read_under_scope(B: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: A = Ts.sblock_alloc_buffer((128, 128)) A_global = Ts.sblock_alloc_buffer((128, 128)) @@ -562,10 +562,10 @@ def cache_read_under_scope(B: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> @Ts.prim_func def cache_read_opaque_access( - A: T.Buffer((128, 128), dtype="float16"), - B: T.Buffer((128, 128), dtype="float16"), - C: T.Buffer((128, 128), dtype="float16"), - D: T.Buffer((128, 128), dtype="float16"), + A: T.Tensor((128, 128), dtype="float16"), + B: T.Tensor((128, 128), dtype="float16"), + C: T.Tensor((128, 128), dtype="float16"), + D: T.Tensor((128, 128), dtype="float16"), ) -> None: A_global = Ts.sblock_alloc_buffer((128, 128), dtype="float16") @@ -700,7 +700,7 @@ def cache_read_multi_consumer_target() -> None: @Ts.prim_func -def continuous_cache_read(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def continuous_cache_read(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) B_shared = Ts.sblock_alloc_buffer((128, 128), scope="shared") B_local = Ts.sblock_alloc_buffer((128, 128), scope="local") @@ -745,8 +745,8 @@ def block_predicate_cache_read() -> None: @Ts.prim_func def cache_read_shape_int64( - A: T.Buffer((T.int64(128), T.int64(128)), dtype="float32"), - C: T.Buffer((T.int64(128), T.int64(128)), dtype="float32"), + A: T.Tensor((T.int64(128), T.int64(128)), dtype="float32"), + C: T.Tensor((T.int64(128), T.int64(128)), dtype="float32"), ) -> None: B = Ts.sblock_alloc_buffer([T.int64(128), T.int64(128)], dtype="float32") A_global = Ts.sblock_alloc_buffer([T.int64(128), T.int64(128)], dtype="float32") @@ -771,7 +771,7 @@ def cache_read_shape_int64( @Ts.prim_func -def cache_read_inplace(data_io: T.Buffer(64, "int32")) -> None: +def cache_read_inplace(data_io: T.Tensor(64, "int32")) -> None: data_1d = Ts.sblock_alloc_buffer([64], dtype="int32") data_io_local = Ts.sblock_alloc_buffer([64], dtype="int32", scope="local") for ax0 in T.serial(64): @@ -800,7 +800,7 @@ def cache_read_inplace(data_io: T.Buffer(64, "int32")) -> None: @Ts.prim_func -def cache_inplace_buffer(data_io: T.Buffer(64, "int32")) -> None: +def cache_inplace_buffer(data_io: T.Tensor(64, "int32")) -> None: data_io_local = Ts.sblock_alloc_buffer([64], dtype="int32", scope="local") data_io_global = Ts.sblock_alloc_buffer([64], dtype="int32") data_io_global_1 = Ts.sblock_alloc_buffer([64], dtype="int32") @@ -837,9 +837,9 @@ def cache_inplace_buffer(data_io: T.Buffer(64, "int32")) -> None: @Ts.prim_func def cache_read_nested_buffer_access( - A: T.Buffer((T.int64(7), T.int64(512)), dtype="float32"), - B: T.Buffer(T.int64(1), dtype="int32"), - C: T.Buffer((T.int64(1), T.int64(512)), dtype="float32"), + A: T.Tensor((T.int64(7), T.int64(512)), dtype="float32"), + B: T.Tensor(T.int64(1), dtype="int32"), + C: T.Tensor((T.int64(1), T.int64(512)), dtype="float32"), ): B_global = Ts.sblock_alloc_buffer((T.int64(1),), "int32") for ax0 in range(T.int64(1)): @@ -860,7 +860,7 @@ def cache_read_nested_buffer_access( @Ts.prim_func -def cache_write_elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def cache_write_elementwise(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) B_global = Ts.sblock_alloc_buffer((128, 128), scope="local") C_local = Ts.sblock_alloc_buffer((128, 128)) @@ -883,7 +883,7 @@ def cache_write_elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> @Ts.prim_func -def cache_write_under_scope(B: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def cache_write_under_scope(B: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: A = Ts.sblock_alloc_buffer((128, 128)) A_global = Ts.sblock_alloc_buffer((128, 128)) @@ -925,10 +925,10 @@ def cache_write_under_scope(B: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> @Ts.prim_func def cache_write_opaque_access( - A: T.Buffer((128, 128), dtype="float16"), - B: T.Buffer((128, 128), dtype="float16"), - C: T.Buffer((128, 128), dtype="float16"), - D: T.Buffer((128, 128), dtype="float16"), + A: T.Tensor((128, 128), dtype="float16"), + B: T.Tensor((128, 128), dtype="float16"), + C: T.Tensor((128, 128), dtype="float16"), + D: T.Tensor((128, 128), dtype="float16"), ) -> None: D_global = Ts.sblock_alloc_buffer((128, 128), dtype="float16") B_global = Ts.sblock_alloc_buffer((128, 128), dtype="float16") @@ -1123,7 +1123,7 @@ def cache_write_multi_consumer_all_consume_cache(): @Ts.prim_func -def continuous_cache_write(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def continuous_cache_write(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) B_shared = Ts.sblock_alloc_buffer((128, 128), scope="shared") @@ -1190,9 +1190,9 @@ def block_predicate_cache_write_output_buf() -> None: @Ts.prim_func def symbolic_matmul_blocked( - A: T.Buffer(((n + 31) // 32 * 32, 4)), # noqa: F821 - B: T.Buffer((4, (n + 31) // 32 * 32)), # noqa: F821 - C: T.Buffer(((n + 31) // 32 * 32, (n + 31) // 32 * 32)), # noqa: F821 + A: T.Tensor(((n + 31) // 32 * 32, 4)), # noqa: F821 + B: T.Tensor((4, (n + 31) // 32 * 32)), # noqa: F821 + C: T.Tensor(((n + 31) // 32 * 32, (n + 31) // 32 * 32)), # noqa: F821 n: T.int32, ): for i0_0, i1_0 in T.grid((n + 31) // 32, (n + 31) // 32): @@ -1218,9 +1218,9 @@ def symbolic_matmul_blocked( @Ts.prim_func def symbolic_matmul_blocked_cache_read( - A: T.Buffer(((n + 31) // 32 * 32, 4)), # noqa: F821 - B: T.Buffer((4, (n + 31) // 32 * 32)), # noqa: F821 - C: T.Buffer(((n + 31) // 32 * 32, (n + 31) // 32 * 32)), # noqa: F821 + A: T.Tensor(((n + 31) // 32 * 32, 4)), # noqa: F821 + B: T.Tensor((4, (n + 31) // 32 * 32)), # noqa: F821 + C: T.Tensor(((n + 31) // 32 * 32, (n + 31) // 32 * 32)), # noqa: F821 n: T.int32, ): for i0_0, i1_0 in T.grid((n + 31) // 32, (n + 31) // 32): @@ -1254,9 +1254,9 @@ def symbolic_matmul_blocked_cache_read( @Ts.prim_func def symbolic_matmul_blocked_cache_write( - A: T.Buffer(((n + 31) // 32 * 32, 4)), # noqa: F821 - B: T.Buffer((4, (n + 31) // 32 * 32)), # noqa: F821 - C: T.Buffer(((n + 31) // 32 * 32, (n + 31) // 32 * 32)), # noqa: F821 + A: T.Tensor(((n + 31) // 32 * 32, 4)), # noqa: F821 + B: T.Tensor((4, (n + 31) // 32 * 32)), # noqa: F821 + C: T.Tensor(((n + 31) // 32 * 32, (n + 31) // 32 * 32)), # noqa: F821 n: T.int32, ): for i0_0, i1_0 in T.grid((n + 31) // 32, (n + 31) // 32): @@ -1659,7 +1659,7 @@ def test_symbolic_matmul_blocked_cache_write(use_block_name): def test_cache_write_with_nested_block_predicate(): @Ts.prim_func - def main(A_buf: T.Buffer((12, 24), "float32"), C_buf: T.Buffer((10, 20), "float32")) -> None: + def main(A_buf: T.Tensor((12, 24), "float32"), C_buf: T.Tensor((10, 20), "float32")) -> None: for i, j in T.grid(12, 24): with Ts.sblock("compute"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -1669,7 +1669,7 @@ def main(A_buf: T.Buffer((12, 24), "float32"), C_buf: T.Buffer((10, 20), "float3 C_buf[vi, vj] = A_buf[vi, vj] * 2.0 @Ts.prim_func - def expected(A_buf: T.Buffer((12, 24), "float32"), C_buf: T.Buffer((10, 20), "float32")): + def expected(A_buf: T.Tensor((12, 24), "float32"), C_buf: T.Tensor((10, 20), "float32")): with Ts.sblock("root"): C_buf_local = Ts.sblock_alloc_buffer((10, 20), scope="local") for i, j in T.grid(12, 24): @@ -1697,7 +1697,7 @@ def expected(A_buf: T.Buffer((12, 24), "float32"), C_buf: T.Buffer((10, 20), "fl def test_cache_read_with_nested_block_predicate(): @Ts.prim_func - def main(A_buf: T.Buffer((12, 24), "float32"), C_buf: T.Buffer((10, 20), "float32")) -> None: + def main(A_buf: T.Tensor((12, 24), "float32"), C_buf: T.Tensor((10, 20), "float32")) -> None: for i, j in T.grid(12, 24): with Ts.sblock("compute"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -1707,7 +1707,7 @@ def main(A_buf: T.Buffer((12, 24), "float32"), C_buf: T.Buffer((10, 20), "float3 C_buf[vi, vj] = A_buf[vi, vj] * 2.0 @Ts.prim_func - def expected(A_buf: T.Buffer((12, 24), "float32"), C_buf: T.Buffer((10, 20), "float32")): + def expected(A_buf: T.Tensor((12, 24), "float32"), C_buf: T.Tensor((10, 20), "float32")): with Ts.sblock("root"): A_buf_local = Ts.sblock_alloc_buffer((10, 20), scope="local") for ax0, ax1 in T.grid(10, 20): @@ -1751,7 +1751,7 @@ def test_cache_write_sibling_nested_block_predicates_use_union(): """ @Ts.prim_func - def main(A_buf: T.Buffer((12, 24), "float32"), C_buf: T.Buffer((12, 24), "float32")) -> None: + def main(A_buf: T.Tensor((12, 24), "float32"), C_buf: T.Tensor((12, 24), "float32")) -> None: for i, j in T.grid(12, 24): with Ts.sblock("compute"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -1794,7 +1794,7 @@ def test_cache_read_sibling_nested_block_predicates_use_union(): """ @Ts.prim_func - def main(A_buf: T.Buffer((12, 24), "float32"), C_buf: T.Buffer((12, 24), "float32")) -> None: + def main(A_buf: T.Tensor((12, 24), "float32"), C_buf: T.Tensor((12, 24), "float32")) -> None: for i, j in T.grid(12, 24): with Ts.sblock("compute"): vi, vj = Ts.axis.remap("SS", [i, j]) 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 f2659a81c287..25120dcf740f 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 @@ -34,7 +34,7 @@ # pylint: disable=no-member,invalid-name,unused-variable,line-too-long,redefined-outer-name,unexpected-keyword-arg,too-many-nested-blocks @Ts.prim_func -def two_elementwise(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), 'float32')) -> None: +def two_elementwise(A: T.Tensor((128, 128), 'float32'), C: T.Tensor((128, 128), 'float32')) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -48,7 +48,7 @@ def two_elementwise(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def two_elementwise_after_compute_at(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), 'float32')) -> None: +def two_elementwise_after_compute_at(A: T.Tensor((128, 128), 'float32'), C: T.Tensor((128, 128), 'float32')) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -64,7 +64,7 @@ def two_elementwise_after_compute_at(A: T.Buffer((128, 128), 'float32'), C: T.Bu C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def blockized_1(A: T.Buffer([128, 128], 'float32'), C: T.Buffer([128, 128], 'float32')) -> None: +def blockized_1(A: T.Tensor([128, 128], 'float32'), C: T.Tensor([128, 128], 'float32')) -> None: B = Ts.sblock_alloc_buffer([128, 128], "float32") @@ -90,7 +90,7 @@ def blockized_1(A: T.Buffer([128, 128], 'float32'), C: T.Buffer([128, 128], 'flo C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def blockized_after_compute_at(A: T.Buffer([128, 128], 'float32'), C: T.Buffer([128, 128], 'float32')) -> None: +def blockized_after_compute_at(A: T.Tensor([128, 128], 'float32'), C: T.Tensor([128, 128], 'float32')) -> None: B = Ts.sblock_alloc_buffer([128, 128], "float32") @@ -117,7 +117,7 @@ def blockized_after_compute_at(A: T.Buffer([128, 128], 'float32'), C: T.Buffer([ C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def blockized_2(A: T.Buffer([128, 128], 'float32'), C: T.Buffer([128, 128], 'float32')) -> None: +def blockized_2(A: T.Tensor([128, 128], 'float32'), C: T.Tensor([128, 128], 'float32')) -> None: B = Ts.sblock_alloc_buffer([128, 128], "float32") @@ -144,7 +144,7 @@ def blockized_2(A: T.Buffer([128, 128], 'float32'), C: T.Buffer([128, 128], 'flo C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def blockized_2_after_reverse_compute_at(A: T.Buffer([128, 128], 'float32'), C: T.Buffer([128, 128], 'float32')) -> None: +def blockized_2_after_reverse_compute_at(A: T.Tensor([128, 128], 'float32'), C: T.Tensor([128, 128], 'float32')) -> None: B = Ts.sblock_alloc_buffer([128, 128], "float32") @@ -173,7 +173,7 @@ def blockized_2_after_reverse_compute_at(A: T.Buffer([128, 128], 'float32'), C: C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def blockized_2_after_compute_at(A: T.Buffer([128, 128], 'float32'), C: T.Buffer([128, 128], 'float32')) -> None: +def blockized_2_after_compute_at(A: T.Tensor([128, 128], 'float32'), C: T.Tensor([128, 128], 'float32')) -> None: B = Ts.sblock_alloc_buffer([128, 128], "float32") @@ -202,7 +202,7 @@ def blockized_2_after_compute_at(A: T.Buffer([128, 128], 'float32'), C: T.Buffer C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def cuda_matmul_0(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048], 'float32'), C: T.Buffer([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable +def cuda_matmul_0(A: T.Tensor([2048, 2048], 'float32'), B: T.Tensor([2048, 2048], 'float32'), C: T.Tensor([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable A_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") B_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") @@ -244,7 +244,7 @@ def cuda_matmul_0(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048] C[v0_4, v1_4] = C_local[v0_4, v1_4] @Ts.prim_func -def cuda_matmul_0_after_compute_at(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048], 'float32'), C: T.Buffer([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable +def cuda_matmul_0_after_compute_at(A: T.Tensor([2048, 2048], 'float32'), B: T.Tensor([2048, 2048], 'float32'), C: T.Tensor([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable A_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") B_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") @@ -288,7 +288,7 @@ def cuda_matmul_0_after_compute_at(A: T.Buffer([2048, 2048], 'float32'), B: T.Bu C[vi, vj] = C_local[vi, vj] @Ts.prim_func -def cuda_matmul_1(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048], 'float32'), C: T.Buffer([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable +def cuda_matmul_1(A: T.Tensor([2048, 2048], 'float32'), B: T.Tensor([2048, 2048], 'float32'), C: T.Tensor([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable A_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") B_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") @@ -334,7 +334,7 @@ def cuda_matmul_1(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048] C[vi, vj] = C_local[vi, vj] @Ts.prim_func -def cuda_matmul_2(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048], 'float32'), C: T.Buffer([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable +def cuda_matmul_2(A: T.Tensor([2048, 2048], 'float32'), B: T.Tensor([2048, 2048], 'float32'), C: T.Tensor([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable A_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") B_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") @@ -381,7 +381,7 @@ def cuda_matmul_2(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048] C[v0, v1] = C_local[v0, v1] @Ts.prim_func -def cuda_matmul_3(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048], 'float32'), C: T.Buffer([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable +def cuda_matmul_3(A: T.Tensor([2048, 2048], 'float32'), B: T.Tensor([2048, 2048], 'float32'), C: T.Tensor([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable A_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") B_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") @@ -429,7 +429,7 @@ def cuda_matmul_3(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048] C[v0, v1] = C_local[v0, v1] @Ts.prim_func -def cuda_matmul_4(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048], 'float32'), C: T.Buffer([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable +def cuda_matmul_4(A: T.Tensor([2048, 2048], 'float32'), B: T.Tensor([2048, 2048], 'float32'), C: T.Tensor([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable A_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") B_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") @@ -478,7 +478,7 @@ def cuda_matmul_4(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048] C[v0, v1] = C_local[v0, v1] @Ts.prim_func -def cuda_matmul_5(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048], 'float32'), C: T.Buffer([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable +def cuda_matmul_5(A: T.Tensor([2048, 2048], 'float32'), B: T.Tensor([2048, 2048], 'float32'), C: T.Tensor([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable A_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") B_shared = Ts.sblock_alloc_buffer([2048, 2048], "float32", scope="shared") @@ -528,7 +528,7 @@ def cuda_matmul_5(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048] C[v0, v1] = C_local[v0, v1] @Ts.prim_func -def tiled(A: T.Buffer([128, 128], 'float32'), C: T.Buffer([128, 128], 'float32')) -> None: +def tiled(A: T.Tensor([128, 128], 'float32'), C: T.Tensor([128, 128], 'float32')) -> None: B = Ts.sblock_alloc_buffer([128, 128], "float32") @@ -543,7 +543,7 @@ def tiled(A: T.Buffer([128, 128], 'float32'), C: T.Buffer([128, 128], 'float32') C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def tiled_after_reverse_compute_at(A: T.Buffer([128, 128], 'float32'), C: T.Buffer([128, 128], 'float32')) -> None: +def tiled_after_reverse_compute_at(A: T.Tensor([128, 128], 'float32'), C: T.Tensor([128, 128], 'float32')) -> None: B = Ts.sblock_alloc_buffer([128, 128], "float32") @@ -560,7 +560,7 @@ def tiled_after_reverse_compute_at(A: T.Buffer([128, 128], 'float32'), C: T.Buff C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def tiled_trivial_binding(A: T.Buffer([1, 128, 128], 'float32'), C: T.Buffer([1, 128, 128], 'float32')) -> None: +def tiled_trivial_binding(A: T.Tensor([1, 128, 128], 'float32'), C: T.Tensor([1, 128, 128], 'float32')) -> None: B = Ts.sblock_alloc_buffer([1, 128, 128], "float32") @@ -575,7 +575,7 @@ def tiled_trivial_binding(A: T.Buffer([1, 128, 128], 'float32'), C: T.Buffer([1, C[0, vi, vj] = B[0, vi, vj] + 1.0 @Ts.prim_func -def tiled_trivial_binding_after_reverse_compute_at(A: T.Buffer([1, 128, 128], 'float32'), C: T.Buffer([1, 128, 128], 'float32')) -> None: +def tiled_trivial_binding_after_reverse_compute_at(A: T.Tensor([1, 128, 128], 'float32'), C: T.Tensor([1, 128, 128], 'float32')) -> None: B = Ts.sblock_alloc_buffer([1, 128, 128], "float32") @@ -592,7 +592,7 @@ def tiled_trivial_binding_after_reverse_compute_at(A: T.Buffer([1, 128, 128], 'f C[0, vi, vj] = B[0, vi, vj] + 1.0 @Ts.prim_func -def factorized(A: T.Buffer([16, 16, 16], 'float32'), B: T.Buffer([16], 'float32')) -> None: +def factorized(A: T.Tensor([16, 16, 16], 'float32'), B: T.Tensor([16], 'float32')) -> None: B_rf_local = Ts.sblock_alloc_buffer([16, 16], "float32", scope="local") for j in T.thread_binding(0, 16, thread = "blockIdx.x"): @@ -612,7 +612,7 @@ def factorized(A: T.Buffer([16, 16, 16], 'float32'), B: T.Buffer([16], 'float32' B[vi] = B[vi] + B_rf_local[vk, vi] @Ts.prim_func -def factorized_after_reverse_compute_at(A: T.Buffer([16, 16, 16], 'float32'), B: T.Buffer([16], 'float32')) -> None: +def factorized_after_reverse_compute_at(A: T.Tensor([16, 16, 16], 'float32'), B: T.Tensor([16], 'float32')) -> None: B_rf_local = Ts.sblock_alloc_buffer([16, 16], "float32", scope="local") for j in T.thread_binding(0, 16, thread = "blockIdx.x"): @@ -634,7 +634,7 @@ def factorized_after_reverse_compute_at(A: T.Buffer([16, 16, 16], 'float32'), B: B[vi] = B[vi] + B_rf_local[vk, vi] @Ts.prim_func -def not_all_compact_data_flow(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), 'float32')): +def not_all_compact_data_flow(A: T.Tensor((128, 128), 'float32'), C: T.Tensor((128, 128), 'float32')): B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -651,7 +651,7 @@ def not_all_compact_data_flow(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((1 C[vi, vj * 2 + 1] = B[vi, vj * 2 + 1] * 2.0 @Ts.prim_func -def not_all_compact_data_flow_after_compute_at(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), 'float32')): +def not_all_compact_data_flow_after_compute_at(A: T.Tensor((128, 128), 'float32'), C: T.Tensor((128, 128), 'float32')): B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -669,7 +669,7 @@ def not_all_compact_data_flow_after_compute_at(A: T.Buffer((128, 128), 'float32' C[vi, vj * 2 + 1] = B[vi, vj * 2 + 1] * 2.0 @Ts.prim_func -def fail_subtree_compact_dataflow(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), 'float32')) -> None: +def fail_subtree_compact_dataflow(A: T.Tensor((128, 128), 'float32'), C: T.Tensor((128, 128), 'float32')) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -690,7 +690,7 @@ def fail_subtree_compact_dataflow(A: T.Buffer((128, 128), 'float32'), C: T.Buffe C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def fail_all_consumers_under_loop(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), 'float32'), D: T.Buffer((128, 128), 'float32')) -> None: +def fail_all_consumers_under_loop(A: T.Tensor((128, 128), 'float32'), C: T.Tensor((128, 128), 'float32'), D: T.Tensor((128, 128), 'float32')) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -708,7 +708,7 @@ def fail_all_consumers_under_loop(A: T.Buffer((128, 128), 'float32'), C: T.Buffe D[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def fail_all_producers_under_loop(A: T.Buffer((128, 128), 'float32'), D: T.Buffer((128, 128), 'float32')) -> None: +def fail_all_producers_under_loop(A: T.Tensor((128, 128), 'float32'), D: T.Tensor((128, 128), 'float32')) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") C = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -727,7 +727,7 @@ def fail_all_producers_under_loop(A: T.Buffer((128, 128), 'float32'), D: T.Buffe D[vi, vj] = B[vi, vj] + C[vi, vj] @Ts.prim_func -def read_out_of_bound(A: T.Buffer([16], 'float32'), C: T.Buffer([16], 'float32')) -> None: +def read_out_of_bound(A: T.Tensor([16], 'float32'), C: T.Tensor([16], 'float32')) -> None: B = Ts.sblock_alloc_buffer([16], "float32") @@ -742,7 +742,7 @@ def read_out_of_bound(A: T.Buffer([16], 'float32'), C: T.Buffer([16], 'float32') C[v] = T.if_then_else(v < 15, T.max(B[v], B[v + 1]), B[v]) @Ts.prim_func -def read_out_of_bound_after_compute_at(A: T.Buffer([16], 'float32'), C: T.Buffer([16], 'float32')) -> None: +def read_out_of_bound_after_compute_at(A: T.Tensor([16], 'float32'), C: T.Tensor([16], 'float32')) -> None: B = Ts.sblock_alloc_buffer([16], "float32") @@ -758,7 +758,7 @@ def read_out_of_bound_after_compute_at(A: T.Buffer([16], 'float32'), C: T.Buffer C[v] = T.if_then_else(v < 15, T.max(B[v], B[v + 1]), B[v]) @Ts.prim_func -def multi_reduction(A: T.Buffer((16, 16), "float32"), C: T.Buffer((), "float32")): +def multi_reduction(A: T.Tensor((16, 16), "float32"), C: T.Tensor((), "float32")): B = Ts.sblock_alloc_buffer((16, ), dtype="float32") for i, k in T.grid(16, 16): with Ts.sblock("B"): @@ -775,8 +775,8 @@ def multi_reduction(A: T.Buffer((16, 16), "float32"), C: T.Buffer((), "float32") @Ts.prim_func def multi_reduction_after_compute_at( - A: T.Buffer((16, 16), "float32"), - C:T.Buffer((), "float32"), + A: T.Tensor((16, 16), "float32"), + C:T.Tensor((), "float32"), ): B = Ts.sblock_alloc_buffer((16, ), dtype="float32") for k in T.grid(16): @@ -793,7 +793,7 @@ def multi_reduction_after_compute_at( C[()] += B[vk] @Ts.prim_func -def tiled_pooling_read_cache(X: T.Buffer([224, 224], dtype='float32'), Y: T.Buffer([224, 224], dtype='float32')) -> None: +def tiled_pooling_read_cache(X: T.Tensor([224, 224], dtype='float32'), Y: T.Tensor([224, 224], dtype='float32')) -> None: cache = Ts.sblock_alloc_buffer([224, 224], dtype="float32") for hh, ww in T.grid(224, 224): @@ -815,7 +815,7 @@ def tiled_pooling_read_cache(X: T.Buffer([224, 224], dtype='float32'), Y: T.Buff cache[h + kh - 1, w + kw - 1], 0.0)) @Ts.prim_func -def tiled_pooling_read_cache_after_compute_at(X: T.Buffer([224, 224], dtype='float32'), Y: T.Buffer([224, 224], dtype='float32')) -> None: +def tiled_pooling_read_cache_after_compute_at(X: T.Tensor([224, 224], dtype='float32'), Y: T.Tensor([224, 224], dtype='float32')) -> None: cache = Ts.sblock_alloc_buffer([224, 224], dtype="float32") for hh_0, ww_0 in T.grid(28, 28): @@ -840,9 +840,9 @@ def tiled_pooling_read_cache_after_compute_at(X: T.Buffer([224, 224], dtype='flo cache[h + kh - 1, w + kw - 1], 0.0)) @Ts.prim_func -def non_uniform_tiled_conv(x: T.Buffer((1, 3, 100, 100), "float32"), - w: T.Buffer((16, 3, 3, 3), "float32"), - y: T.Buffer((1, 16, 98, 98), "float32")) -> None: +def non_uniform_tiled_conv(x: T.Tensor((1, 3, 100, 100), "float32"), + w: T.Tensor((16, 3, 3, 3), "float32"), + y: T.Tensor((1, 16, 98, 98), "float32")) -> None: x_global = Ts.sblock_alloc_buffer([1, 3, 100, 100], dtype="float32") for ax0, ax1, ax2, ax3 in T.grid(1, 3, 100, 100): with Ts.sblock("cache"): @@ -862,9 +862,9 @@ def non_uniform_tiled_conv(x: T.Buffer((1, 3, 100, 100), "float32"), x_global[nn, cc // 16 * 3 + rc, hh + rh, ww + rw] * w[cc, rc, rh, rw] @Ts.prim_func -def non_uniform_tiled_conv_after_compute_at(x: T.Buffer((1, 3, 100, 100), "float32"), - w: T.Buffer((16, 3, 3, 3), "float32"), - y: T.Buffer((1, 16, 98, 98), "float32")) -> None: +def non_uniform_tiled_conv_after_compute_at(x: T.Tensor((1, 3, 100, 100), "float32"), + w: T.Tensor((16, 3, 3, 3), "float32"), + y: T.Tensor((1, 16, 98, 98), "float32")) -> None: x_global = Ts.sblock_alloc_buffer([1, 3, 100, 100], dtype="float32") for h_o, w_o in T.grid(7, 7): for ax0, ax1, ax2 in T.grid(3, 17, 17): @@ -889,9 +889,9 @@ def non_uniform_tiled_conv_after_compute_at(x: T.Buffer((1, 3, 100, 100), "float x_global[nn, cc // 16 * 3 + rc, hh + rh, ww + rw] * w[cc, rc, rh, rw] @Ts.prim_func -def concat_two_elemwise(x: T.Buffer((16,), "float32"), - y: T.Buffer((8,), "float32"), - T_concat: T.Buffer((24,), "float32")) -> None: +def concat_two_elemwise(x: T.Tensor((16,), "float32"), + y: T.Tensor((8,), "float32"), + T_concat: T.Tensor((24,), "float32")) -> None: T_add_1 = Ts.sblock_alloc_buffer([16], dtype="float32") T_add_2 = Ts.sblock_alloc_buffer([8], dtype="float32") for i in T.serial(16): @@ -908,9 +908,9 @@ def concat_two_elemwise(x: T.Buffer((16,), "float32"), T_concat[ax] = T.if_then_else(16 <= ax, T_add_2[ax - 16], T_add_1[ax]) @Ts.prim_func -def concat_two_elemwise_after_compute_at(x: T.Buffer((16,), "float32"), - y: T.Buffer((8,), "float32"), - T_concat: T.Buffer((24,), "float32")) -> None: +def concat_two_elemwise_after_compute_at(x: T.Tensor((16,), "float32"), + y: T.Tensor((8,), "float32"), + T_concat: T.Tensor((24,), "float32")) -> None: T_add_1 = Ts.sblock_alloc_buffer([16], dtype="float32") T_add_2 = Ts.sblock_alloc_buffer([8], dtype="float32") for i in T.serial(24): @@ -927,7 +927,7 @@ def concat_two_elemwise_after_compute_at(x: T.Buffer((16,), "float32"), T_concat[ax] = T.if_then_else(16 <= ax, T_add_2[ax - 16], T_add_1[ax]) @Ts.prim_func -def floordiv_and_floormod_indices(X: T.Buffer([16, 16]), Y: T.Buffer([256])) -> None: +def floordiv_and_floormod_indices(X: T.Tensor([16, 16]), Y: T.Tensor([256])) -> None: temp = Ts.sblock_alloc_buffer([16, 16]) for i, j in T.grid(16, 16): @@ -940,7 +940,7 @@ def floordiv_and_floormod_indices(X: T.Buffer([16, 16]), Y: T.Buffer([256])) -> Y[v_i] = temp[v_i // 16, v_i % 16] @Ts.prim_func -def floordiv_and_floormod_indices_after_reverse_compute_at(X: T.Buffer([16, 16], dtype='float32'), Y: T.Buffer([256], dtype='float32')) -> None: +def floordiv_and_floormod_indices_after_reverse_compute_at(X: T.Tensor([16, 16], dtype='float32'), Y: T.Tensor([256], dtype='float32')) -> None: temp = Ts.sblock_alloc_buffer([16, 16], dtype="float32") for i in T.serial(0, 16): @@ -954,8 +954,8 @@ def floordiv_and_floormod_indices_after_reverse_compute_at(X: T.Buffer([16, 16], Y[v_i] = temp[v_i // 16, v_i % 16] @Ts.prim_func -def recursive_floordiv_floormod(A: T.Buffer((16, 64, 1, 8, 8, 32), "float32"), - C: T.Buffer((3, 512, 512), "float32")) -> None: +def recursive_floordiv_floormod(A: T.Tensor((16, 64, 1, 8, 8, 32), "float32"), + C: T.Tensor((3, 512, 512), "float32")) -> None: T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): B = Ts.sblock_alloc_buffer((1, 128, 16, 8, 2, 32, 2), "float32") @@ -973,7 +973,7 @@ def recursive_floordiv_floormod(A: T.Buffer((16, 64, 1, 8, 8, 32), "float32"), C[v1, v2, v3] = B[v1 // 8, v2 // 4, v3 // 32, v1, v2 % 4 // 2, v3 % 32, v2 % 2] * 2 @Ts.prim_func -def recursive_floordiv_floormod_after_reverse_compute_at(A: T.Buffer((16, 64, 1, 8, 8, 32), "float32"), C: T.Buffer((3, 512, 512), "float32")) -> None: +def recursive_floordiv_floormod_after_reverse_compute_at(A: T.Tensor((16, 64, 1, 8, 8, 32), "float32"), C: T.Tensor((3, 512, 512), "float32")) -> None: T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): B = Ts.sblock_alloc_buffer((1, 128, 16, 8, 2, 32, 2)) @@ -994,7 +994,7 @@ def recursive_floordiv_floormod_after_reverse_compute_at(A: T.Buffer((16, 64, 1, C[v1, v2, v3] = B[v1 // 8, v2 // 4, v3 // 32, v1, v2 % 4 // 2, v3 % 32, v2 % 2] * T.float32(2) @Ts.prim_func -def tiled_repeat_op(x: T.Buffer((4,), "float32"), T_repeat: T.Buffer((64,), "float32")) -> None: +def tiled_repeat_op(x: T.Tensor((4,), "float32"), T_repeat: T.Tensor((64,), "float32")) -> None: T_add = Ts.sblock_alloc_buffer([4], dtype="float32") for i0 in T.serial(4): with Ts.sblock("T_add"): @@ -1006,7 +1006,7 @@ def tiled_repeat_op(x: T.Buffer((4,), "float32"), T_repeat: T.Buffer((64,), "flo T_repeat[ax0] = T_add[ax0 // 16] @Ts.prim_func -def tiled_repeat_op_after_compute_at(x: T.Buffer((4,), "float32"), T_repeat: T.Buffer((64,), "float32")) -> None: +def tiled_repeat_op_after_compute_at(x: T.Tensor((4,), "float32"), T_repeat: T.Tensor((64,), "float32")) -> None: T_add = Ts.sblock_alloc_buffer([4], dtype="float32") for i0_0 in T.serial(8): with Ts.sblock("T_add"): @@ -1018,7 +1018,7 @@ def tiled_repeat_op_after_compute_at(x: T.Buffer((4,), "float32"), T_repeat: T.B T_repeat[ax0] = T_add[ax0 // 16] @Ts.prim_func -def static_bound(A: T.Buffer((32, 1), "float32"), C: T.Buffer((32, 1), "float32")) -> None: +def static_bound(A: T.Tensor((32, 1), "float32"), C: T.Tensor((32, 1), "float32")) -> None: B = Ts.sblock_alloc_buffer((32, 1), "float32") for i, j in T.grid(32, 1): with Ts.sblock("B"): @@ -1033,7 +1033,7 @@ def static_bound(A: T.Buffer((32, 1), "float32"), C: T.Buffer((32, 1), "float32" C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def static_bound_after_compute_at(A: T.Buffer((32, 1), "float32"), C: T.Buffer((32, 1), "float32")) -> None: +def static_bound_after_compute_at(A: T.Tensor((32, 1), "float32"), C: T.Tensor((32, 1), "float32")) -> None: B = Ts.sblock_alloc_buffer((32, 1), "float32") for i in range(32): for ax0, ax1 in T.grid(1, 1): @@ -1180,7 +1180,7 @@ def test_compute_at_tiled_repeat_op(use_block_name): def test_compute_at_rev_iter(): @Ts.prim_func - def before(X: T.Buffer((10, 10), "float32"), Z: T.Buffer((10, 10), "float32")): + def before(X: T.Tensor((10, 10), "float32"), Z: T.Tensor((10, 10), "float32")): Y = Ts.sblock_alloc_buffer([10, 10], "float32") for i, j in T.grid(10, 10): with Ts.sblock("b0"): @@ -1192,7 +1192,7 @@ def before(X: T.Buffer((10, 10), "float32"), Z: T.Buffer((10, 10), "float32")): Z[vi, vj] = Y[vj, vi] + 2.0 @Ts.prim_func - def after(X: T.Buffer((10, 10), "float32"), Z: T.Buffer((10, 10), "float32")): + def after(X: T.Tensor((10, 10), "float32"), Z: T.Tensor((10, 10), "float32")): Y = Ts.sblock_alloc_buffer([10, 10], "float32") for i in range(10): for j in range(10): @@ -1311,8 +1311,8 @@ def test_compute_at_simplify_symbolic_predicate(): class Before: @Ts.prim_func def main( - X: T.Buffer((T.int64(8), n * 32), "float32"), # noqa: F821 - Y: T.Buffer((T.int64(8), n * 32), "float32"), # noqa: F821 + X: T.Tensor((T.int64(8), n * 32), "float32"), # noqa: F821 + Y: T.Tensor((T.int64(8), n * 32), "float32"), # noqa: F821 n: T.int64, ): for i, k in T.grid(T.int64(8), n * 32): @@ -1324,8 +1324,8 @@ def main( class After: @Ts.prim_func def main( - X: T.Buffer((T.int64(8), n * T.int64(32))), # noqa: F821 - Y: T.Buffer((T.int64(8), n * T.int64(32))), # noqa: F821 + X: T.Tensor((T.int64(8), n * T.int64(32))), # noqa: F821 + Y: T.Tensor((T.int64(8), n * T.int64(32))), # noqa: F821 n: T.int64, ): X_global = Ts.sblock_alloc_buffer((T.int64(8), n * T.int64(32))) @@ -1354,7 +1354,7 @@ def main( def test_compute_at_non_perfect_channel_group(use_block_name): @Ts.prim_func def grouped_channel_bias( - X: T.Buffer((720, 8, 8), "float32"), Y: T.Buffer((720, 8, 8), "float32") + X: T.Tensor((720, 8, 8), "float32"), Y: T.Tensor((720, 8, 8), "float32") ): B = Ts.sblock_alloc_buffer([45], dtype="float32", scope="") for i in T.grid(45): @@ -1369,7 +1369,7 @@ def grouped_channel_bias( @Ts.prim_func def grouped_channel_bias_non_perfect_tiled( - X: T.Buffer((720, 8, 8), "float32"), Y: T.Buffer((720, 8, 8), "float32") + X: T.Tensor((720, 8, 8), "float32"), Y: T.Tensor((720, 8, 8), "float32") ): B = Ts.sblock_alloc_buffer([45], dtype="float32") for c_o in range(2): @@ -1461,9 +1461,9 @@ def _create_prim_func(): def test_compute_at_to_index(): @Ts.prim_func def multi_producers_conv( - data: T.Buffer((1, 3, 224, 224), "int8"), - w: T.Buffer((16, 3, 7, 7), "int8"), - conv: T.Buffer((1, 16, 112, 112), "int32"), + data: T.Tensor((1, 3, 224, 224), "int8"), + w: T.Tensor((16, 3, 7, 7), "int8"), + conv: T.Tensor((1, 16, 112, 112), "int32"), ) -> None: pad = Ts.sblock_alloc_buffer([1, 3, 230, 230], dtype="int8") wbuf = Ts.sblock_alloc_buffer([16, 3, 7, 7], dtype="int8") @@ -1499,9 +1499,9 @@ def multi_producers_conv( @Ts.prim_func def multi_producers_after_compute_at( - data: T.Buffer((1, 3, 224, 224), "int8"), - w: T.Buffer((16, 3, 7, 7), "int8"), - conv: T.Buffer((1, 16, 112, 112), "int32"), + data: T.Tensor((1, 3, 224, 224), "int8"), + w: T.Tensor((16, 3, 7, 7), "int8"), + conv: T.Tensor((1, 16, 112, 112), "int32"), ) -> None: pad = Ts.sblock_alloc_buffer([1, 3, 230, 230], dtype="int8") wbuf = Ts.sblock_alloc_buffer([16, 3, 7, 7], dtype="int8") @@ -1547,7 +1547,7 @@ def multi_producers_after_compute_at( def test_reverse_compute_at_to_index(): @Ts.prim_func - def main(A: T.Buffer((128, 128), "float32"), D: T.Buffer((128, 128), "float32")) -> None: + def main(A: T.Tensor((128, 128), "float32"), D: T.Tensor((128, 128), "float32")) -> None: B = Ts.sblock_alloc_buffer([128, 128], dtype="float32") C = Ts.sblock_alloc_buffer([128, 128], dtype="float32") for i_0, j_0, i_1 in T.grid(8, 8, 16): @@ -1574,7 +1574,7 @@ def main(A: T.Buffer((128, 128), "float32"), D: T.Buffer((128, 128), "float32")) @Ts.prim_func def main_reverse_compute_at( - A: T.Buffer((128, 128), "float32"), D: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), D: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer([128, 128], dtype="float32") C = Ts.sblock_alloc_buffer([128, 128], dtype="float32") @@ -1610,7 +1610,7 @@ def main_reverse_compute_at( def test_reverse_compute_at_with_unit_loop(): @Ts.prim_func - def main(A: T.Buffer((128, 128), "float32"), D: T.Buffer((1, 2, 1), "float32")) -> None: + def main(A: T.Tensor((128, 128), "float32"), D: T.Tensor((1, 2, 1), "float32")) -> None: B = Ts.sblock_alloc_buffer([128, 128], dtype="float32") for i_0, j_0, i_1 in T.grid(T.int64(8), T.int64(8), T.int64(16)): for j_1 in T.serial(T.int64(16)): @@ -1629,7 +1629,7 @@ def main(A: T.Buffer((128, 128), "float32"), D: T.Buffer((1, 2, 1), "float32")) @Ts.prim_func def main_reverse_compute_at( - A: T.Buffer((128, 128), "float32"), D: T.Buffer((1, 2, 1), "float32") + A: T.Tensor((128, 128), "float32"), D: T.Tensor((1, 2, 1), "float32") ): B = Ts.sblock_alloc_buffer([128, 128], dtype="float32") for i_0, j_0, i_1 in T.grid(T.int64(8), T.int64(8), T.int64(16)): @@ -1662,7 +1662,7 @@ def main_reverse_compute_at( def test_reverse_compute_at_layout_trans(): @Ts.prim_func - def before(A: T.Buffer((1, 3, 5, 5, 16), "float32"), C: T.Buffer((1, 6, 5, 5, 8), "float32")): + def before(A: T.Tensor((1, 3, 5, 5, 16), "float32"), C: T.Tensor((1, 6, 5, 5, 8), "float32")): B = Ts.sblock_alloc_buffer((1, 3, 5, 5, 16)) for i0, i1, i2, i3, i4 in T.grid(1, 3, 5, 5, 16): with Ts.sblock("compute"): @@ -1678,7 +1678,7 @@ def before(A: T.Buffer((1, 3, 5, 5, 16), "float32"), C: T.Buffer((1, 6, 5, 5, 8) ] @Ts.prim_func - def after(A: T.Buffer((1, 3, 5, 5, 16), "float32"), C: T.Buffer((1, 6, 5, 5, 8), "float32")): + def after(A: T.Tensor((1, 3, 5, 5, 16), "float32"), C: T.Tensor((1, 6, 5, 5, 8), "float32")): B = Ts.sblock_alloc_buffer((1, 3, 5, 5, 16)) for i0, i1 in T.grid(1, 3): for i2, i3, i4 in T.grid(5, 5, 16): @@ -1707,7 +1707,7 @@ def test_shape_var_as_bound(): n = T.dynamic("n", "int32") @Ts.prim_func - def before(A: T.Buffer((32, 1, 128)), B: T.Buffer((32, n, 128)), C: T.Buffer((32, 1, n))): + def before(A: T.Tensor((32, 1, 128)), B: T.Tensor((32, n, 128)), C: T.Tensor((32, 1, n))): # with Ts.sblock("root"): C_rf = Ts.sblock_alloc_buffer((128, 32, 1, n)) @@ -1736,7 +1736,7 @@ def before(A: T.Buffer((32, 1, 128)), B: T.Buffer((32, n, 128)), C: T.Buffer((32 n = T.dynamic("n", "int32") @Ts.prim_func - def expected(A: T.Buffer((32, 1, 128), "float32"), B: T.Buffer((32, n, 128)), C: T.Buffer((32, 1, n))): + def expected(A: T.Tensor((32, 1, 128), "float32"), B: T.Tensor((32, n, 128)), C: T.Tensor((32, 1, n))): # with Ts.sblock("root"): C_rf = Ts.sblock_alloc_buffer((128, 32, 1, n)) 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 786758d3ed5f..3662f6ee5031 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 @@ -33,7 +33,7 @@ @Ts.prim_func -def elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -48,7 +48,7 @@ def elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: @Ts.prim_func def elementwise_multi_producer_consumer( - A: T.Buffer((128, 128)), C: T.Buffer((128, 128)), D: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), C: T.Tensor((128, 128)), D: T.Tensor((128, 128)) ) -> None: B = Ts.sblock_alloc_buffer((128, 128)) @@ -68,7 +68,7 @@ def elementwise_multi_producer_consumer( @Ts.prim_func def elementwise_multi_consumer_inlined( - A: T.Buffer((128, 128)), C: T.Buffer((128, 128)), D: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), C: T.Tensor((128, 128)), D: T.Tensor((128, 128)) ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("C"): @@ -81,7 +81,7 @@ def elementwise_multi_consumer_inlined( @Ts.prim_func -def elementwise_standalone(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise_standalone(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -95,7 +95,7 @@ def elementwise_standalone(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> @Ts.prim_func -def elementwise_standalone_dce(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise_standalone_dce(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): with Ts.sblock("C"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -103,7 +103,7 @@ def elementwise_standalone_dce(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) @Ts.prim_func -def elementwise_under_loop(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise_under_loop(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i in T.serial(0, 128): for j in T.serial(0, 128): @@ -117,7 +117,7 @@ def elementwise_under_loop(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> @Ts.prim_func -def elementwise_inlined(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise_inlined(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): with Ts.sblock("C"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -125,7 +125,7 @@ def elementwise_inlined(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> Non @Ts.prim_func -def fail_multi_reader_writer(A: T.Buffer((128, 128)), D: T.Buffer((128, 128))) -> None: +def fail_multi_reader_writer(A: T.Tensor((128, 128)), D: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) C = Ts.sblock_alloc_buffer((128, 128)) @@ -141,7 +141,7 @@ def fail_multi_reader_writer(A: T.Buffer((128, 128)), D: T.Buffer((128, 128))) - @Ts.prim_func -def elementwise_multi_reverse_loads(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise_multi_reverse_loads(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -156,7 +156,7 @@ def elementwise_multi_reverse_loads(A: T.Buffer((128, 128)), C: T.Buffer((128, 1 @Ts.prim_func def elementwise_multi_reverse_loads_inlined( - A: T.Buffer((128, 128)), C: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), C: T.Tensor((128, 128)) ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -166,7 +166,7 @@ def elementwise_multi_reverse_loads_inlined( @Ts.prim_func def elementwise_reverse_affine_load( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((8, 32, 8, 8), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((8, 32, 8, 8), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -184,7 +184,7 @@ def elementwise_reverse_affine_load( @Ts.prim_func def elementwise_reverse_affine_load_inlined( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((8, 32, 8, 8), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((8, 32, 8, 8), "float32") ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -199,9 +199,9 @@ def elementwise_reverse_affine_load_inlined( @Ts.prim_func def elementwise_reverse_affine_load_unit_iter( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((8, 16, 1), "float32"), - D: T.Buffer((1, 8, 16, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((8, 16, 1), "float32"), + D: T.Tensor((1, 8, 16, 128), "float32"), ) -> None: C = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -216,9 +216,9 @@ def elementwise_reverse_affine_load_unit_iter( @Ts.prim_func def elementwise_reverse_affine_load_unit_iter_inlined( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((8, 16, 1), "float32"), - D: T.Buffer((1, 8, 16, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((8, 16, 1), "float32"), + D: T.Tensor((1, 8, 16, 128), "float32"), ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -228,9 +228,9 @@ def elementwise_reverse_affine_load_unit_iter_inlined( @Ts.prim_func def elementwise_reverse_affine_load_unit_iter_simplified( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((8, 16, 1), "float32"), - D: T.Buffer((1, 8, 16, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((8, 16, 1), "float32"), + D: T.Tensor((1, 8, 16, 128), "float32"), ) -> None: C = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -245,9 +245,9 @@ def elementwise_reverse_affine_load_unit_iter_simplified( @Ts.prim_func def elementwise_reverse_affine_load_unit_iter_simplified_inlined( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((8, 16, 1), "float32"), - D: T.Buffer((1, 8, 16, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((8, 16, 1), "float32"), + D: T.Tensor((1, 8, 16, 128), "float32"), ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -257,7 +257,7 @@ def elementwise_reverse_affine_load_unit_iter_simplified_inlined( @Ts.prim_func def elementwise_reverse_affine_chain( - A: T.Buffer((128, 128), "float32"), D: T.Buffer((1, 8, 16, 128), "float32") + A: T.Tensor((128, 128), "float32"), D: T.Tensor((1, 8, 16, 128), "float32") ): B = Ts.sblock_alloc_buffer((128, 128)) C = Ts.sblock_alloc_buffer((8, 16, 128)) @@ -277,7 +277,7 @@ def elementwise_reverse_affine_chain( @Ts.prim_func def elementwise_reverse_affine_chain_inlined( - A: T.Buffer((128, 128), "float32"), D: T.Buffer((1, 8, 16, 128), "float32") + A: T.Tensor((128, 128), "float32"), D: T.Tensor((1, 8, 16, 128), "float32") ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -287,8 +287,8 @@ def elementwise_reverse_affine_chain_inlined( @Ts.prim_func def elementwise_multi_reverse_affine_load( - A: T.Buffer((128, 128), "float32"), - C: T.Buffer((8, 16, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + C: T.Tensor((8, 16, 128), "float32"), ) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -303,8 +303,8 @@ def elementwise_multi_reverse_affine_load( @Ts.prim_func def elementwise_multi_reverse_affine_load_inlined( - A: T.Buffer((128, 128), "float32"), - C: T.Buffer((8, 16, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + C: T.Tensor((8, 16, 128), "float32"), ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -314,7 +314,7 @@ def elementwise_multi_reverse_affine_load_inlined( @Ts.prim_func def elementwise_reverse_non_affine_load( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((8, 16, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((8, 16, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -328,7 +328,7 @@ def elementwise_reverse_non_affine_load( @Ts.prim_func -def opaque_access_load(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def opaque_access_load(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -345,7 +345,7 @@ def opaque_access_load(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None @Ts.prim_func -def opaque_access_store(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def opaque_access_store(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -363,7 +363,7 @@ def opaque_access_store(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> Non @Ts.prim_func -def buffer_matched(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def buffer_matched(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -378,7 +378,7 @@ def buffer_matched(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: @Ts.prim_func -def elementwise_predicate(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise_predicate(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -393,7 +393,7 @@ def elementwise_predicate(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> N @Ts.prim_func -def elementwise_predicate_inlined(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise_predicate_inlined(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): with Ts.sblock("C"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -402,7 +402,7 @@ def elementwise_predicate_inlined(A: T.Buffer((128, 128)), C: T.Buffer((128, 128 @Ts.prim_func -def elementwise_multi_loads(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise_multi_loads(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -416,7 +416,7 @@ def elementwise_multi_loads(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> @Ts.prim_func -def elementwise_multi_loads_inlined(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise_multi_loads_inlined(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): with Ts.sblock("C"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -424,7 +424,7 @@ def elementwise_multi_loads_inlined(A: T.Buffer((128, 128)), C: T.Buffer((128, 1 @Ts.prim_func -def access_opaque_ptr_then_elemwise(A: T.Buffer([1024]), B: T.Buffer([1024])) -> None: +def access_opaque_ptr_then_elemwise(A: T.Tensor([1024]), B: T.Tensor([1024])) -> None: A_cache = Ts.sblock_alloc_buffer([1024]) BB = Ts.sblock_alloc_buffer([1024]) with Ts.sblock("opaque"): @@ -445,7 +445,7 @@ def access_opaque_ptr_then_elemwise(A: T.Buffer([1024]), B: T.Buffer([1024])) -> @Ts.prim_func def access_opaque_ptr_then_elemwise_inline( - A: T.Buffer([1024], dtype="float32"), B: T.Buffer([1024], dtype="float32") + A: T.Tensor([1024], dtype="float32"), B: T.Tensor([1024], dtype="float32") ) -> None: A_cache = Ts.sblock_alloc_buffer([1024], dtype="float32") with Ts.sblock("opaque"): @@ -464,9 +464,9 @@ def access_opaque_ptr_then_elemwise_inline( @Ts.prim_func def matmul_relu( - A: T.Buffer([512, 512], dtype="float32"), - B: T.Buffer([512, 512], dtype="float32"), - compute: T.Buffer([512, 512], dtype="float32"), + A: T.Tensor([512, 512], dtype="float32"), + B: T.Tensor([512, 512], dtype="float32"), + compute: T.Tensor([512, 512], dtype="float32"), ) -> None: C = Ts.sblock_alloc_buffer([512, 512], dtype="float32") for i0, i1, i2 in T.grid(512, 512, 512): @@ -487,7 +487,7 @@ def matmul_relu( @Ts.prim_func def elementwise_output( - A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), B: T.Tensor((128, 128)), C: T.Tensor((128, 128)) ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -501,8 +501,8 @@ def elementwise_output( @Ts.prim_func def inline_block_with_init( - A: T.Buffer((1, 512, 7, 7), "float32"), - B: T.Buffer((1, 512, 1, 1), "float32"), + A: T.Tensor((1, 512, 7, 7), "float32"), + B: T.Tensor((1, 512, 1, 1), "float32"), ) -> None: B_rf = Ts.sblock_alloc_buffer([1, 512, 1, 1, 49], dtype="float32") for i0, i1, i2, i3, i4, i5 in T.grid(1, 512, 1, 1, 49, 1): @@ -538,9 +538,9 @@ def inline_block_with_init( @Ts.prim_func def exp_exp_opaque_access_with_tvm_access_ptr( - lookup_table: T.Buffer((1024,), "int8"), - x: T.Buffer((16,), "float16"), - compute: T.Buffer((16,), "float16"), + lookup_table: T.Tensor((1024,), "int8"), + x: T.Tensor((16,), "float16"), + compute: T.Tensor((16,), "float16"), ) -> None: compute_1 = Ts.sblock_alloc_buffer([16], dtype="float16") for i0 in T.serial(16): @@ -562,9 +562,9 @@ def exp_exp_opaque_access_with_tvm_access_ptr( @Ts.prim_func def exp_exp_opaque_access_with_tvm_access_ptr_inlined( - lookup_table: T.Buffer((1024,), "int8"), - x: T.Buffer((16,), "float16"), - compute: T.Buffer((16,), "float16"), + lookup_table: T.Tensor((1024,), "int8"), + x: T.Tensor((16,), "float16"), + compute: T.Tensor((16,), "float16"), ) -> None: for i0 in T.serial(16): with Ts.sblock("compute_1"): @@ -581,7 +581,7 @@ def exp_exp_opaque_access_with_tvm_access_ptr_inlined( @Ts.prim_func def elementwise_overcomputed_producer( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((127, 127), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((127, 127), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -596,7 +596,7 @@ def elementwise_overcomputed_producer( @Ts.prim_func def elementwise_overcomputed_producer_reverse_inlined( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((127, 127), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((127, 127), "float32") ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -607,7 +607,7 @@ def elementwise_overcomputed_producer_reverse_inlined( @Ts.prim_func def elementwise_overcomputed_producer_simplify_predicate( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((127, 127), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((127, 127), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i in T.grid(16384): @@ -623,7 +623,7 @@ def elementwise_overcomputed_producer_simplify_predicate( @Ts.prim_func def elementwise_overcomputed_producer_simplify_predicate_reverse_inlined( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((127, 127), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((127, 127), "float32") ) -> None: for i in T.grid(16384): with Ts.sblock("B"): @@ -635,7 +635,7 @@ def elementwise_overcomputed_producer_simplify_predicate_reverse_inlined( @Ts.prim_func def elementwise_overcomputed_producer_injective_load( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((127, 127), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((127, 127), "float32") ) -> None: B = Ts.sblock_alloc_buffer((8, 8, 16, 16)) for i0, j0, i1, j1 in T.grid(8, 8, 16, 16): @@ -650,7 +650,7 @@ def elementwise_overcomputed_producer_injective_load( @Ts.prim_func def elementwise_overcomputed_producer_injective_load_reverse_inlined( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((127, 127), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((127, 127), "float32") ) -> None: for i0, j0, i1, j1 in T.grid(8, 8, 16, 16): with Ts.sblock("B"): @@ -661,7 +661,7 @@ def elementwise_overcomputed_producer_injective_load_reverse_inlined( @Ts.prim_func def elementwise_producer_not_cover_consumer( - A: T.Buffer((128, 128), "float32"), D: T.Buffer((256, 128), "float32") + A: T.Tensor((128, 128), "float32"), D: T.Tensor((256, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -676,7 +676,7 @@ def elementwise_producer_not_cover_consumer( @Ts.prim_func def elementwise_producer_is_reduction( - A: T.Buffer((128, 128), "float32"), D: T.Buffer((128), "float32") + A: T.Tensor((128, 128), "float32"), D: T.Tensor((128), "float32") ) -> None: B = Ts.sblock_alloc_buffer(128) for i, j in T.grid(128, 128): @@ -692,7 +692,7 @@ def elementwise_producer_is_reduction( @Ts.prim_func -def elementwise_predicate_producer(A: T.Buffer((128, 128)), C: T.Buffer((127, 128))) -> None: +def elementwise_predicate_producer(A: T.Tensor((128, 128)), C: T.Tensor((127, 128))) -> None: B = Ts.sblock_alloc_buffer((127, 128)) for i, j in T.grid(128, 128): @@ -708,7 +708,7 @@ def elementwise_predicate_producer(A: T.Buffer((128, 128)), C: T.Buffer((127, 12 @Ts.prim_func def elementwise_predicate_producer_inlined( - A: T.Buffer((128, 128)), C: T.Buffer((127, 128)) + A: T.Tensor((128, 128)), C: T.Tensor((127, 128)) ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -723,7 +723,7 @@ def elementwise_predicate_producer_inlined( @tvm.script.ir_module class Conv2dInt8_TensorCore_with_predicate_before: @Ts.prim_func - def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), "int8"), p2: T.Buffer((1, 1, 1, 256), "int32"), p3: T.Buffer((1, 1, 1, 256), "int32"), p4: T.Buffer(256, "int32"), p5: T.Buffer(256, "int32"), p6: T.Buffer(256, "int32"), p7: T.Buffer((), "int32"), p8: T.Buffer(1, "int32"), p9: T.Buffer((16, 56, 56, 256), "int32"), compute: T.Buffer((16, 56, 56, 256), "int32")): + def main(p0: T.Tensor((16, 56, 56, 64), "int8"), p1: T.Tensor((256, 1, 1, 64), "int8"), p2: T.Tensor((1, 1, 1, 256), "int32"), p3: T.Tensor((1, 1, 1, 256), "int32"), p4: T.Tensor(256, "int32"), p5: T.Tensor(256, "int32"), p6: T.Tensor(256, "int32"), p7: T.Tensor((), "int32"), p8: T.Tensor(1, "int32"), p9: T.Tensor((16, 56, 56, 256), "int32"), compute: T.Tensor((16, 56, 56, 256), "int32")): # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -844,7 +844,7 @@ def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), " @tvm.script.ir_module class Conv2dInt8_TensorCore_with_predicate_after: @Ts.prim_func - def main(p0: T.Buffer((16, 56, 56, 64), "int8"), p1: T.Buffer((256, 1, 1, 64), "int8"), p2: T.Buffer((1, 1, 1, 256), "int32"), p3: T.Buffer((1, 1, 1, 256), "int32"), p4: T.Buffer((256,), "int32"), p5: T.Buffer((256,), "int32"), p6: T.Buffer((256,), "int32"), p7: T.Buffer((), "int32"), p8: T.Buffer((1,), "int32"), p9: T.Buffer((16, 56, 56, 256), "int32"), compute: T.Buffer((16, 56, 56, 256), "int32")): + def main(p0: T.Tensor((16, 56, 56, 64), "int8"), p1: T.Tensor((256, 1, 1, 64), "int8"), p2: T.Tensor((1, 1, 1, 256), "int32"), p3: T.Tensor((1, 1, 1, 256), "int32"), p4: T.Tensor((256,), "int32"), p5: T.Tensor((256,), "int32"), p6: T.Tensor((256,), "int32"), p7: T.Tensor((), "int32"), p8: T.Tensor((1,), "int32"), p9: T.Tensor((16, 56, 56, 256), "int32"), compute: T.Tensor((16, 56, 56, 256), "int32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): Ts.reads() @@ -1289,7 +1289,7 @@ def test_compute_inline_softmax(): m = T.dynamic("m") @Ts.prim_func - def before(lv44: T.Buffer((T.int64(1), T.int64(32), n, m)), var_compute_intermediate: T.Buffer((T.int64(1), T.int64(32), n, m), 'float16')): + def before(lv44: T.Tensor((T.int64(1), T.int64(32), n, m)), var_compute_intermediate: T.Tensor((T.int64(1), T.int64(32), n, m), 'float16')): T.func_attr({"tirx.noalias": True}) T_softmax_maxelem = Ts.sblock_alloc_buffer((T.int64(1), T.int64(32), n)) @@ -1336,7 +1336,7 @@ def before(lv44: T.Buffer((T.int64(1), T.int64(32), n, m)), var_compute_intermed m = T.dynamic("m") @Ts.prim_func - def after(lv44: T.Buffer((T.int64(1), T.int64(32), n, m)), var_compute_intermediate: T.Buffer((T.int64(1), T.int64(32), n, m), 'float16')): + def after(lv44: T.Tensor((T.int64(1), T.int64(32), n, m)), var_compute_intermediate: T.Tensor((T.int64(1), T.int64(32), n, m), 'float16')): T.func_attr({"tirx.noalias": True}) # with Ts.sblock("root"): @@ -1384,7 +1384,7 @@ def test_reverse_compute_inline_layer_norm(): n = T.dynamic("n") @Ts.prim_func - def before(lv6: T.Buffer((T.int64(1), n, T.int64(2560))), weight1: T.Buffer((T.int64(2560),), "float32"), bias: T.Buffer((T.int64(2560),), "float32"), var_compute_intermediate: T.Buffer((T.int64(1), n, T.int64(2560)), 'float16')): + def before(lv6: T.Tensor((T.int64(1), n, T.int64(2560))), weight1: T.Tensor((T.int64(2560),), "float32"), bias: T.Tensor((T.int64(2560),), "float32"), var_compute_intermediate: T.Tensor((T.int64(1), n, T.int64(2560)), 'float16')): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) A_red_temp_v0_shared = Ts.sblock_alloc_buffer((T.int64(1), n), scope="shared") @@ -1425,7 +1425,7 @@ def before(lv6: T.Buffer((T.int64(1), n, T.int64(2560))), weight1: T.Buffer((T.i n = T.dynamic("n") @Ts.prim_func - def after(lv6: T.Buffer((T.int64(1), n, T.int64(2560))), weight1: T.Buffer((T.int64(2560),), "float32"), bias: T.Buffer((T.int64(2560),), "float32"), var_compute_intermediate: T.Buffer((T.int64(1), n, T.int64(2560)), 'float16')): + def after(lv6: T.Tensor((T.int64(1), n, T.int64(2560))), weight1: T.Tensor((T.int64(2560),), "float32"), bias: T.Tensor((T.int64(2560),), "float32"), var_compute_intermediate: T.Tensor((T.int64(1), n, T.int64(2560)), 'float16')): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): @@ -1466,8 +1466,8 @@ def after(lv6: T.Buffer((T.int64(1), n, T.int64(2560))), weight1: T.Buffer((T.in def test_reverse_compute_inline_slicing_then_cachewrite(): @Ts.prim_func def before( - x: T.Buffer((1, 16, 7, 7), "float32"), - T_strided_slice_with_axes: T.Buffer((1, 12, 7, 7), "float32"), + x: T.Tensor((1, 16, 7, 7), "float32"), + T_strided_slice_with_axes: T.Tensor((1, 12, 7, 7), "float32"), ): T_add = Ts.sblock_alloc_buffer((1, 16, 7, 7)) for ax0, ax1, ax2, ax3 in T.grid(1, 16, 7, 7): @@ -1483,8 +1483,8 @@ def before( @Ts.prim_func def after( - x: T.Buffer((1, 16, 7, 7), "float32"), - T_strided_slice_with_axes: T.Buffer((1, 12, 7, 7), "float32"), + x: T.Tensor((1, 16, 7, 7), "float32"), + T_strided_slice_with_axes: T.Tensor((1, 12, 7, 7), "float32"), ): T_strided_slice_with_axes_global = Ts.sblock_alloc_buffer((1, 12, 7, 7)) for ax0, ax1, ax2, ax3 in T.grid(1, 16, 7, 7): @@ -1510,9 +1510,9 @@ def after( def test_inline_with_reduction(): @Ts.prim_func def before( - T_softmax_norm: T.Buffer((T.int64(6), T.int64(1), T.int64(1)), "float32"), - T_reshape_2: T.Buffer((T.int64(6), T.int64(1), T.int64(64)), "float32"), - T_transpose: T.Buffer((T.int64(1), T.int64(1), T.int64(6), T.int64(64)), "float32"), + T_softmax_norm: T.Tensor((T.int64(6), T.int64(1), T.int64(1)), "float32"), + T_reshape_2: T.Tensor((T.int64(6), T.int64(1), T.int64(64)), "float32"), + T_transpose: T.Tensor((T.int64(1), T.int64(1), T.int64(6), T.int64(64)), "float32"), ): T_batch_matmul_NN = Ts.sblock_alloc_buffer((T.int64(6), T.int64(1), T.int64(64))) for ax0, ax1 in T.grid(T.int64(6), T.int64(64)): @@ -1537,9 +1537,9 @@ def before( @Ts.prim_func def after( - T_softmax_norm: T.Buffer((T.int64(6), T.int64(1), T.int64(1)), "float32"), - T_reshape_2: T.Buffer((T.int64(6), T.int64(1), T.int64(64)), "float32"), - T_transpose: T.Buffer((T.int64(1), T.int64(1), T.int64(6), T.int64(64)), "float32"), + T_softmax_norm: T.Tensor((T.int64(6), T.int64(1), T.int64(1)), "float32"), + T_reshape_2: T.Tensor((T.int64(6), T.int64(1), T.int64(64)), "float32"), + T_transpose: T.Tensor((T.int64(1), T.int64(1), T.int64(6), T.int64(64)), "float32"), ): for ax0, ax1 in T.grid(T.int64(6), T.int64(64)): with Ts.sblock("bmm"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_decompose_padding.py b/tests/python/s_tir/schedule/test_tir_schedule_decompose_padding.py index d51cb3d417d2..88c5d066f95d 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_decompose_padding.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_decompose_padding.py @@ -49,8 +49,8 @@ def check_decompose_padding(origin, scheduled, expected, check_run=False): def test_int64_indices_batch_decompose_padding(): @Ts.prim_func def before_decompose( - x: T.Buffer((T.int64(1), T.int64(128), T.int64(128)), "int32"), - y: T.Buffer((T.int64(1), T.int64(140), T.int64(128)), "int32"), + x: T.Tensor((T.int64(1), T.int64(128), T.int64(128)), "int32"), + y: T.Tensor((T.int64(1), T.int64(140), T.int64(128)), "int32"), ): for b, i, j in T.grid(T.int64(1), T.int64(140), T.int64(128)): with Ts.sblock("block"): @@ -59,8 +59,8 @@ def before_decompose( @Ts.prim_func def after_decompose( - x: T.Buffer((T.int64(1), T.int64(128), T.int64(128)), "int32"), - y: T.Buffer((T.int64(1), T.int64(140), T.int64(128)), "int32"), + x: T.Tensor((T.int64(1), T.int64(128), T.int64(128)), "int32"), + y: T.Tensor((T.int64(1), T.int64(140), T.int64(128)), "int32"), ): # with Ts.sblock("root"): for b, i in T.grid(T.int64(1), T.int64(140)): @@ -91,14 +91,14 @@ def after_decompose( def test_1d_decompose_padding(): @Ts.prim_func - def before_decompose(x: T.Buffer(128, "int32"), y: T.Buffer(140, "int32")): + def before_decompose(x: T.Tensor(128, "int32"), y: T.Tensor(140, "int32")): for i in range(140): with Ts.sblock("block"): vi = Ts.axis.remap("S", [i]) y[vi] = T.if_then_else(vi >= 6 and vi < 134, x[vi - 6], 0) @Ts.prim_func - def after_decompose(x: T.Buffer(128, "int32"), y: T.Buffer(140, "int32")): + def after_decompose(x: T.Tensor(128, "int32"), y: T.Tensor(140, "int32")): for i in T.serial(140): with Ts.sblock("block_pad_const"): vi = Ts.axis.spatial(140, i) @@ -120,7 +120,7 @@ def after_decompose(x: T.Buffer(128, "int32"), y: T.Buffer(140, "int32")): @Ts.prim_func def sum_pool_2d( - x: T.Buffer((1, 16, 225, 225), "int8"), tensor: T.Buffer((1, 16, 225, 225), "int8") + x: T.Tensor((1, 16, 225, 225), "int8"), tensor: T.Tensor((1, 16, 225, 225), "int8") ): pad_temp = Ts.sblock_alloc_buffer([1, 16, 231, 231], dtype="int8") for i0, i1, i2, i3 in T.grid(1, 16, 231, 231): @@ -146,7 +146,7 @@ def test_decompose_hw_padding_direct(): @Ts.prim_func def pooling_decompose_0( - x: T.Buffer((1, 16, 225, 225), "int8"), tensor: T.Buffer((1, 16, 225, 225), "int8") + x: T.Tensor((1, 16, 225, 225), "int8"), tensor: T.Tensor((1, 16, 225, 225), "int8") ): pad_temp = Ts.sblock_alloc_buffer([1, 16, 231, 231], dtype="int8") for i0, i1, i2, i3 in T.grid(1, 16, 231, 231): @@ -177,7 +177,7 @@ def test_decompose_hw_padding_tiled(): @Ts.prim_func def pooling_decompose_1( - x: T.Buffer((1, 16, 225, 225), "int8"), tensor: T.Buffer((1, 16, 225, 225), "int8") + x: T.Tensor((1, 16, 225, 225), "int8"), tensor: T.Tensor((1, 16, 225, 225), "int8") ) -> None: pad_temp = Ts.sblock_alloc_buffer([1, 16, 231, 231], dtype="int8") for i0, i2_0, i3_0 in T.grid(1, 3, 3): @@ -237,7 +237,7 @@ def test_decompose_hw_padding_tiled_and_lift_pad(): @Ts.prim_func def pooling_decompose_2( - x: T.Buffer((1, 16, 225, 225), "int8"), tensor: T.Buffer((1, 16, 225, 225), "int8") + x: T.Tensor((1, 16, 225, 225), "int8"), tensor: T.Tensor((1, 16, 225, 225), "int8") ) -> None: pad_temp = Ts.sblock_alloc_buffer([1, 16, 231, 231], dtype="int8") for i0, i2_0, i3_0, ax0, ax1, ax2 in T.grid(1, 3, 3, 16, 81, 81): @@ -297,7 +297,7 @@ def test_decompose_hw_padding_non_perfect_tiled(): @Ts.prim_func def pooling_decompose_3( - x: T.Buffer((1, 16, 225, 225), "int8"), tensor: T.Buffer((1, 16, 225, 225), "int8") + x: T.Tensor((1, 16, 225, 225), "int8"), tensor: T.Tensor((1, 16, 225, 225), "int8") ) -> None: pad_temp = Ts.sblock_alloc_buffer([1, 16, 231, 231], dtype="int8") for i0, i2_0, i3_0 in T.grid(1, 3, 3): @@ -361,8 +361,8 @@ def test_decompose_wrt_single_child_subtree(): @Ts.prim_func def pad_op( - x: T.Buffer((1, 16, 225, 225), "int8"), - y: T.Buffer((1, 16, 231, 231), dtype="int8"), + x: T.Tensor((1, 16, 225, 225), "int8"), + y: T.Tensor((1, 16, 231, 231), dtype="int8"), ): for i0, i1, i2, i3 in T.grid(1, 16, 231, 231): with Ts.sblock("pad_temp"): @@ -375,7 +375,7 @@ def pad_op( @Ts.prim_func def pad_op_after( - x: T.Buffer((1, 16, 225, 225), "int8"), y: T.Buffer((1, 16, 231, 231), "int8") + x: T.Tensor((1, 16, 225, 225), "int8"), y: T.Tensor((1, 16, 231, 231), "int8") ): for i0, i1 in T.grid(1, 16): for i2, i3 in T.grid(231, 231): @@ -401,7 +401,7 @@ def test_not_to_decompose_trivial_predicate(): @Ts.prim_func def trivial_pad( - x: T.Buffer((1, 16, 225, 225), "int8"), y: T.Buffer([1, 16, 225, 225], dtype="int8") + x: T.Tensor((1, 16, 225, 225), "int8"), y: T.Tensor([1, 16, 225, 225], dtype="int8") ): for i0, i1, i2, i3 in T.grid(1, 16, 225, 225): with Ts.sblock("pad_temp"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_error.py b/tests/python/s_tir/schedule/test_tir_schedule_error.py index 8c239bd8f2a1..f5437065d1da 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_error.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_error.py @@ -30,7 +30,7 @@ @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j in T.grid(128, 128): with Ts.sblock("init"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -43,8 +43,8 @@ def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 1 @Ts.prim_func def two_kernels( - A: T.Buffer((1, seq_len * 8), "int32"), # noqa: F821 - B: T.Buffer((1, seq_len * 8), "int32", align=8), # noqa: F821 + A: T.Tensor((1, seq_len * 8), "int32"), # noqa: F821 + B: T.Tensor((1, seq_len * 8), "int32", align=8), # noqa: F821 seq_len: T.int32, ): T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/s_tir/schedule/test_tir_schedule_for_kind.py b/tests/python/s_tir/schedule/test_tir_schedule_for_kind.py index e8e5a63504df..641afc396500 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_for_kind.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_for_kind.py @@ -34,7 +34,7 @@ @Ts.prim_func -def element_wise(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: +def element_wise(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -42,7 +42,7 @@ def element_wise(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: @Ts.prim_func -def element_wise_parallelized(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: +def element_wise_parallelized(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i0 in T.parallel(0, 128): for i1 in T.serial(0, 128): with Ts.sblock("B"): @@ -51,7 +51,7 @@ def element_wise_parallelized(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) @Ts.prim_func -def element_wise_i_bound(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: +def element_wise_i_bound(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i0 in T.thread_binding(0, 128, thread="threadIdx.x"): for i1 in T.serial(0, 128): with Ts.sblock("B"): @@ -60,7 +60,7 @@ def element_wise_i_bound(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> No @Ts.prim_func -def element_wise_compute_at_split(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def element_wise_compute_at_split(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i in T.serial(0, 128): for j0 in T.serial(0, 128): @@ -76,7 +76,7 @@ def element_wise_compute_at_split(A: T.Buffer((128, 128)), C: T.Buffer((128, 128 @Ts.prim_func def element_wise_compute_at_split_vectorized( - A: T.Buffer((128, 128)), C: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), C: T.Tensor((128, 128)) ) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i in T.serial(0, 128): @@ -93,7 +93,7 @@ def element_wise_compute_at_split_vectorized( @Ts.prim_func -def element_wise_split_predicate(A: T.Buffer([128, 128]), B: T.Buffer([128, 128])) -> None: +def element_wise_split_predicate(A: T.Tensor([128, 128]), B: T.Tensor([128, 128])) -> None: for i, j_0, j_1 in T.grid(128, 13, 10): with Ts.sblock("B"): Ts.where(j_0 * 10 + j_1 < 128) @@ -104,7 +104,7 @@ def element_wise_split_predicate(A: T.Buffer([128, 128]), B: T.Buffer([128, 128] @Ts.prim_func def element_wise_split_predicate_parallelized( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]) ) -> None: for i in T.serial(0, 128): for j_0 in T.parallel(0, 13): @@ -118,7 +118,7 @@ def element_wise_split_predicate_parallelized( @Ts.prim_func def element_wise_split_predicate_vectorized( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]) ) -> None: for i in T.vectorized(0, 128): for j_0, j_1 in T.grid(13, 10): @@ -131,7 +131,7 @@ def element_wise_split_predicate_vectorized( @Ts.prim_func def element_wise_compute_at_split_j0_j1o_bound( - A: T.Buffer((128, 128)), C: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), C: T.Tensor((128, 128)) ) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i in T.serial(0, 128): @@ -148,7 +148,7 @@ def element_wise_compute_at_split_j0_j1o_bound( @Ts.prim_func -def matmul(A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def matmul(A: T.Tensor((128, 128)), B: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("C"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) @@ -158,7 +158,7 @@ def matmul(A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 1 @Ts.prim_func -def rowsum(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for i, k in T.grid(128, 128): with Ts.sblock("B"): vi, vk = Ts.axis.remap("SR", [i, k]) @@ -168,7 +168,7 @@ def rowsum(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: @Ts.prim_func -def rowsum_unrolled(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum_unrolled(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for i0 in T.unroll(0, 128): for i1 in T.serial(0, 128): with Ts.sblock("B"): @@ -179,7 +179,7 @@ def rowsum_unrolled(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: @Ts.prim_func -def rowsum_not_quasi_affine(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum_not_quasi_affine(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for i, k in T.grid(128, 16): with Ts.sblock("B"): vi = Ts.axis.S(128, i) @@ -190,7 +190,7 @@ def rowsum_not_quasi_affine(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> Non @Ts.prim_func -def rowsum_not_compact_data_flow(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum_not_compact_data_flow(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for i, k in T.grid(128, 16): with Ts.sblock("B"): vi, vk = Ts.axis.remap("SR", [i, k]) @@ -200,7 +200,7 @@ def rowsum_not_compact_data_flow(A: T.Buffer((128, 128)), B: T.Buffer((128,))) - @Ts.prim_func -def rowsum_cross_thread_reduction(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum_cross_thread_reduction(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for i0 in T.serial(0, 128): for i1 in T.thread_binding(0, 128, thread="threadIdx.x"): with Ts.sblock("B"): @@ -211,7 +211,7 @@ def rowsum_cross_thread_reduction(A: T.Buffer((128, 128)), B: T.Buffer((128,))) @Ts.prim_func -def opaque_block(A: T.Buffer((16,))) -> None: +def opaque_block(A: T.Tensor((16,))) -> None: for i in T.serial(0, 15): with Ts.sblock("opaque"): A[i + 1] = A[i + 1] + A[i] @@ -219,7 +219,7 @@ def opaque_block(A: T.Buffer((16,))) -> None: @Ts.prim_func def block_inside_init( - A: T.Buffer([128, 128, 128], dtype="float32"), B: T.Buffer([128, 128], dtype="float32") + A: T.Tensor([128, 128, 128], dtype="float32"), B: T.Tensor([128, 128], dtype="float32") ) -> None: for i in T.serial(0, 128): with Ts.sblock("outer"): @@ -238,7 +238,7 @@ def block_inside_init( @Ts.prim_func def thread_bound_block_inside_init( - A: T.Buffer([128, 128, 128], dtype="float32"), B: T.Buffer([128, 128], dtype="float32") + A: T.Tensor([128, 128, 128], dtype="float32"), B: T.Tensor([128, 128], dtype="float32") ) -> None: for i in T.thread_binding(0, 128, thread="threadIdx.x"): with Ts.sblock("outer"): @@ -257,9 +257,9 @@ def thread_bound_block_inside_init( @Ts.prim_func def decomposed_gemm( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ): local = Ts.sblock_alloc_buffer((16, 16), "float32") for i, j in T.grid(4, 4): @@ -283,9 +283,9 @@ def decomposed_gemm( @Ts.prim_func def decomposed_gemm_after_vectorize( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ): local = Ts.sblock_alloc_buffer((16, 16), "float32") for i, j in T.grid(4, 4): @@ -310,7 +310,7 @@ def decomposed_gemm_after_vectorize( @Ts.prim_func def nested_block_bind( - A: T.Buffer((16, 16, 16, 16), "float32"), B: T.Buffer((16, 16, 16), "float32") + A: T.Tensor((16, 16, 16, 16), "float32"), B: T.Tensor((16, 16, 16), "float32") ): for i, j in T.grid(16, 16): with Ts.sblock("outer"): @@ -325,7 +325,7 @@ def nested_block_bind( @Ts.prim_func def thread_bound_nested_block( - A: T.Buffer((16, 16, 16, 16), "float32"), B: T.Buffer((16, 16, 16), "float32") + A: T.Tensor((16, 16, 16, 16), "float32"), B: T.Tensor((16, 16, 16), "float32") ) -> None: for i in T.serial(16): for j in T.thread_binding(16, thread="blockIdx.x"): @@ -342,7 +342,7 @@ def thread_bound_nested_block( @Ts.prim_func def nested_block_bind_after_cache_read( - A: T.Buffer((16, 16), "float32"), B: T.Buffer((16,), "float32") + A: T.Tensor((16, 16), "float32"), B: T.Tensor((16,), "float32") ) -> None: for i in T.serial(16): with Ts.sblock("outer"): @@ -363,7 +363,7 @@ def nested_block_bind_after_cache_read( @Ts.prim_func def thread_bound_nested_block_after_cache_read( - A: T.Buffer((16, 16), "float32"), B: T.Buffer((16,), "float32") + A: T.Tensor((16, 16), "float32"), B: T.Tensor((16,), "float32") ) -> None: for i in T.thread_binding(16, thread="blockIdx.x"): with Ts.sblock("outer"): @@ -384,9 +384,9 @@ def thread_bound_nested_block_after_cache_read( @Ts.prim_func def decomposed_gemm_parallelize_init( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: local = Ts.sblock_alloc_buffer([16, 16], dtype="float32") for i, j in T.grid(4, 4): @@ -416,7 +416,7 @@ def decomposed_gemm_parallelize_init( @Ts.prim_func -def scatter_compute(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): +def scatter_compute(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): for i in T.grid(8): with Ts.sblock("first_half"): vi = Ts.axis.spatial(16, 8 + i) @@ -430,7 +430,7 @@ def scatter_compute(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32") @Ts.prim_func def scatter_compute_parallelize( - A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32") + A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32") ) -> None: # body # with Ts.sblock("root") @@ -646,8 +646,8 @@ def test_scatter_parallelize(): def test_bind_thread_iter_var_dtype(): @Ts.prim_func(private=True) def before( - A: T.Buffer((T.int64(128), T.int64(128))), - B: T.Buffer((T.int64(128), T.int64(128))), + A: T.Tensor((T.int64(128), T.int64(128))), + B: T.Tensor((T.int64(128), T.int64(128))), ) -> None: for i, j in T.grid(T.int64(128), T.int64(128)): with Ts.sblock("B"): @@ -656,8 +656,8 @@ def before( @Ts.prim_func(private=True) def expected( - A: T.Buffer((T.int64(128), T.int64(128))), - B: T.Buffer((T.int64(128), T.int64(128))), + A: T.Tensor((T.int64(128), T.int64(128))), + B: T.Tensor((T.int64(128), T.int64(128))), ) -> None: for i0 in T.thread_binding(T.int64(128), thread="threadIdx.x"): # Use T.serial with explicit int64 min so the inner sblock iter_var dom diff --git a/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue.py b/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue.py index ec4a62c12242..16468b6a07bd 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue.py @@ -34,10 +34,10 @@ @Ts.prim_func def matmul_bias_before( - A: T.Buffer((16, 16), "int8"), - B: T.Buffer((16, 16), "int8"), - C: T.Buffer((16, 16), "int32"), - D: T.Buffer((16, 16), "int32"), + A: T.Tensor((16, 16), "int8"), + B: T.Tensor((16, 16), "int8"), + C: T.Tensor((16, 16), "int32"), + D: T.Tensor((16, 16), "int32"), ) -> None: temp = Ts.sblock_alloc_buffer((16, 16), dtype="int32") for i, j, k in T.grid(16, 16, 16): @@ -54,10 +54,10 @@ def matmul_bias_before( @Ts.prim_func def matmul_bias_expected( - A: T.Buffer((16, 16), "int8"), - B: T.Buffer((16, 16), "int8"), - C: T.Buffer((16, 16), "int32"), - D: T.Buffer((16, 16), "int32"), + A: T.Tensor((16, 16), "int8"), + B: T.Tensor((16, 16), "int8"), + C: T.Tensor((16, 16), "int32"), + D: T.Tensor((16, 16), "int32"), ) -> None: temp = Ts.sblock_alloc_buffer((16, 16), dtype="int32") for i, j, k in T.grid(16, 16, 16): @@ -72,10 +72,10 @@ def matmul_bias_expected( @Ts.prim_func def matmul_bias_fp32_before( - A: T.Buffer((32, 32), "float32"), - B: T.Buffer((32, 32), "float32"), - C: T.Buffer((32, 32), "float32"), - D: T.Buffer((32, 32), "float32"), + A: T.Tensor((32, 32), "float32"), + B: T.Tensor((32, 32), "float32"), + C: T.Tensor((32, 32), "float32"), + D: T.Tensor((32, 32), "float32"), ) -> None: temp = Ts.sblock_alloc_buffer((32, 32), dtype="float32") for i, j, k in T.grid(32, 32, 32): @@ -92,10 +92,10 @@ def matmul_bias_fp32_before( @Ts.prim_func def matmul_bias_fp32_expected( - A: T.Buffer((32, 32), "float32"), - B: T.Buffer((32, 32), "float32"), - C: T.Buffer((32, 32), "float32"), - D: T.Buffer((32, 32), "float32"), + A: T.Tensor((32, 32), "float32"), + B: T.Tensor((32, 32), "float32"), + C: T.Tensor((32, 32), "float32"), + D: T.Tensor((32, 32), "float32"), ) -> None: temp = Ts.sblock_alloc_buffer((32, 32), dtype="float32") for i, j, k in T.grid(32, 32, 32): @@ -110,11 +110,11 @@ def matmul_bias_fp32_expected( @Ts.prim_func def matmul_bias_multiple_epilogue_before( - A: T.Buffer((16, 16), "int8"), - B: T.Buffer((16, 16), "int8"), - C: T.Buffer((16, 16), "int32"), - D: T.Buffer((16, 16), "int32"), - E: T.Buffer((16, 16), "int32"), + A: T.Tensor((16, 16), "int8"), + B: T.Tensor((16, 16), "int8"), + C: T.Tensor((16, 16), "int32"), + D: T.Tensor((16, 16), "int32"), + E: T.Tensor((16, 16), "int32"), ) -> None: temp = Ts.sblock_alloc_buffer((16, 16), dtype="int32") for i, j, k in T.grid(16, 16, 16): @@ -135,11 +135,11 @@ def matmul_bias_multiple_epilogue_before( @Ts.prim_func def matmul_bias_multiple_epilogue_expected( - A: T.Buffer((16, 16), "int8"), - B: T.Buffer((16, 16), "int8"), - C: T.Buffer((16, 16), "int32"), - D: T.Buffer((16, 16), "int32"), - E: T.Buffer((16, 16), "int32"), + A: T.Tensor((16, 16), "int8"), + B: T.Tensor((16, 16), "int8"), + C: T.Tensor((16, 16), "int32"), + D: T.Tensor((16, 16), "int32"), + E: T.Tensor((16, 16), "int32"), ) -> None: temp = Ts.sblock_alloc_buffer((16, 16), dtype="int32") for i, j, k in T.grid(16, 16, 16): @@ -219,11 +219,11 @@ def test_fuse_reduction_epilogue_multiple_epilogue(): @Ts.prim_func def matmul_bias_invalid_multiple_use_before( - A: T.Buffer((16, 16), "int8"), - B: T.Buffer((16, 16), "int8"), - C1: T.Buffer((16, 16), "int32"), - C2: T.Buffer((16, 16), "int32"), - D: T.Buffer((16, 16), "int32"), + A: T.Tensor((16, 16), "int8"), + B: T.Tensor((16, 16), "int8"), + C1: T.Tensor((16, 16), "int32"), + C2: T.Tensor((16, 16), "int32"), + D: T.Tensor((16, 16), "int32"), ) -> None: """Epilogue uses the reduction result twice; fusion must be rejected.""" temp = Ts.sblock_alloc_buffer((16, 16), dtype="int32") @@ -249,10 +249,10 @@ def test_fuse_reduction_epilogue_reject_multiple_use(): @Ts.prim_func def matmul_bias_invalid_scaling_before( - A: T.Buffer((16, 16), "int8"), - B: T.Buffer((16, 16), "int8"), - C: T.Buffer((16, 16), "int32"), - D: T.Buffer((16, 16), "int32"), + A: T.Tensor((16, 16), "int8"), + B: T.Tensor((16, 16), "int8"), + C: T.Tensor((16, 16), "int32"), + D: T.Tensor((16, 16), "int32"), ) -> None: """Epilogue scales the reduction result; fusion must be rejected.""" temp = Ts.sblock_alloc_buffer((16, 16), dtype="int32") diff --git a/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue_clipping.py b/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue_clipping.py index 2e5f319b730b..f73fcd9d10d0 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue_clipping.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue_clipping.py @@ -34,9 +34,9 @@ @Ts.prim_func def matmul_clipping_before( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), lower: T.float32, upper: T.float32, ) -> None: @@ -57,9 +57,9 @@ def matmul_clipping_before( @Ts.prim_func def matmul_clipping_expected( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), lower: T.float32, upper: T.float32, ) -> None: @@ -85,9 +85,9 @@ def test_matmul_clipping(): @Ts.prim_func def matmul_clipping_before_per_iteration( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), ) -> None: """Original function with per-iteration clipping (same semantics as fused).""" temp = Ts.sblock_alloc_buffer((16, 16), dtype="float32") @@ -156,10 +156,10 @@ def test_matmul_clipping_correctness_unified(): @Ts.prim_func def matmul_clipping_multiple_epilogue_before( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), - E: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), + E: T.Tensor((16, 16), "float32"), lower: T.float32, upper: T.float32, ) -> None: @@ -185,10 +185,10 @@ def matmul_clipping_multiple_epilogue_before( @Ts.prim_func def matmul_clipping_multiple_epilogue_expected( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), - E: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), + E: T.Tensor((16, 16), "float32"), lower: T.float32, upper: T.float32, ) -> None: @@ -247,9 +247,9 @@ def test_matmul_clipping_commutative_variants(pattern_func): @Ts.prim_func def test_func( - A: T.Buffer((8, 8), "float32"), - B: T.Buffer((8, 8), "float32"), - D: T.Buffer((8, 8), "float32"), + A: T.Tensor((8, 8), "float32"), + B: T.Tensor((8, 8), "float32"), + D: T.Tensor((8, 8), "float32"), ) -> None: temp = Ts.sblock_alloc_buffer((8, 8), dtype="float32") for i, j, k in T.grid(8, 8, 8): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue_relu.py b/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue_relu.py index a25863dd2ec6..1d42f23a1e91 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue_relu.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_fuse_reduction_epilogue_relu.py @@ -34,10 +34,10 @@ @Ts.prim_func def matmul_bias_relu_before( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), ) -> None: """Original function with separate reduction and epilogue blocks (Bias + ReLU).""" temp = Ts.sblock_alloc_buffer((16, 16), dtype="float32") @@ -56,10 +56,10 @@ def matmul_bias_relu_before( @Ts.prim_func def matmul_bias_relu_before_per_iteration( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), ) -> None: """Original function with per-iteration ReLU (same semantics as fused).""" temp = Ts.sblock_alloc_buffer((16, 16), dtype="float32") @@ -82,10 +82,10 @@ def matmul_bias_relu_before_per_iteration( @Ts.prim_func def matmul_bias_relu_expected( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), ) -> None: """Expected function after fusion (Bias + ReLU).""" temp = Ts.sblock_alloc_buffer((16, 16), dtype="float32") @@ -157,11 +157,11 @@ def test_matmul_bias_relu_correctness_unified(): @Ts.prim_func def matmul_bias_relu_multiple_epilogue_before( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), - E: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), + E: T.Tensor((16, 16), "float32"), ) -> None: """Original function with separate reduction and multiple epilogue blocks (one with ReLU, one without).""" temp = Ts.sblock_alloc_buffer((16, 16), dtype="float32") @@ -185,11 +185,11 @@ def matmul_bias_relu_multiple_epilogue_before( @Ts.prim_func def matmul_bias_relu_multiple_epilogue_expected( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), - D: T.Buffer((16, 16), "float32"), - E: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), + D: T.Tensor((16, 16), "float32"), + E: T.Tensor((16, 16), "float32"), ) -> None: """Expected function after fusion (Bias + ReLU) with multiple epilogue blocks.""" temp = Ts.sblock_alloc_buffer((16, 16), dtype="float32") diff --git a/tests/python/s_tir/schedule/test_tir_schedule_merge.py b/tests/python/s_tir/schedule/test_tir_schedule_merge.py index e7309f7d81e2..f764c9d05fee 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_merge.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_merge.py @@ -32,7 +32,7 @@ @Ts.prim_func -def elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128)), D: T.Buffer((64, 64))) -> None: +def elementwise(A: T.Tensor((128, 128)), C: T.Tensor((128, 128)), D: T.Tensor((64, 64))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -58,7 +58,7 @@ def elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128)), D: T.Buffer((6 @Ts.prim_func def elementwise_merged( - A: T.Buffer((128, 128)), C: T.Buffer((128, 128)), D: T.Buffer((64, 64)) + A: T.Tensor((128, 128)), C: T.Tensor((128, 128)), D: T.Tensor((64, 64)) ) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -86,7 +86,7 @@ def elementwise_merged( @Ts.prim_func def elementwise_merged2( - A: T.Buffer((128, 128)), C: T.Buffer((128, 128)), D: T.Buffer((64, 64)) + A: T.Tensor((128, 128)), C: T.Tensor((128, 128)), D: T.Tensor((64, 64)) ) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -136,7 +136,7 @@ def test_merge2(): def test_merge_fail_not_only_child(): @Ts.prim_func - def elementwise_with_seq(A: T.Buffer((128, 128, 128)), C: T.Buffer((128, 128, 128))) -> None: + def elementwise_with_seq(A: T.Tensor((128, 128, 128)), C: T.Tensor((128, 128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128, 128)) D = Ts.sblock_alloc_buffer((128, 128, 128)) for i, j in T.grid(128, 128): @@ -166,7 +166,7 @@ def elementwise_with_seq(A: T.Buffer((128, 128, 128)), C: T.Buffer((128, 128, 12 def test_merge_fail_not_start_with_zero(): @Ts.prim_func def elementwise_loops_not_start_with_zero( - A: T.Buffer((128, 128, 128)), C: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), C: T.Tensor((128, 128, 128)) ) -> None: B = Ts.sblock_alloc_buffer((128, 128, 128)) for i, j in T.grid(128, 128): @@ -192,7 +192,7 @@ def elementwise_loops_not_start_with_zero( def test_merge_fail_not_same_extent(): @Ts.prim_func def elementwise_loops_not_same_extent( - A: T.Buffer((128, 128, 128)), C: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), C: T.Tensor((128, 128, 128)) ) -> None: B = Ts.sblock_alloc_buffer((64, 128, 128)) for i, j in T.grid(64, 128): @@ -218,7 +218,7 @@ def elementwise_loops_not_same_extent( def test_merge_fail_not_same_level(): @Ts.prim_func def elementwise_not_same_level( - A: T.Buffer((128, 128, 128)), C: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), C: T.Tensor((128, 128, 128)) ) -> None: B = Ts.sblock_alloc_buffer((128, 128, 128)) for i, j in T.grid(128, 128): @@ -244,7 +244,7 @@ def elementwise_not_same_level( def test_merge_fail_with_different_scope(): @Ts.prim_func def elementwise_with_different_scope( - A: T.Buffer((128, 128, 128)), C: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), C: T.Tensor((128, 128, 128)) ) -> None: B = Ts.sblock_alloc_buffer((128, 128, 128)) with Ts.sblock("A"): 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 e99efe06cb90..80d3970570a5 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 @@ -33,9 +33,9 @@ @Ts.prim_func def matmul_before( - A: T.Buffer((128, 127), "float32"), - B: T.Buffer((127, 127), "float32"), - C: T.Buffer((128, 127), "float32"), + A: T.Tensor((128, 127), "float32"), + B: T.Tensor((127, 127), "float32"), + C: T.Tensor((128, 127), "float32"), ) -> None: A_shared = Ts.sblock_alloc_buffer((128, 127), "float32", scope="shared") B_shared = Ts.sblock_alloc_buffer((127, 127), "float32", scope="shared") @@ -62,9 +62,9 @@ def matmul_before( @Ts.prim_func def matmul_expected( - A: T.Buffer((128, 127), "float32"), - B: T.Buffer((127, 127), "float32"), - C: T.Buffer((128, 127), "float32"), + A: T.Tensor((128, 127), "float32"), + B: T.Tensor((127, 127), "float32"), + C: T.Tensor((128, 127), "float32"), ) -> None: A_shared_padded = Ts.sblock_alloc_buffer([128, 128], dtype="float32", scope="shared") B_shared_padded = Ts.sblock_alloc_buffer([128, 128], dtype="float32", scope="shared") @@ -109,9 +109,9 @@ def test_pad_matmul(): @Ts.prim_func def matmul_before( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((n, 128), "float32"), - C: T.Buffer((128, n), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((n, 128), "float32"), + C: T.Tensor((128, n), "float32"), ) -> None: for i0, i1, i2 in T.grid(128, n, 128): with Ts.sblock("C"): @@ -124,9 +124,9 @@ def matmul_before( @Ts.prim_func def matmul_after( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((n, 128), "float32"), - C: T.Buffer((128, n), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((n, 128), "float32"), + C: T.Tensor((128, n), "float32"), ): B_pad = Ts.sblock_alloc_buffer(((n + 31) // 32 * 32, 128)) C_pad = Ts.sblock_alloc_buffer((128, (n + 31) // 32 * 32)) @@ -159,10 +159,10 @@ def test_pad_matmul_2(): @Ts.prim_func def before( - A: T.Buffer((1, n, 4096)), - B: T.Buffer((11008, 4096)), - M: T.Buffer((1, n, 11008)), - D: T.Buffer((1, n, 11008)), + A: T.Tensor((1, n, 4096)), + B: T.Tensor((11008, 4096)), + M: T.Tensor((1, n, 11008)), + D: T.Tensor((1, n, 11008)), ): T.func_attr({"tirx.noalias": True}) @@ -184,10 +184,10 @@ def before( @Ts.prim_func def after( - A: T.Buffer((1, n, 4096)), - B: T.Buffer((11008, 4096)), - M: T.Buffer((1, n, 11008)), - D: T.Buffer((1, n, 11008)), + A: T.Tensor((1, n, 4096)), + B: T.Tensor((11008, 4096)), + M: T.Tensor((1, n, 11008)), + D: T.Tensor((1, n, 11008)), ): T.func_attr({"tirx.noalias": True}) @@ -230,9 +230,9 @@ def test_pad_rms(): @Ts.prim_func def before( - A: T.Buffer((1, n, 4096)), - W: T.Buffer((4096,), "float32"), - Result: T.Buffer((1, n, 4096), "float32"), + A: T.Tensor((1, n, 4096)), + W: T.Tensor((4096,), "float32"), + Result: T.Tensor((1, n, 4096), "float32"), ): T.func_attr({"tirx.noalias": True}) @@ -257,7 +257,7 @@ def before( @Ts.prim_func def after( - A: T.Buffer((1, n, 4096)), W: T.Buffer((4096,), "float32"), Result: T.Buffer((1, n, 4096)) + A: T.Tensor((1, n, 4096)), W: T.Tensor((4096,), "float32"), Result: T.Tensor((1, n, 4096)) ): T.func_attr({"tirx.noalias": True}) diff --git a/tests/python/s_tir/schedule/test_tir_schedule_partition.py b/tests/python/s_tir/schedule/test_tir_schedule_partition.py index a126eaf33fe8..215aa455fffb 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_partition.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_partition.py @@ -35,7 +35,7 @@ @Ts.prim_func -def elementwise(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128))) -> None: +def elementwise(A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128))) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("B"): vi, vj, vk = Ts.axis.remap("SSS", [i, j, k]) @@ -44,8 +44,8 @@ def elementwise(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128))) -> N @Ts.prim_func def elementwise_symbolic( - A: T.Buffer((128, 128, n)), # noqa: F821 - B: T.Buffer((128, 128, n)), # noqa: F821 + A: T.Tensor((128, 128, n)), # noqa: F821 + B: T.Tensor((128, 128, n)), # noqa: F821 n: T.int32, ) -> None: for i, j, k in T.grid(128, 128, n): @@ -55,7 +55,7 @@ def elementwise_symbolic( @Ts.prim_func -def elementwise_with_anno(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128))) -> None: +def elementwise_with_anno(A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128))) -> None: for i, j in T.grid(128, 128): for k in T.serial(0, 128, annotations={"useless_annotation": True}): with Ts.sblock("B"): @@ -67,7 +67,7 @@ def elementwise_with_anno(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 1 @Ts.prim_func def elementwise_with_thread_binding( - A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128)) ) -> None: for i, j in T.grid(128, 128): for k in T.thread_binding(0, 128, thread="threadIdx.x"): @@ -80,7 +80,7 @@ def elementwise_with_thread_binding( @Ts.prim_func def elementwise_with_opaque_block( - A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128)) ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("opaque"): @@ -95,7 +95,7 @@ def elementwise_with_opaque_block( @Ts.prim_func def elementwise_partition_with_opaque_block( - A: T.Buffer([128, 128, 128]), B: T.Buffer([128, 128, 128]) + A: T.Tensor([128, 128, 128]), B: T.Tensor([128, 128, 128]) ) -> None: with Ts.sblock("root"): Ts.reads() @@ -132,7 +132,7 @@ def elementwise_partition_with_opaque_block( @Ts.prim_func def elementwise_loop_partition_case0( - A: T.Buffer([128, 128, 128]), B: T.Buffer([128, 128, 128]) + A: T.Tensor([128, 128, 128]), B: T.Tensor([128, 128, 128]) ) -> None: with Ts.sblock("root"): Ts.reads() @@ -210,7 +210,7 @@ def elementwise_loop_partition_case0( @Ts.prim_func def elementwise_loop_partition_case1( - A: T.Buffer([128, 128, 128]), B: T.Buffer([128, 128, 128]) + A: T.Tensor([128, 128, 128]), B: T.Tensor([128, 128, 128]) ) -> None: with Ts.sblock("root"): Ts.reads() @@ -275,7 +275,7 @@ def elementwise_loop_partition_case1( @Ts.prim_func -def opaque_access(A: T.Buffer([16, 16], "float32"), B: T.Buffer([16, 16], "float32")) -> None: +def opaque_access(A: T.Tensor([16, 16], "float32"), B: T.Tensor([16, 16], "float32")) -> None: for i, j in T.grid(16, 16): with Ts.sblock("A"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -291,7 +291,7 @@ def opaque_access(A: T.Buffer([16, 16], "float32"), B: T.Buffer([16, 16], "float @Ts.prim_func -def opaque_access_loop_partition(A: T.Buffer((16, 16)), B: T.Buffer((16, 16))) -> None: +def opaque_access_loop_partition(A: T.Tensor((16, 16)), B: T.Tensor((16, 16))) -> None: for i in range(16): with Ts.sblock("A_j_common"): Ts.reads() diff --git a/tests/python/s_tir/schedule/test_tir_schedule_read_write_at.py b/tests/python/s_tir/schedule/test_tir_schedule_read_write_at.py index 9cd9bd4b20e9..7303ce312db9 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_read_write_at.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_read_write_at.py @@ -50,7 +50,7 @@ # pylint: disable=no-member,invalid-name,unused-variable,line-too-long,redefined-outer-name,unexpected-keyword-arg,too-many-nested-blocks,not-callable @Ts.prim_func -def cuda_matmul(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048], 'float32'), C: T.Buffer([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable +def cuda_matmul(A: T.Tensor([2048, 2048], 'float32'), B: T.Tensor([2048, 2048], 'float32'), C: T.Tensor([2048, 2048], 'float32')) -> None: # pylint: disable=undefined-loop-variable for by in T.thread_binding(0, 32, thread = "blockIdx.y"): for bx in T.thread_binding(0, 32, thread = "blockIdx.x"): @@ -72,7 +72,7 @@ def cuda_matmul(A: T.Buffer([2048, 2048], 'float32'), B: T.Buffer([2048, 2048], C[vi, vj] = C[vi, vj] + A[vi, vk] * B[vk, vj] @Ts.prim_func -def cuda_matmul_read_at_a(A: T.Buffer([2048, 2048], dtype='float32'), B: T.Buffer([2048, 2048], dtype='float32'), C: T.Buffer([2048, 2048], dtype='float32')) -> None: +def cuda_matmul_read_at_a(A: T.Tensor([2048, 2048], dtype='float32'), B: T.Tensor([2048, 2048], dtype='float32'), C: T.Tensor([2048, 2048], dtype='float32')) -> None: A_shared = Ts.sblock_alloc_buffer([2048, 2048], dtype="float32", scope="shared") for by in T.thread_binding(0, 32, thread="blockIdx.y"): @@ -103,7 +103,7 @@ def cuda_matmul_read_at_a(A: T.Buffer([2048, 2048], dtype='float32'), B: T.Buffe C[vi, vj] = C[vi, vj] + A_shared[vi, vk] * B[vk, vj] @Ts.prim_func -def cuda_matmul_read_at_ab(A: T.Buffer([2048, 2048], dtype='float32'), B: T.Buffer([2048, 2048], dtype='float32'), C: T.Buffer([2048, 2048], dtype='float32')) -> None: +def cuda_matmul_read_at_ab(A: T.Tensor([2048, 2048], dtype='float32'), B: T.Tensor([2048, 2048], dtype='float32'), C: T.Tensor([2048, 2048], dtype='float32')) -> None: A_shared = Ts.sblock_alloc_buffer([2048, 2048], dtype="float32", scope="shared") B_shared = Ts.sblock_alloc_buffer([2048, 2048], dtype="float32", scope="shared") @@ -143,7 +143,7 @@ def cuda_matmul_read_at_ab(A: T.Buffer([2048, 2048], dtype='float32'), B: T.Buff C[vi, vj] = C[vi, vj] + A_shared[vi, vk] * B_shared[vk, vj] @Ts.prim_func -def cuda_matmul_write_at_c(A: T.Buffer([2048, 2048], dtype='float32'), B: T.Buffer([2048, 2048], dtype='float32'), C: T.Buffer([2048, 2048], dtype='float32')) -> None: +def cuda_matmul_write_at_c(A: T.Tensor([2048, 2048], dtype='float32'), B: T.Tensor([2048, 2048], dtype='float32'), C: T.Tensor([2048, 2048], dtype='float32')) -> None: A_shared = Ts.sblock_alloc_buffer([2048, 2048], dtype="float32", scope="shared") B_shared = Ts.sblock_alloc_buffer([2048, 2048], dtype="float32", scope="shared") diff --git a/tests/python/s_tir/schedule/test_tir_schedule_reduction.py b/tests/python/s_tir/schedule/test_tir_schedule_reduction.py index 72642ae05248..35974310a9fc 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_reduction.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_reduction.py @@ -36,7 +36,7 @@ @Ts.prim_func -def rowsum_blockized(A: T.Buffer([32, 4, 128]), B: T.Buffer([32, 4])) -> None: +def rowsum_blockized(A: T.Tensor([32, 4, 128]), B: T.Tensor([32, 4])) -> None: for i0, i2_0 in T.grid(32, 16): with Ts.sblock("blockized_B"): io, ko = Ts.axis.remap("SR", [i0, i2_0]) @@ -53,7 +53,7 @@ def rowsum_blockized(A: T.Buffer([32, 4, 128]), B: T.Buffer([32, 4])) -> None: @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) @@ -64,7 +64,7 @@ def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 1 @Ts.prim_func def matmul_decompose0( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("init"): @@ -79,8 +79,8 @@ def matmul_decompose0( @Ts.prim_func def matmul_decompose1( - A: T.Buffer([32, 4, 128], elem_offset=0, align=64, offset_factor=1), - B: T.Buffer([32, 4], elem_offset=0, align=64, offset_factor=1), + A: T.Tensor([32, 4, 128], elem_offset=0, align=64, offset_factor=1), + B: T.Tensor([32, 4], elem_offset=0, align=64, offset_factor=1), ) -> None: for i0 in T.serial(0, 32): with Ts.sblock("blockized_B_init"): @@ -101,9 +101,9 @@ def matmul_decompose1( @Ts.prim_func def matmul_decompose2( - A: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), - B: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), - C: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), + A: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), + B: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), + C: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), ) -> None: for i0, i1 in T.grid(128, 128): with Ts.sblock("update_init"): @@ -117,7 +117,7 @@ def matmul_decompose2( @Ts.prim_func def matmul_decompose_fail3( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i, k, j in T.grid(128, 128, 128): with Ts.sblock("update"): @@ -129,9 +129,9 @@ def matmul_decompose_fail3( @Ts.prim_func def matmul_decompose4( - A: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), - B: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), - C: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), + A: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), + B: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), + C: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), ) -> None: # body with Ts.sblock("root"): @@ -154,7 +154,7 @@ def matmul_decompose4( @Ts.prim_func def matmul_with_annotation( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): @@ -167,7 +167,7 @@ def matmul_with_annotation( @Ts.prim_func def matmul_decompose_with_annotation( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("init"): @@ -184,7 +184,7 @@ def matmul_decompose_with_annotation( @Ts.prim_func def colsum_with_vectorization( - A: T.Buffer([128, 32], dtype="float32"), B: T.Buffer([32], dtype="float32") + A: T.Tensor([128, 32], dtype="float32"), B: T.Tensor([32], dtype="float32") ) -> None: for k in T.serial(0, 128): for i in T.vectorized(0, 32): @@ -197,7 +197,7 @@ def colsum_with_vectorization( @Ts.prim_func def colsum_decompose_with_vectorization( - A: T.Buffer([128, 32], dtype="float32"), B: T.Buffer([32], dtype="float32") + A: T.Tensor([128, 32], dtype="float32"), B: T.Tensor([32], dtype="float32") ) -> None: for i in T.vectorized(0, 32): with Ts.sblock("B_init"): @@ -295,7 +295,7 @@ def test_decompose_reduction_ref_hash_check(): def test_decompose_reduction_nested_block(): @Ts.prim_func - def nested_block(A: T.Buffer((1, 64), "float32"), B: T.Buffer((1,), "float32")): + def nested_block(A: T.Tensor((1, 64), "float32"), B: T.Tensor((1,), "float32")): for i, ko in T.grid(1, 2): with Ts.sblock("outer"): vi, vko = Ts.axis.remap("SR", [i, ko]) @@ -312,7 +312,7 @@ def nested_block(A: T.Buffer((1, 64), "float32"), B: T.Buffer((1,), "float32")): B[vi] += C[vki] @Ts.prim_func - def decomposed_nested_block(A: T.Buffer((1, 64), "float32"), B: T.Buffer((1,), "float32")): + def decomposed_nested_block(A: T.Tensor((1, 64), "float32"), B: T.Tensor((1,), "float32")): for i in range(1): with Ts.sblock("outer_init"): vi = Ts.axis.spatial(1, i) @@ -351,7 +351,7 @@ def test_decompose_reduction_with_thread_binding(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((32, 16), "float32"), B: T.Buffer((32,), "float32")): + def main(A: T.Tensor((32, 16), "float32"), B: T.Tensor((32,), "float32")): for t in T.thread_binding(0, 32, thread="threadIdx.x"): for r in T.serial(16): with Ts.sblock("B"): @@ -363,7 +363,7 @@ def main(A: T.Buffer((32, 16), "float32"), B: T.Buffer((32,), "float32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((32, 16), "float32"), B: T.Buffer((32,), "float32")): + def main(A: T.Tensor((32, 16), "float32"), B: T.Tensor((32,), "float32")): for t_init in T.thread_binding(0, 32, thread="threadIdx.x"): with Ts.sblock("B_init"): vi = Ts.axis.remap("S", [t_init]) @@ -385,7 +385,7 @@ def test_decompose_reduction_preserves_general_spatial_predicates(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((8, 8), "float32"), B: T.Buffer((8,), "float32")): + def main(A: T.Tensor((8, 8), "float32"), B: T.Tensor((8,), "float32")): for i, k in T.grid(10, 10): with Ts.sblock("B"): Ts.where(1 <= i and i < 9 and 1 <= k and k < 9) @@ -398,7 +398,7 @@ def main(A: T.Buffer((8, 8), "float32"), B: T.Buffer((8,), "float32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((8, 8), "float32"), B: T.Buffer((8,), "float32")): + def main(A: T.Tensor((8, 8), "float32"), B: T.Tensor((8,), "float32")): for i_init in range(10): with Ts.sblock("B_init"): Ts.where(1 <= i_init and i_init < 9) @@ -421,7 +421,7 @@ def test_decompose_reduction_drops_mixed_rfactor_bound(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((20,), "float32"), B: T.Buffer((), "float32")): + def main(A: T.Tensor((20,), "float32"), B: T.Tensor((), "float32")): for k in range(20): with Ts.sblock("B"): vk = Ts.axis.reduce(20, k) @@ -432,7 +432,7 @@ def main(A: T.Buffer((20,), "float32"), B: T.Buffer((), "float32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((20,), "float32"), B: T.Buffer((), "float32")): + def main(A: T.Tensor((20,), "float32"), B: T.Tensor((), "float32")): B_rf = Ts.sblock_alloc_buffer((16,), elem_offset=T.int64(0)) for k_1_init in range(16): with Ts.sblock("B_rf_init"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_reindex.py b/tests/python/s_tir/schedule/test_tir_schedule_reindex.py index 05042b40bd8e..162c41b790ee 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_reindex.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_reindex.py @@ -32,7 +32,7 @@ @Ts.prim_func def transpose_elementwise( - A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32") ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -42,7 +42,7 @@ def transpose_elementwise( @Ts.prim_func def transpose_elementwise_reindex_read( - A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32") ) -> None: A_reindex = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): @@ -57,9 +57,9 @@ def transpose_elementwise_reindex_read( @Ts.prim_func def conv2d_nhwc( - Input: T.Buffer((1, 224, 224, 3), "float32"), - Weight: T.Buffer((7, 7, 3, 64), "float32"), - Conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32"), + Input: T.Tensor((1, 224, 224, 3), "float32"), + Weight: T.Tensor((7, 7, 3, 64), "float32"), + Conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32"), ) -> None: PadInput = Ts.sblock_alloc_buffer([1, 230, 230, 3], dtype="float32") for i0, i1, i2, i3 in T.grid(1, 230, 230, 3): @@ -83,9 +83,9 @@ def conv2d_nhwc( @Ts.prim_func def conv2d_nhwc_reindex_data( - Input: T.Buffer((1, 224, 224, 3), "float32"), - Weight: T.Buffer((7, 7, 3, 64), "float32"), - Conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32"), + Input: T.Tensor((1, 224, 224, 3), "float32"), + Weight: T.Tensor((7, 7, 3, 64), "float32"), + Conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32"), ) -> None: PadInput = Ts.sblock_alloc_buffer([1, 230, 230, 3], dtype="float32") ReindexInput = Ts.sblock_alloc_buffer([1, 112, 112, 7, 7, 3], dtype="float32") @@ -113,9 +113,9 @@ def conv2d_nhwc_reindex_data( @Ts.prim_func def conv2d_nhwc_reindex_weight( - inputs: T.Buffer([1, 224, 224, 3], dtype="float32"), - weight: T.Buffer([7, 7, 3, 64], dtype="float32"), - conv2d_nhwc: T.Buffer([1, 112, 112, 64], dtype="float32"), + inputs: T.Tensor([1, 224, 224, 3], dtype="float32"), + weight: T.Tensor([7, 7, 3, 64], dtype="float32"), + conv2d_nhwc: T.Tensor([1, 112, 112, 64], dtype="float32"), ) -> None: PadInput = Ts.sblock_alloc_buffer([1, 230, 230, 3], dtype="float32") weight_reindex = Ts.sblock_alloc_buffer([64, 7, 7, 3], dtype="float32") @@ -154,9 +154,9 @@ def conv2d_nhwc_reindex_weight( @Ts.prim_func def matmul( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - C: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + C: T.Tensor((512, 512), "float32"), ) -> None: for i0, i1, i2 in T.grid(512, 512, 512): with Ts.sblock("matmul"): @@ -170,9 +170,9 @@ def matmul( @Ts.prim_func def matmul_reindex_write( - A: T.Buffer((512, 512), "float32"), - B: T.Buffer((512, 512), "float32"), - C: T.Buffer((512, 512), "float32"), + A: T.Tensor((512, 512), "float32"), + B: T.Tensor((512, 512), "float32"), + C: T.Tensor((512, 512), "float32"), ) -> None: C_reindex = Ts.sblock_alloc_buffer([512, 512], dtype="float32") for i0, i1, i2 in T.grid(512, 512, 512): @@ -192,7 +192,7 @@ def matmul_reindex_write( @Ts.prim_func -def multiple_read(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32")) -> None: +def multiple_read(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32")) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -201,9 +201,9 @@ def multiple_read(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "f @Ts.prim_func def mixed_dtype( - p0: T.Buffer((T.int64(2), 1280), "float16"), - p1: T.Buffer((1280, 1280), "float16"), - T_matmul_NT: T.Buffer((T.int64(2), 1280), "float16"), + p0: T.Tensor((T.int64(2), 1280), "float16"), + p1: T.Tensor((1280, 1280), "float16"), + T_matmul_NT: T.Tensor((T.int64(2), 1280), "float16"), ) -> None: for i0, i1, i2 in T.grid(T.int64(2), 1280, 1280): with Ts.sblock("T_matmul_NT"): @@ -218,9 +218,9 @@ def mixed_dtype( @Ts.prim_func def mixed_dtype_reindex_write( - p0: T.Buffer((T.int64(2), 1280), "float16"), - p1: T.Buffer((1280, 1280), "float16"), - T_matmul_NT: T.Buffer((T.int64(2), 1280), "float16"), + p0: T.Tensor((T.int64(2), 1280), "float16"), + p1: T.Tensor((1280, 1280), "float16"), + T_matmul_NT: T.Tensor((T.int64(2), 1280), "float16"), ) -> None: T_matmul_NT_reindex = Ts.sblock_alloc_buffer([T.int64(2), 1280], dtype="float16") for i0, i1, i2 in T.grid(T.int64(2), 1280, 1280): @@ -243,9 +243,9 @@ def mixed_dtype_reindex_write( @Ts.prim_func def matmul_unit_dim( - A: T.Buffer((1, 512), "float32"), - B: T.Buffer((512, 1), "float32"), - C: T.Buffer((1, 1), "float32"), + A: T.Tensor((1, 512), "float32"), + B: T.Tensor((512, 1), "float32"), + C: T.Tensor((1, 1), "float32"), ) -> None: for i0, i1, i2 in T.grid(1, 1, 512): with Ts.sblock("matmul"): @@ -259,9 +259,9 @@ def matmul_unit_dim( @Ts.prim_func def matmul_unit_dim_reindex_write( - A: T.Buffer((1, 512), "float32"), - B: T.Buffer((512, 1), "float32"), - C: T.Buffer((1, 1), "float32"), + A: T.Tensor((1, 512), "float32"), + B: T.Tensor((512, 1), "float32"), + C: T.Tensor((1, 1), "float32"), ) -> None: C_reindex = Ts.sblock_alloc_buffer([1, 1], dtype="float32") for i0, i1, i2 in T.grid(1, 1, 512): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_reorder.py b/tests/python/s_tir/schedule/test_tir_schedule_reorder.py index 6938cf3ef1f0..296350a782c1 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_reorder.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_reorder.py @@ -34,7 +34,7 @@ @Ts.prim_func -def elementwise(A: T.Buffer((128, 128, 128, 128)), B: T.Buffer((128, 128, 128, 128))) -> None: +def elementwise(A: T.Tensor((128, 128, 128, 128)), B: T.Tensor((128, 128, 128, 128))) -> None: for i, j, k, l in T.grid(128, 128, 128, 128): with Ts.sblock("B"): vi, vj, vk, vl = Ts.axis.remap("SSSS", [i, j, k, l]) @@ -43,7 +43,7 @@ def elementwise(A: T.Buffer((128, 128, 128, 128)), B: T.Buffer((128, 128, 128, 1 @Ts.prim_func def elementwise_not_affine( - A: T.Buffer((128, 128, 128, 128)), B: T.Buffer((128, 128, 128, 128)) + A: T.Tensor((128, 128, 128, 128)), B: T.Tensor((128, 128, 128, 128)) ) -> None: for i, j, k, l in T.grid(128, 128, 128, 8): with Ts.sblock("B"): @@ -54,7 +54,7 @@ def elementwise_not_affine( @Ts.prim_func def elementwise_dependent_loop( - A: T.Buffer((128, 128, 128, 128)), B: T.Buffer((128, 128, 128, 128)) + A: T.Tensor((128, 128, 128, 128)), B: T.Tensor((128, 128, 128, 128)) ) -> None: for i in T.serial(0, 128): for j, k, l in T.grid(128, i, 128): @@ -65,7 +65,7 @@ def elementwise_dependent_loop( @Ts.prim_func def elementwise_predicate( - A: T.Buffer((128, 128, 128, 128)), B: T.Buffer((128, 128, 128, 128)) + A: T.Tensor((128, 128, 128, 128)), B: T.Tensor((128, 128, 128, 128)) ) -> None: for i, j, k, l in T.grid(128, 128, 128, 128): with Ts.sblock("B"): @@ -76,7 +76,7 @@ def elementwise_predicate( @Ts.prim_func def elementwise_non_single_branch( - A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128)) ) -> None: C = Ts.sblock_alloc_buffer((128, 128, 128)) @@ -93,7 +93,7 @@ def elementwise_non_single_branch( @Ts.prim_func def elementwise_with_loops_not_same_scope( - A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128)) ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("A"): @@ -108,7 +108,7 @@ def elementwise_with_loops_not_same_scope( @Ts.prim_func def elementwise_with_wrong_block_var_type( - A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128)) ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("B"): @@ -121,7 +121,7 @@ def elementwise_with_wrong_block_var_type( @Ts.prim_func def elementwise_reordered( - A: T.Buffer((128, 128, 128, 128)), B: T.Buffer((128, 128, 128, 128)) + A: T.Tensor((128, 128, 128, 128)), B: T.Tensor((128, 128, 128, 128)) ) -> None: for l, j, k, i in T.grid(128, 128, 128, 128): with Ts.sblock("B"): @@ -131,7 +131,7 @@ def elementwise_reordered( @Ts.prim_func def elementwise_reordered2( - A: T.Buffer((128, 128, 128, 128)), B: T.Buffer((128, 128, 128, 128)) + A: T.Tensor((128, 128, 128, 128)), B: T.Tensor((128, 128, 128, 128)) ) -> None: for k, j, i, l in T.grid(128, 128, 128, 128): with Ts.sblock("B"): @@ -141,7 +141,7 @@ def elementwise_reordered2( @Ts.prim_func def elementwise_reordered_with_predicate( - A: T.Buffer((128, 128, 128, 128)), B: T.Buffer((128, 128, 128, 128)) + A: T.Tensor((128, 128, 128, 128)), B: T.Tensor((128, 128, 128, 128)) ) -> None: for l, j, k, i in T.grid(128, 128, 128, 128): with Ts.sblock("B"): @@ -151,7 +151,7 @@ def elementwise_reordered_with_predicate( @Ts.prim_func -def opaque_access(A: T.Buffer([16, 16], "float32"), B: T.Buffer([16, 16], "float32")) -> None: +def opaque_access(A: T.Tensor([16, 16], "float32"), B: T.Tensor([16, 16], "float32")) -> None: for i, j in T.grid(16, 16): with Ts.sblock("A"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -168,7 +168,7 @@ def opaque_access(A: T.Buffer([16, 16], "float32"), B: T.Buffer([16, 16], "float @Ts.prim_func def opaque_access_reorder( - A: T.Buffer([16, 16], "float32"), B: T.Buffer([16, 16], "float32") + A: T.Tensor([16, 16], "float32"), B: T.Tensor([16, 16], "float32") ) -> None: for j, i in T.grid(16, 16): with Ts.sblock("A"): @@ -219,7 +219,7 @@ def test_reorder_with_opaque_access(): def test_reorder_overlapped_access(): @Ts.prim_func - def overlapped_access(A: T.Buffer((14, 4), "float32"), B: T.Buffer((14, 4), "float32")): + def overlapped_access(A: T.Tensor((14, 4), "float32"), B: T.Tensor((14, 4), "float32")): # example to write first axis multiple times for v0, v1, v2 in T.grid(6, 4, 4): with Ts.sblock("block"): @@ -228,7 +228,7 @@ def overlapped_access(A: T.Buffer((14, 4), "float32"), B: T.Buffer((14, 4), "flo B[i, j] = A[i, j] + 1.0 @Ts.prim_func - def overlapped_access_reorder(A: T.Buffer((14, 4), "float32"), B: T.Buffer((14, 4), "float32")): + def overlapped_access_reorder(A: T.Tensor((14, 4), "float32"), B: T.Tensor((14, 4), "float32")): # example to write first axis multiple times for v0, v2, v1 in T.grid(6, 4, 4): with Ts.sblock("block"): @@ -245,7 +245,7 @@ def overlapped_access_reorder(A: T.Buffer((14, 4), "float32"), B: T.Buffer((14, def test_reorder_with_partial_affineness(): @Ts.prim_func - def non_affine_func(A: T.Buffer((14, 4), "float32"), B: T.Buffer((14, 4), "float32")): + def non_affine_func(A: T.Tensor((14, 4), "float32"), B: T.Tensor((14, 4), "float32")): for v0, v1, v2 in T.grid(6, 4, 4): with Ts.sblock("block"): i = Ts.axis.spatial(14, v0 * v0 + v1) @@ -253,7 +253,7 @@ def non_affine_func(A: T.Buffer((14, 4), "float32"), B: T.Buffer((14, 4), "float B[i, j] = A[i, j] + 1.0 @Ts.prim_func - def non_affine_func_reorder(A: T.Buffer((14, 4), "float32"), B: T.Buffer((14, 4), "float32")): + def non_affine_func_reorder(A: T.Tensor((14, 4), "float32"), B: T.Tensor((14, 4), "float32")): for v0, v2, v1 in T.grid(6, 4, 4): with Ts.sblock("block"): i = Ts.axis.spatial(14, v0 * v0 + v1) @@ -273,7 +273,7 @@ def non_affine_func_reorder(A: T.Buffer((14, 4), "float32"), B: T.Buffer((14, 4) def test_reorder_with_cascade_tiled_ops(): @Ts.prim_func def cascade_pool_ops( - x: T.Buffer((1, 16, 112, 112), "float32"), y2: T.Buffer((1, 16, 108, 108), "float32") + x: T.Tensor((1, 16, 112, 112), "float32"), y2: T.Tensor((1, 16, 108, 108), "float32") ) -> None: y1 = Ts.sblock_alloc_buffer([1, 16, 110, 110], dtype="float32") for n, c, h, w, kh, kw in T.grid(1, 16, 110, 110, 3, 3): @@ -291,7 +291,7 @@ def cascade_pool_ops( @Ts.prim_func def cascade_pool_ops_tile_reordered( - x: T.Buffer((1, 16, 112, 112), "float32"), y2: T.Buffer((1, 16, 108, 108), "float32") + x: T.Tensor((1, 16, 112, 112), "float32"), y2: T.Tensor((1, 16, 108, 108), "float32") ) -> None: y1 = Ts.sblock_alloc_buffer([1, 16, 110, 110], dtype="float32") for n, c, h_o in T.grid(1, 16, 27): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_reorder_block_iter_var.py b/tests/python/s_tir/schedule/test_tir_schedule_reorder_block_iter_var.py index 2eea8612a7a6..dc9648bc48b9 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_reorder_block_iter_var.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_reorder_block_iter_var.py @@ -28,9 +28,9 @@ @Ts.prim_func def matmul( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("C"): @@ -42,9 +42,9 @@ def matmul( @Ts.prim_func def matmul_after_reorder_block_iter_var( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ): for i, j, k in T.grid(128, 128, 128): with Ts.sblock("C"): 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 0519fd5fe2a5..f3a1d54be7c8 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_rfactor.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_rfactor.py @@ -33,9 +33,9 @@ @Ts.prim_func def transformed_matmul( - A: T.Buffer([128, 128], dtype="float32"), - B: T.Buffer([128, 128], dtype="float32"), - C: T.Buffer([128, 128], dtype="float32"), + A: T.Tensor([128, 128], dtype="float32"), + B: T.Tensor([128, 128], dtype="float32"), + C: T.Tensor([128, 128], dtype="float32"), ) -> None: for i0, i1, i2_outer, i2_inner_outer, i2_inner_inner in T.grid(128, 128, 4, 8, 4): with Ts.sblock("update"): @@ -50,9 +50,9 @@ def transformed_matmul( @Ts.prim_func def transformed_matmul_with_let( - A: T.Buffer([128, 128], dtype="float32"), - B: T.Buffer([128, 128], dtype="float32"), - C: T.Buffer([128, 128], dtype="float32"), + A: T.Tensor([128, 128], dtype="float32"), + B: T.Tensor([128, 128], dtype="float32"), + C: T.Tensor([128, 128], dtype="float32"), ) -> None: for i0, i1, i2_outer, i2_inner_outer, i2_inner_inner in T.grid(128, 128, 4, 8, 4): with Ts.sblock("update"): @@ -68,9 +68,9 @@ def transformed_matmul_with_let( @Ts.prim_func def matmul_rfactor( - A: T.Buffer([128, 128], dtype="float32"), - B: T.Buffer([128, 128], dtype="float32"), - C: T.Buffer([128, 128], dtype="float32"), + A: T.Tensor([128, 128], dtype="float32"), + B: T.Tensor([128, 128], dtype="float32"), + C: T.Tensor([128, 128], dtype="float32"), ) -> None: C_rf = Ts.sblock_alloc_buffer([4, 128, 128], dtype="float32") @@ -98,7 +98,7 @@ def matmul_rfactor( @Ts.prim_func def matmul_not_stage_pipeline( - A: T.Buffer([256, 256]), B: T.Buffer([256, 256]), D: T.Buffer([256, 256]) + A: T.Tensor([256, 256]), B: T.Tensor([256, 256]), D: T.Tensor([256, 256]) ) -> None: C = Ts.sblock_alloc_buffer([256, 256]) @@ -117,7 +117,7 @@ def matmul_not_stage_pipeline( @Ts.prim_func def matmul_not_same_buffer_access( - A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 128)) + A: T.Tensor((128, 128)), B: T.Tensor((128, 128)), C: T.Tensor((128, 128)) ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("C"): @@ -129,10 +129,10 @@ def matmul_not_same_buffer_access( @Ts.prim_func def matmul_loop_multiple_children( - A: T.Buffer([128, 128]), - B: T.Buffer([128, 128]), - C: T.Buffer([128, 128]), - D: T.Buffer([128, 128]), + A: T.Tensor([128, 128]), + B: T.Tensor([128, 128]), + C: T.Tensor([128, 128]), + D: T.Tensor([128, 128]), ) -> None: for k, i, j in T.grid(128, 128, 128): with Ts.sblock("C"): @@ -148,7 +148,7 @@ def matmul_loop_multiple_children( @Ts.prim_func -def square_sum(A: T.Buffer([16, 256, 256]), C: T.Buffer([16])) -> None: +def square_sum(A: T.Tensor([16, 256, 256]), C: T.Tensor([16])) -> None: for b0, i0, j0 in T.grid(16, 256, 256): with Ts.sblock("C"): b, i, j = Ts.axis.remap("SRR", [b0, i0, j0]) @@ -158,7 +158,7 @@ def square_sum(A: T.Buffer([16, 256, 256]), C: T.Buffer([16])) -> None: @Ts.prim_func -def square_sum_rfactor(A: T.Buffer([16, 256, 256]), C: T.Buffer([16])) -> None: +def square_sum_rfactor(A: T.Tensor([16, 256, 256]), C: T.Tensor([16])) -> None: C_rf = Ts.sblock_alloc_buffer([16, 256]) for i0, i1, i2 in T.grid(16, 256, 256): @@ -177,7 +177,7 @@ def square_sum_rfactor(A: T.Buffer([16, 256, 256]), C: T.Buffer([16])) -> None: @Ts.prim_func -def transformed_square_sum_square_root(A: T.Buffer([16, 256, 256]), D: T.Buffer([16])) -> None: +def transformed_square_sum_square_root(A: T.Tensor([16, 256, 256]), D: T.Tensor([16])) -> None: C = Ts.sblock_alloc_buffer([16]) for i0, i1_i2_fused_outer, i1_i2_fused_inner in T.grid(16, 65536, 1): @@ -199,7 +199,7 @@ def transformed_square_sum_square_root(A: T.Buffer([16, 256, 256]), D: T.Buffer( @Ts.prim_func -def square_sum_square_root_rfactor(A: T.Buffer([16, 256, 256]), D: T.Buffer([16])) -> None: +def square_sum_square_root_rfactor(A: T.Tensor([16, 256, 256]), D: T.Tensor([16])) -> None: C = Ts.sblock_alloc_buffer([16]) C_rf = Ts.sblock_alloc_buffer([1, 16]) @@ -227,7 +227,7 @@ def square_sum_square_root_rfactor(A: T.Buffer([16, 256, 256]), D: T.Buffer([16] @Ts.prim_func def transformed_square_sum_square_root_factor_one_1( - A: T.Buffer([16, 256, 256]), D: T.Buffer([16]) + A: T.Tensor([16, 256, 256]), D: T.Tensor([16]) ) -> None: C = Ts.sblock_alloc_buffer([16]) @@ -247,7 +247,7 @@ def transformed_square_sum_square_root_factor_one_1( @Ts.prim_func def square_sum_square_root_factor_one_1_rfactor( - A: T.Buffer((16, 256, 256), "float32"), D: T.Buffer((16,), "float32") + A: T.Tensor((16, 256, 256), "float32"), D: T.Tensor((16,), "float32") ) -> None: C = Ts.sblock_alloc_buffer([16], dtype="float32") C_rf = Ts.sblock_alloc_buffer([1, 16], dtype="float32") @@ -274,7 +274,7 @@ def square_sum_square_root_factor_one_1_rfactor( @Ts.prim_func def transformed_square_sum_square_root_factor_one_2( - A: T.Buffer([16, 256, 256]), D: T.Buffer([16]) + A: T.Tensor([16, 256, 256]), D: T.Tensor([16]) ) -> None: C = Ts.sblock_alloc_buffer([16]) @@ -294,7 +294,7 @@ def transformed_square_sum_square_root_factor_one_2( @Ts.prim_func def square_sum_square_root_factor_one_2_rfactor( - A: T.Buffer((16, 256, 256), "float32"), D: T.Buffer((16,), "float32") + A: T.Tensor((16, 256, 256), "float32"), D: T.Tensor((16,), "float32") ) -> None: C = Ts.sblock_alloc_buffer([16], dtype="float32") C_rf = Ts.sblock_alloc_buffer([16, 1], dtype="float32") @@ -320,7 +320,7 @@ def square_sum_square_root_factor_one_2_rfactor( @Ts.prim_func -def square_sum_with_annotation(A: T.Buffer([16, 256, 256]), C: T.Buffer([16])) -> None: +def square_sum_with_annotation(A: T.Tensor([16, 256, 256]), C: T.Tensor([16])) -> None: for b0, i0, j0 in T.grid(16, 256, 256): with Ts.sblock("C"): Ts.sblock_attr({"test_annotation": 1}) @@ -331,7 +331,7 @@ def square_sum_with_annotation(A: T.Buffer([16, 256, 256]), C: T.Buffer([16])) - @Ts.prim_func -def square_sum_with_annotation_rfactor(A: T.Buffer([16, 256, 256]), C: T.Buffer([16])) -> None: +def square_sum_with_annotation_rfactor(A: T.Tensor([16, 256, 256]), C: T.Tensor([16])) -> None: C_rf = Ts.sblock_alloc_buffer([16, 256]) for i0, i1, i2 in T.grid(16, 256, 256): @@ -352,7 +352,7 @@ def square_sum_with_annotation_rfactor(A: T.Buffer([16, 256, 256]), C: T.Buffer( @Ts.prim_func -def element_wise(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: +def element_wise(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -360,7 +360,7 @@ def element_wise(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: @Ts.prim_func -def rowsum(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for i, k in T.grid(128, 128): with Ts.sblock("B"): vi, vk = Ts.axis.remap("SR", [i, k]) @@ -370,7 +370,7 @@ def rowsum(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: @Ts.prim_func -def rowsum_not_quasi_affine(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum_not_quasi_affine(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for i, k in T.grid(128, 16): with Ts.sblock("B"): vi = Ts.axis.S(128, i) @@ -381,7 +381,7 @@ def rowsum_not_quasi_affine(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> Non @Ts.prim_func -def rowsum_not_dominant(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> None: +def rowsum_not_dominant(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))) -> None: for i, k in T.grid(128, 128): with Ts.sblock("B"): vi, vk = Ts.axis.remap("SR", [i, k]) @@ -391,7 +391,7 @@ def rowsum_not_dominant(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))) -> Non @Ts.prim_func -def rowsum_not_serial(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum_not_serial(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for i in T.serial(0, 128): for k in T.parallel(0, 128): with Ts.sblock("B"): @@ -402,7 +402,7 @@ def rowsum_not_serial(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: @Ts.prim_func -def rowsum_wrong_reduce_pattern1(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum_wrong_reduce_pattern1(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for i, k in T.grid(128, 128): with Ts.sblock("B"): vi, vk = Ts.axis.remap("SR", [i, k]) @@ -412,7 +412,7 @@ def rowsum_wrong_reduce_pattern1(A: T.Buffer((128, 128)), B: T.Buffer((128,))) - @Ts.prim_func -def rowsum_wrong_reduce_pattern2(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum_wrong_reduce_pattern2(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for i, k in T.grid(128, 128): with Ts.sblock("B"): vi, vk = Ts.axis.remap("SR", [i, k]) @@ -422,7 +422,7 @@ def rowsum_wrong_reduce_pattern2(A: T.Buffer((128, 128)), B: T.Buffer((128,))) - @Ts.prim_func -def rowsum_init_not_bufferstore(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum_init_not_bufferstore(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for i, k in T.grid(128, 128): with Ts.sblock("B"): vi, vk = Ts.axis.remap("SR", [i, k]) @@ -433,7 +433,7 @@ def rowsum_init_not_bufferstore(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> @Ts.prim_func -def rowsum_transformed(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: +def rowsum_transformed(A: T.Tensor((128, 128)), B: T.Tensor((128,))) -> None: for io, ii_ko_fused, ki in T.grid(32, 128, 4): with Ts.sblock("B"): vi = Ts.axis.S(128, io * 4 + T.floordiv(ii_ko_fused, 32)) @@ -444,7 +444,7 @@ def rowsum_transformed(A: T.Buffer((128, 128)), B: T.Buffer((128,))) -> None: @Ts.prim_func -def rowsum_zero_dim(A: T.Buffer([128]), B: T.Buffer([])) -> None: +def rowsum_zero_dim(A: T.Tensor([128]), B: T.Tensor([])) -> None: for k0 in range(128): with Ts.sblock("B"): k = Ts.axis.R(128, k0) @@ -454,7 +454,7 @@ def rowsum_zero_dim(A: T.Buffer([128]), B: T.Buffer([])) -> None: @Ts.prim_func -def rowsum_zero_dim_rfactor(A: T.Buffer([128]), B: T.Buffer([])) -> None: +def rowsum_zero_dim_rfactor(A: T.Tensor([128]), B: T.Tensor([])) -> None: B_rf = Ts.sblock_alloc_buffer([128], elem_offset=T.int64(0)) for i in range(128): @@ -472,7 +472,7 @@ def rowsum_zero_dim_rfactor(A: T.Buffer([128]), B: T.Buffer([])) -> None: @Ts.prim_func def rowsum_predicate( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: for i, k_0, k_1 in T.grid(128, 13, 10): with Ts.sblock("B"): @@ -486,7 +486,7 @@ def rowsum_predicate( @Ts.prim_func def rowsum_predicate_rfactor( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: B_rf = Ts.sblock_alloc_buffer([128, 13], dtype="float32") for i, k_0, k_1 in T.grid(128, 13, 10): @@ -505,7 +505,7 @@ def rowsum_predicate_rfactor( @Ts.prim_func -def multiple_reduction_blocks(A: T.Buffer((16, 16, 16)), F: T.Buffer((16, 16))) -> None: +def multiple_reduction_blocks(A: T.Tensor((16, 16, 16)), F: T.Tensor((16, 16))) -> None: C = Ts.sblock_alloc_buffer((16, 16)) D = Ts.sblock_alloc_buffer((16, 16)) E = Ts.sblock_alloc_buffer((16, 16)) @@ -544,7 +544,7 @@ def multiple_reduction_blocks(A: T.Buffer((16, 16, 16)), F: T.Buffer((16, 16))) @Ts.prim_func -def multiple_reduction_blocks_rfactor(A: T.Buffer([16, 16, 16]), F: T.Buffer([16, 16])) -> None: +def multiple_reduction_blocks_rfactor(A: T.Tensor([16, 16, 16]), F: T.Tensor([16, 16])) -> None: C = Ts.sblock_alloc_buffer([16, 16]) D = Ts.sblock_alloc_buffer([16, 16]) E = Ts.sblock_alloc_buffer([16, 16]) @@ -591,8 +591,8 @@ def multiple_reduction_blocks_rfactor(A: T.Buffer([16, 16, 16]), F: T.Buffer([16 @Ts.prim_func def rfactor_spatial_only( - A: T.Buffer((1, 512, 7, 7), "float32"), - B: T.Buffer((1, 512, 1, 1), "float32"), + A: T.Tensor((1, 512, 7, 7), "float32"), + B: T.Tensor((1, 512, 1, 1), "float32"), ) -> None: for _i0, i1, _i2, _i3, i4, _i5 in T.grid(1, 512, 1, 1, 49, 1): with Ts.sblock("acc"): @@ -613,8 +613,8 @@ def rfactor_spatial_only( @Ts.prim_func def rfactor_spatial_only_after( - A: T.Buffer((1, 512, 7, 7), "float32"), - B: T.Buffer((1, 512, 1, 1), "float32"), + A: T.Tensor((1, 512, 7, 7), "float32"), + B: T.Tensor((1, 512, 1, 1), "float32"), ) -> None: # body # with Ts.sblock("root") @@ -641,10 +641,10 @@ def rfactor_spatial_only_after( @Ts.prim_func def argmax_split( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -667,10 +667,10 @@ def argmax_split( @Ts.prim_func def argmin_split_init_update_reordered( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmin_v0: T.Buffer((128,), "int32"), - argmin_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmin_v0: T.Tensor((128,), "int32"), + argmin_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmin"): @@ -693,10 +693,10 @@ def argmin_split_init_update_reordered( @Ts.prim_func def argmax_split_different_shape( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((256,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((256,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -719,10 +719,10 @@ def argmax_split_different_shape( @Ts.prim_func def argmax_split_different_indices( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -745,10 +745,10 @@ def argmax_split_different_indices( @Ts.prim_func def argmax_split_init_not_bufferstore( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -772,10 +772,10 @@ def argmax_split_init_not_bufferstore( @Ts.prim_func def argmax_split_init_buffer_duplicate( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -798,10 +798,10 @@ def argmax_split_init_buffer_duplicate( @Ts.prim_func def argmax_split_bind_fewer_than_init( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -821,10 +821,10 @@ def argmax_split_bind_fewer_than_init( @Ts.prim_func def argmax_split_bind_more_than_init( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -846,10 +846,10 @@ def argmax_split_bind_more_than_init( @Ts.prim_func def argmax_split_let_body_neither_seqstmt_nor_bufferstore( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -871,10 +871,10 @@ def argmax_split_let_body_neither_seqstmt_nor_bufferstore( @Ts.prim_func def argmax_split_init_update_inconsistent_bufferstore_number( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -898,10 +898,10 @@ def argmax_split_init_update_inconsistent_bufferstore_number( @Ts.prim_func def argmax_split_body_seq_not_bufferstore( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -924,10 +924,10 @@ def argmax_split_body_seq_not_bufferstore( @Ts.prim_func def argmax_split_body_bufferstore_value_not_var( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -954,10 +954,10 @@ def argmax_split_body_bufferstore_value_not_var( @Ts.prim_func(check_well_formed=False) def argmax_split_body_bufferstore_value_unbound_var( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -980,10 +980,10 @@ def argmax_split_body_bufferstore_value_unbound_var( @Ts.prim_func def argmax_split_one_let_var_used_multi_times( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "int32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "int32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "int32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "int32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -1006,10 +1006,10 @@ def argmax_split_one_let_var_used_multi_times( @Ts.prim_func def argmax_split_body_one_buffer_updated_multi_times( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "int32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "int32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "int32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "int32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -1032,11 +1032,11 @@ def argmax_split_body_one_buffer_updated_multi_times( @Ts.prim_func def argmax_split_init_buffer_not_match( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v0_1: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v0_1: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0, i1_1 in T.grid(128, 4, 32): with Ts.sblock("argmax"): @@ -1059,10 +1059,10 @@ def argmax_split_init_buffer_not_match( @Ts.prim_func def argmax_split_rfactor( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: argmax_v0_rf = Ts.sblock_alloc_buffer([128, 32], dtype="int32") argmax_v1_rf = Ts.sblock_alloc_buffer([128, 32], dtype="float32") @@ -1106,10 +1106,10 @@ def argmax_split_rfactor( @Ts.prim_func def argmin_split_rfactor( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmin_v0: T.Buffer((128,), "int32"), - argmin_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmin_v0: T.Tensor((128,), "int32"), + argmin_v1: T.Tensor((128,), "float32"), ) -> None: argmin_v0_rf = Ts.sblock_alloc_buffer([128, 32], dtype="int32") argmin_v1_rf = Ts.sblock_alloc_buffer([128, 32], dtype="float32") @@ -1153,7 +1153,7 @@ def argmin_split_rfactor( @Ts.prim_func def argmax_topi_rfactor( - placeholder: T.Buffer((1, 32), "int32"), placeholder_red: T.Buffer(1, "int32") + placeholder: T.Tensor((1, 32), "int32"), placeholder_red: T.Tensor(1, "int32") ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) placeholder_red_temp_v0 = Ts.sblock_alloc_buffer([1], dtype="int32") @@ -1220,7 +1220,7 @@ def argmax_topi_rfactor( @Ts.prim_func def argmin_topi_rfactor( - placeholder: T.Buffer((1, 32), "int32"), placeholder_red: T.Buffer(1, "int32") + placeholder: T.Tensor((1, 32), "int32"), placeholder_red: T.Tensor(1, "int32") ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) placeholder_red_temp_v0 = Ts.sblock_alloc_buffer([1], dtype="int32") @@ -1287,7 +1287,7 @@ def argmin_topi_rfactor( @Ts.prim_func def argmax_topi_select_last_rfactor( - placeholder: T.Buffer((1, 32), "int32"), placeholder_red: T.Buffer(1, "int32") + placeholder: T.Tensor((1, 32), "int32"), placeholder_red: T.Tensor(1, "int32") ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) placeholder_red_temp_v0 = Ts.sblock_alloc_buffer([1], dtype="int32") @@ -1354,7 +1354,7 @@ def argmax_topi_select_last_rfactor( @Ts.prim_func def argmin_topi_select_last_rfactor( - placeholder: T.Buffer((1, 32), "int32"), placeholder_red: T.Buffer(1, "int32") + placeholder: T.Tensor((1, 32), "int32"), placeholder_red: T.Tensor(1, "int32") ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) placeholder_red_temp_v0 = Ts.sblock_alloc_buffer([1], dtype="int32") @@ -1843,9 +1843,9 @@ def test_reduction_rfactor_int64(): # fmt: off @Ts.prim_func def before( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(128)), "float32"), - C: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(128)), "float32"), + C: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): for i0, i1, i2_outer, i2_inner_outer, i2_inner_inner in T.grid( T.int64(128), T.int64(128), T.int64(4), T.int64(8), T.int64(4) @@ -1861,9 +1861,9 @@ def before( C[vi, vj] = C[vi, vj] + (A[vi, vk] * B[vj, vk]) @Ts.prim_func - def expected(A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(128)), "float32"), - C: T.Buffer((T.int64(128), T.int64(128)), "float32"), + def expected(A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(128)), "float32"), + C: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): C_rf = Ts.sblock_alloc_buffer((T.int64(4), T.int64(128), T.int64(128)), "float32") diff --git a/tests/python/s_tir/schedule/test_tir_schedule_rolling_buffer.py b/tests/python/s_tir/schedule/test_tir_schedule_rolling_buffer.py index de92cfad529f..7080edfa376c 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_rolling_buffer.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_rolling_buffer.py @@ -66,7 +66,7 @@ def _tile_nd(s, tile, block_name): def test_1d_rolling_buffer(): @Ts.prim_func - def before(A: T.Buffer((4, 12), "int32"), C: T.Buffer((4, 8), "int32")): + def before(A: T.Tensor((4, 12), "int32"), C: T.Tensor((4, 8), "int32")): B = Ts.sblock_alloc_buffer((4, 10), "int32") for c in T.serial(4): for i in T.serial(0, 10): @@ -85,7 +85,7 @@ def before(A: T.Buffer((4, 12), "int32"), C: T.Buffer((4, 8), "int32")): C[cc, vi] = C[cc, vi] + B[cc, vi + vk] @Ts.prim_func - def expected(A: T.Buffer((4, 12), "int32"), C: T.Buffer((4, 8), "int32")): + def expected(A: T.Tensor((4, 12), "int32"), C: T.Tensor((4, 8), "int32")): B = Ts.sblock_alloc_buffer([4, 6], dtype="int32") for c, i_0 in T.grid(4, 2): for ax0, ax1 in T.grid(6, 3): @@ -119,7 +119,7 @@ def expected(A: T.Buffer((4, 12), "int32"), C: T.Buffer((4, 8), "int32")): @Ts.prim_func -def cascade_2_max_pool2d(A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "int8")): +def cascade_2_max_pool2d(A: T.Tensor((1, 12, 12, 16), "int8"), C: T.Tensor((1, 8, 8, 16), "int8")): B = Ts.sblock_alloc_buffer([1, 10, 10, 16], dtype="int8") for i0, i1, i2, i3, i4, i5 in T.grid(1, 10, 10, 16, 3, 3): with Ts.sblock("B"): @@ -137,7 +137,7 @@ def cascade_2_max_pool2d(A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 8 @Ts.prim_func def cascade_3_max_pool2d_with_stride( - A: T.Buffer((1, 24, 24, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "int8") + A: T.Tensor((1, 24, 24, 16), "int8"), C: T.Tensor((1, 8, 8, 16), "int8") ): B_0 = Ts.sblock_alloc_buffer([1, 22, 22, 16], dtype="int8") B_1 = Ts.sblock_alloc_buffer([1, 10, 10, 16], dtype="int8") @@ -169,7 +169,7 @@ def cascade_3_max_pool2d_with_stride( def test_cascade_max_pool2d_w_tiled(): @Ts.prim_func - def expected(A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "int8")): + def expected(A: T.Tensor((1, 12, 12, 16), "int8"), C: T.Tensor((1, 8, 8, 16), "int8")): B = Ts.sblock_alloc_buffer([1, 10, 6, 16], dtype="int8") for i0_0, i1_0, i2_0, i3_0 in T.grid(1, 1, 2, 1): for ax0, ax1, ax2, ax3, ax4 in T.grid(10, 6, 16, 3, 3): @@ -210,7 +210,7 @@ def expected(A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "i def test_cascade_max_pool2d_h_tiled(): @Ts.prim_func - def expected(A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "int8")): + def expected(A: T.Tensor((1, 12, 12, 16), "int8"), C: T.Tensor((1, 8, 8, 16), "int8")): B = Ts.sblock_alloc_buffer([1, 6, 10, 16], dtype="int8") for i0_0, i1_0, i2_0, i3_0 in T.grid(1, 2, 1, 1): for ax0, ax1, ax2, ax3, ax4 in T.grid(6, 10, 16, 3, 3): @@ -251,7 +251,7 @@ def expected(A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "i def test_cascade_max_pool2d_h_w_c_tiled(): @Ts.prim_func - def expected(A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "int8")): + def expected(A: T.Tensor((1, 12, 12, 16), "int8"), C: T.Tensor((1, 8, 8, 16), "int8")): B = Ts.sblock_alloc_buffer([1, 6, 10, 16], dtype="int8") for i0_0, i1_0, i2_0, i3_0 in T.grid(1, 2, 2, 2): for ax0, ax1, ax2, ax3, ax4 in T.grid(6, 6, 8, 3, 3): @@ -293,7 +293,7 @@ def expected(A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "i def test_cascade_max_pool2d_non_perfect_tiled(): @Ts.prim_func - def expected(A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "int8")) -> None: + def expected(A: T.Tensor((1, 12, 12, 16), "int8"), C: T.Tensor((1, 8, 8, 16), "int8")) -> None: B = Ts.sblock_alloc_buffer([1, 8, 10, 16], dtype="int8") for i0_0, i1_0, i2_0, i3_0 in T.grid(1, 2, 2, 1): for ax0, ax1, ax2, ax3, ax4 in T.grid(8, 8, 16, 3, 3): @@ -340,7 +340,7 @@ def expected(A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "i def test_cascade_3_max_pool2d_with_stride(): @Ts.prim_func - def expected(A: T.Buffer((1, 24, 24, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "int8")) -> None: + def expected(A: T.Tensor((1, 24, 24, 16), "int8"), C: T.Tensor((1, 8, 8, 16), "int8")) -> None: B_0 = Ts.sblock_alloc_buffer([1, 13, 22, 16], dtype="int8") B_1 = Ts.sblock_alloc_buffer([1, 6, 10, 16], dtype="int8") for i0_0, i1_0, i2_0, i3_0 in T.grid(1, 2, 2, 1): @@ -401,7 +401,7 @@ def expected(A: T.Buffer((1, 24, 24, 16), "int8"), C: T.Buffer((1, 8, 8, 16), "i def test_upscale(): @Ts.prim_func - def before(A: T.Buffer((1, 16, 16, 16), "int8"), C: T.Buffer((1, 24, 24, 16), "int8")) -> None: + def before(A: T.Tensor((1, 16, 16, 16), "int8"), C: T.Tensor((1, 24, 24, 16), "int8")) -> None: B = Ts.sblock_alloc_buffer([1, 14, 14, 16], dtype="int8") for i0_0, i1_0, i2_0, i3_0 in T.grid(1, 5, 5, 1): for ax0, ax1, ax2, ax3, ax4 in T.grid(5, 5, 16, 3, 3): @@ -437,7 +437,7 @@ def before(A: T.Buffer((1, 16, 16, 16), "int8"), C: T.Buffer((1, 24, 24, 16), "i @Ts.prim_func def expected( - A: T.Buffer((1, 16, 16, 16), "int8"), C: T.Buffer((1, 24, 24, 16), "int8") + A: T.Tensor((1, 16, 16, 16), "int8"), C: T.Tensor((1, 24, 24, 16), "int8") ) -> None: B = Ts.sblock_alloc_buffer([1, 5, 14, 16], dtype="int8") for i0_0, i1_0, i2_0, i3_0 in T.grid(1, 5, 5, 1): @@ -485,7 +485,7 @@ def expected( def test_fail_rolling_buffer_multi_writers(): @Ts.prim_func def func_multi_writers( - A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 12, 12, 16), "int8") + A: T.Tensor((1, 12, 12, 16), "int8"), C: T.Tensor((1, 12, 12, 16), "int8") ): B = Ts.sblock_alloc_buffer([1, 12, 12, 16], dtype="int8") for i0, i1, i2, i3 in T.grid(1, 3, 3, 1): @@ -530,7 +530,7 @@ def func_multi_writers( def test_fail_rolling_buffer_not_match(): @Ts.prim_func def func_non_overlap( - A: T.Buffer((1, 12, 12, 16), "int8"), C: T.Buffer((1, 12, 12, 16), "int8") + A: T.Tensor((1, 12, 12, 16), "int8"), C: T.Tensor((1, 12, 12, 16), "int8") ): B = Ts.sblock_alloc_buffer([1, 12, 12, 16], dtype="int8") for i0_0, i1_0, i2_0, i3_0 in T.grid(1, 3, 3, 1): 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 3820cd0aa210..fb7dbe8aadf1 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_sampling.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_sampling.py @@ -31,7 +31,7 @@ @Ts.prim_func -def elementwise(A: T.Buffer((128, 257, 1470)), B: T.Buffer((128, 257, 1470))) -> None: +def elementwise(A: T.Tensor((128, 257, 1470)), B: T.Tensor((128, 257, 1470))) -> None: for i, j, k in T.grid(128, 257, 1470): with Ts.sblock("B"): vi, vj, vk = Ts.axis.remap("SSS", [i, j, k]) @@ -40,9 +40,9 @@ def elementwise(A: T.Buffer((128, 257, 1470)), B: T.Buffer((128, 257, 1470))) -> @Ts.prim_func def tiled_conv2d_with_padding( - inputs: T.Buffer((1, 224, 224, 3), "float32"), - weight: T.Buffer((7, 7, 3, 64), "float32"), - conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32"), + inputs: T.Tensor((1, 224, 224, 3), "float32"), + weight: T.Tensor((7, 7, 3, 64), "float32"), + conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32"), ) -> None: PadInput = Ts.sblock_alloc_buffer([1, 230, 230, 3], dtype="float32") for i0, i1, i2, i3 in T.grid(1, 230, 230, 3): @@ -216,7 +216,7 @@ def test_sample_perfect_tile_on_dynamic_loops(): n = T.dynamic("n", "int32") @Ts.prim_func - def workload(A: T.Buffer((n, 1024))) -> None: + def workload(A: T.Tensor((n, 1024))) -> None: for i, j in T.grid(n, 1024): with Ts.sblock("B"): vi, vj = Ts.axis.remap("SS", [i, j]) diff --git a/tests/python/s_tir/schedule/test_tir_schedule_set_dtype.py b/tests/python/s_tir/schedule/test_tir_schedule_set_dtype.py index c0d9e5e9d0f3..a0de519aed0d 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_set_dtype.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_set_dtype.py @@ -33,7 +33,7 @@ # pylint: disable=no-member,invalid-name,unused-variable,unexpected-keyword-arg @Ts.prim_func -def element_wise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")) -> None: +def element_wise(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")) -> None: B = Ts.sblock_alloc_buffer((128, 128), dtype="float32") for i, j in T.grid(128, 128): @@ -46,7 +46,7 @@ def element_wise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "fl C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def element_wise_set_dtype(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")): +def element_wise_set_dtype(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")): B = Ts.sblock_alloc_buffer((128, 128), "float16") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -62,7 +62,7 @@ def element_wise_set_dtype(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, C[vi, vj] = T.cast(B[vi, vj], "float32") + 1.0 @Ts.prim_func -def element_wise_subregion_match(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")) -> None: +def element_wise_subregion_match(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")) -> None: B = Ts.sblock_alloc_buffer((128, 128), dtype="float32") for i, j in T.grid(128, 128): @@ -77,7 +77,7 @@ def element_wise_subregion_match(A: T.Buffer((128, 128), "float32"), C: T.Buffer C[vi, vj] = B_subregion1[()] + 1.0 @Ts.prim_func -def element_wise_subregion_match_set_dtype(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")) -> None: +def element_wise_subregion_match_set_dtype(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float16") for i, j in T.grid(128, 128): with Ts.sblock("B"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_set_scope.py b/tests/python/s_tir/schedule/test_tir_schedule_set_scope.py index a6cad4d3afaa..30f7d4f25120 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_set_scope.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_set_scope.py @@ -32,7 +32,7 @@ # pylint: disable=no-member,invalid-name,unused-variable,unexpected-keyword-arg @Ts.prim_func -def element_wise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")) -> None: +def element_wise(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")) -> None: B = Ts.sblock_alloc_buffer((128, 128), dtype="float32") for i, j in T.grid(128, 128): @@ -45,7 +45,7 @@ def element_wise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "fl C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def element_wise_set_scope(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")) -> None: +def element_wise_set_scope(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")) -> None: B_shared = Ts.sblock_alloc_buffer([128, 128], dtype="float32", scope="shared") for i, j in T.grid(128, 128): @@ -58,7 +58,7 @@ def element_wise_set_scope(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, C[vi, vj] = B_shared[vi, vj] + T.float32(1) @Ts.prim_func -def element_wise_subregion_match(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")) -> None: +def element_wise_subregion_match(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")) -> None: B = Ts.sblock_alloc_buffer((128, 128), dtype="float32") for i, j in T.grid(128, 128): @@ -73,7 +73,7 @@ def element_wise_subregion_match(A: T.Buffer((128, 128), "float32"), C: T.Buffer C[vi, vj] = B_subregion1[()] + 1.0 @Ts.prim_func -def element_wise_subregion_match_set_scope(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")) -> None: +def element_wise_subregion_match_set_scope(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")) -> None: B_shared = Ts.sblock_alloc_buffer([128, 128], dtype="float32", scope="shared") for i, j in T.grid(128, 128): 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 06aa8c363a84..a201bb047436 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 @@ -35,7 +35,7 @@ @Ts.prim_func -def elementwise(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128))) -> None: +def elementwise(A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128))) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("B"): vi, vj, vk = Ts.axis.remap("SSS", [i, j, k]) @@ -43,7 +43,7 @@ def elementwise(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128))) -> N @Ts.prim_func -def elementwise_dependent_loops(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128))) -> None: +def elementwise_dependent_loops(A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128))) -> None: for i in T.serial(0, 128): for j, k in T.grid(i, 128): with Ts.sblock("B"): @@ -55,8 +55,8 @@ def elementwise_dependent_loops(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, @Ts.prim_func def elementwise_symbolic( - A: T.Buffer((128, 128, n)), # noqa: F821 - B: T.Buffer((128, 128, n)), # noqa: F821 + A: T.Tensor((128, 128, n)), # noqa: F821 + B: T.Tensor((128, 128, n)), # noqa: F821 n: T.int32, ) -> None: for i, j, k in T.grid(128, 128, n): @@ -67,8 +67,8 @@ def elementwise_symbolic( @Ts.prim_func def elementwise_symbolic_fused( - A: T.Buffer((128, 128, n)), # noqa: F821 - B: T.Buffer((128, 128, n)), # noqa: F821 + A: T.Tensor((128, 128, n)), # noqa: F821 + B: T.Tensor((128, 128, n)), # noqa: F821 n: T.int32, ) -> None: for i_j_k_fused in T.serial(0, (n * 16384)): @@ -83,8 +83,8 @@ def elementwise_symbolic_fused( @Ts.prim_func def elementwise_symbolic_split( - A: T.Buffer((128, 128, n)), # noqa: F821 - B: T.Buffer((128, 128, n)), # noqa: F821 + A: T.Tensor((128, 128, n)), # noqa: F821 + B: T.Tensor((128, 128, n)), # noqa: F821 n: T.int32, ) -> None: for i, j, k0, k1 in T.grid(128, 128, 10, T.floordiv((n + 9), 10)): @@ -98,7 +98,7 @@ def elementwise_symbolic_split( @Ts.prim_func -def elementwise_with_seq(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128))) -> None: +def elementwise_with_seq(A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128))) -> None: C = Ts.sblock_alloc_buffer((128, 128, 128)) for i, j in T.grid(128, 128): for k in T.serial(0, 128): @@ -112,7 +112,7 @@ def elementwise_with_seq(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 12 @Ts.prim_func -def elementwise_with_anno(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128))) -> None: +def elementwise_with_anno(A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128))) -> None: for i, j in T.grid(128, 128): for k in T.serial(0, 128, annotations={"useless_annotation": True}): with Ts.sblock("B"): @@ -124,7 +124,7 @@ def elementwise_with_anno(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 1 @Ts.prim_func def elementwise_with_thread_binding( - A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128)) ) -> None: for i, j in T.grid(128, 128): for k in T.thread_binding(0, 128, thread="threadIdx.x"): @@ -137,7 +137,7 @@ def elementwise_with_thread_binding( @Ts.prim_func def elementwise_with_starting_point( - A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128)) ) -> None: for i, j in T.grid(128, 128): for k in T.serial(10, 128): @@ -150,7 +150,7 @@ def elementwise_with_starting_point( @Ts.prim_func def elementwise_with_opaque_block( - A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128)) ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("opaque"): @@ -164,7 +164,7 @@ def elementwise_with_opaque_block( @Ts.prim_func -def elementwise_fused(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128))) -> None: +def elementwise_fused(A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128))) -> None: for fused in T.serial(0, 2097152): with Ts.sblock("B"): vi = Ts.axis.S(128, T.floordiv(fused, 16384)) @@ -176,7 +176,7 @@ def elementwise_fused(A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128)) @Ts.prim_func -def elementwise_split_case0(A: T.Buffer([128, 128, 128]), B: T.Buffer([128, 128, 128])) -> None: +def elementwise_split_case0(A: T.Tensor([128, 128, 128]), B: T.Tensor([128, 128, 128])) -> None: for i1, i2, i3, j1, j2, k1, k2 in T.grid(2, 1, 64, 4, 32, 16, 8): with Ts.sblock("B"): vi = Ts.axis.S(128, i1 * 64 + i2 * 64 + i3) @@ -188,7 +188,7 @@ def elementwise_split_case0(A: T.Buffer([128, 128, 128]), B: T.Buffer([128, 128, @Ts.prim_func -def elementwise_split_case1(A: T.Buffer([128, 128, 128]), B: T.Buffer([128, 128, 128])) -> None: +def elementwise_split_case1(A: T.Tensor([128, 128, 128]), B: T.Tensor([128, 128, 128])) -> None: for i1, i2, i3, j1, j2, j3, k1, k2, k3 in T.grid(2, 1, 64, 2, 1, 64, 2, 1, 64): with Ts.sblock("B"): vi = Ts.axis.S(128, i1 * 64 + i2 * 64 + i3) @@ -201,7 +201,7 @@ def elementwise_split_case1(A: T.Buffer([128, 128, 128]), B: T.Buffer([128, 128, @Ts.prim_func def elementwise_split_with_predicate( - A: T.Buffer([128, 128, 128]), B: T.Buffer([128, 128, 128]) + A: T.Tensor([128, 128, 128]), B: T.Tensor([128, 128, 128]) ) -> None: for i0, i1, i2, j0, j1, k0, k1 in T.grid(1000, 2, 3, 1, 129, 3, 43): with Ts.sblock("B"): @@ -216,7 +216,7 @@ def elementwise_split_with_predicate( @Ts.prim_func def elementwise_fuse_with_opaque_block( - A: T.Buffer([128, 128, 128]), B: T.Buffer([128, 128, 128]) + A: T.Tensor([128, 128, 128]), B: T.Tensor([128, 128, 128]) ) -> None: for i_j_k_fused in T.serial(0, 2097152): with Ts.sblock("opaque"): @@ -249,7 +249,7 @@ def elementwise_fuse_with_opaque_block( @Ts.prim_func def elementwise_split_with_opaque_block( - A: T.Buffer([128, 128, 128]), B: T.Buffer([128, 128, 128]) + A: T.Tensor([128, 128, 128]), B: T.Tensor([128, 128, 128]) ) -> None: for i0, i1, j, k in T.grid(8, 16, 128, 128): with Ts.sblock("opaque"): @@ -264,7 +264,7 @@ def elementwise_split_with_opaque_block( @Ts.prim_func -def opaque_access(A: T.Buffer([16, 16], "float32"), B: T.Buffer([16, 16], "float32")) -> None: +def opaque_access(A: T.Tensor([16, 16], "float32"), B: T.Tensor([16, 16], "float32")) -> None: for i, j in T.grid(16, 16): with Ts.sblock("A"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -280,7 +280,7 @@ def opaque_access(A: T.Buffer([16, 16], "float32"), B: T.Buffer([16, 16], "float @Ts.prim_func -def opaque_access_fused(A: T.Buffer([16, 16]), B: T.Buffer([16, 16])) -> None: +def opaque_access_fused(A: T.Tensor([16, 16]), B: T.Tensor([16, 16])) -> None: for i_j_fused in T.serial(0, 256): with Ts.sblock("A"): vi = Ts.axis.S(16, T.floordiv(i_j_fused, 16)) @@ -298,7 +298,7 @@ def opaque_access_fused(A: T.Buffer([16, 16]), B: T.Buffer([16, 16])) -> None: @Ts.prim_func -def opaque_access_split(A: T.Buffer((16, 16)), B: T.Buffer((16, 16))) -> None: +def opaque_access_split(A: T.Tensor((16, 16)), B: T.Tensor((16, 16))) -> None: for i, j0, j1 in T.grid(16, 4, 4): with Ts.sblock("A"): vi = Ts.axis.S(16, i) @@ -316,7 +316,7 @@ def opaque_access_split(A: T.Buffer((16, 16)), B: T.Buffer((16, 16))) -> None: @Ts.prim_func -def elementwise_not_affine(A: T.Buffer((127, 128)), B: T.Buffer((127, 128))) -> None: +def elementwise_not_affine(A: T.Tensor((127, 128)), B: T.Tensor((127, 128))) -> None: for i in T.serial(0, 4): for j, k in T.grid(T.min(31, 126 - i * 32) + 1, 128): with Ts.sblock("B"): @@ -326,7 +326,7 @@ def elementwise_not_affine(A: T.Buffer((127, 128)), B: T.Buffer((127, 128))) -> @Ts.prim_func -def elementwise_not_affine_fused(A: T.Buffer([127, 128]), B: T.Buffer([127, 128])) -> None: +def elementwise_not_affine_fused(A: T.Tensor([127, 128]), B: T.Tensor([127, 128])) -> None: for i in T.grid(4): for j_k_fused in T.serial(0, T.min(31, 126 - i * 32) * 128 + 128): with Ts.sblock("B"): @@ -380,7 +380,7 @@ def test_split_with_dynamic_inferred_factor(): M = T.dynamic("M", "int32") @Ts.prim_func - def before(A: T.Buffer((N, 128, M)), B: T.Buffer((N, 128, M))) -> None: + def before(A: T.Tensor((N, 128, M)), B: T.Tensor((N, 128, M))) -> None: for i, j, k in T.grid(N, 128, M): with Ts.sblock("B"): vi, vj, vk = Ts.axis.remap("SSS", [i, j, k]) @@ -390,7 +390,7 @@ def before(A: T.Buffer((N, 128, M)), B: T.Buffer((N, 128, M))) -> None: M = T.dynamic("M", "int32") @Ts.prim_func - def expected(A: T.Buffer((N, 128, M)), B: T.Buffer((N, 128, M))) -> None: + def expected(A: T.Tensor((N, 128, M)), B: T.Tensor((N, 128, M))) -> None: 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): with Ts.sblock("B"): vi = Ts.axis.spatial(N, i_0 * 16 + i_1) @@ -554,9 +554,9 @@ def test_fuse_not_affine(): def test_add_unit_loop_above_block(): @Ts.prim_func def zero_dim( - A: T.Buffer((), "int32"), - B: T.Buffer((), "int32"), - C: T.Buffer((), "int32"), + A: T.Tensor((), "int32"), + B: T.Tensor((), "int32"), + C: T.Tensor((), "int32"), ) -> None: with Ts.sblock("C"): vi = Ts.axis.spatial(1, 0) @@ -564,9 +564,9 @@ def zero_dim( @Ts.prim_func def zero_dim_added( - A: T.Buffer((), "int32"), - B: T.Buffer((), "int32"), - C: T.Buffer((), "int32"), + A: T.Tensor((), "int32"), + B: T.Tensor((), "int32"), + C: T.Tensor((), "int32"), ) -> None: for u in range(1): with Ts.sblock("C"): @@ -582,9 +582,9 @@ def zero_dim_added( def test_add_unit_loop_above_loop(): @Ts.prim_func def zero_dim( - A: T.Buffer((), "int32"), - B: T.Buffer((), "int32"), - C: T.Buffer((), "int32"), + A: T.Tensor((), "int32"), + B: T.Tensor((), "int32"), + C: T.Tensor((), "int32"), ) -> None: for u in range(1): with Ts.sblock("C"): @@ -593,9 +593,9 @@ def zero_dim( @Ts.prim_func def zero_dim_added( - A: T.Buffer((), "int32"), - B: T.Buffer((), "int32"), - C: T.Buffer((), "int32"), + A: T.Tensor((), "int32"), + B: T.Tensor((), "int32"), + C: T.Tensor((), "int32"), ) -> None: for u1, u2 in T.grid(1, 1): with Ts.sblock("C"): @@ -676,7 +676,7 @@ def test_split_int64_factors(): def test_unsupported_target_scalable_split(): @Ts.prim_func - def before(A: T.Buffer((128,), "float32")): + def before(A: T.Tensor((128,), "float32")): T.func_attr({"global_symbol": "my_module", "tirx.noalias": True}) for i in T.serial(128): with Ts.sblock("A"): @@ -693,14 +693,14 @@ def before(A: T.Buffer((128,), "float32")): def test_fused_symbolic_2D_tiling(): @Ts.prim_func - def before(A: T.Buffer((M, N)), B: T.Buffer((M, N)), M: T.int32, N: T.int32) -> None: # noqa: F821 + def before(A: T.Tensor((M, N)), B: T.Tensor((M, N)), M: T.int32, N: T.int32) -> None: # noqa: F821 for i, j in T.grid(M, N): with Ts.sblock("B"): vi, vj = Ts.axis.remap("SS", [i, j]) B[vi, vj] = A[vi, vj] * 2.0 @Ts.prim_func - def expected(A: T.Buffer((M, N)), B: T.Buffer((M, N)), M: T.int32, N: T.int32) -> None: # noqa: F821 + def expected(A: T.Tensor((M, N)), B: T.Tensor((M, N)), M: T.int32, N: T.int32) -> None: # noqa: F821 for i_0_j_0_fused, i_1, j_1 in T.grid(((M + 63) // 64) * ((N + 15) // 16), 64, 16): with Ts.sblock("B"): vi = Ts.axis.spatial(M, i_0_j_0_fused // ((N + 15) // 16) * 64 + i_1) diff --git a/tests/python/s_tir/schedule/test_tir_schedule_state.py b/tests/python/s_tir/schedule/test_tir_schedule_state.py index 8495eddc6710..c11a334e63f0 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_state.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_state.py @@ -33,7 +33,7 @@ @Ts.prim_func -def elementwise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")) -> None: +def elementwise(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -46,7 +46,7 @@ def elementwise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "flo @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j in T.grid(128, 128): with Ts.sblock("init"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -59,7 +59,7 @@ def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 1 @Ts.prim_func def block_in_opaque_block( - A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32") ) -> None: for i in range(128): with Ts.sblock("B"): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_state_cached_flags.py b/tests/python/s_tir/schedule/test_tir_schedule_state_cached_flags.py index 04f6ce389ec5..996fe0cf9820 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_state_cached_flags.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_state_cached_flags.py @@ -32,7 +32,7 @@ # fmt: off @Ts.prim_func -def elementwise(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), 'float32')) -> None: +def elementwise(A: T.Tensor((128, 128), 'float32'), C: T.Tensor((128, 128), 'float32')) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): @@ -45,7 +45,7 @@ def elementwise(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), 'flo C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j in T.grid(128, 128): with Ts.sblock("init"): @@ -57,7 +57,7 @@ def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 1 C[vi, vj] = C[vi, vj] + A[vi, vk] * B[vj, vk] @Ts.prim_func -def block_in_opaque_block(A: T.Buffer((128, 128), 'float32'), B: T.Buffer((128, 128), 'float32')) -> None: +def block_in_opaque_block(A: T.Tensor((128, 128), 'float32'), B: T.Tensor((128, 128), 'float32')) -> None: for i in range(128): with Ts.sblock("B"): @@ -83,7 +83,7 @@ def block_in_opaque_block(A: T.Buffer((128, 128), 'float32'), B: T.Buffer((128, B[vi, vj] = A[vi, vj] * 2.0 @Ts.prim_func -def write_after_read(A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def write_after_read(A: T.Tensor((128, 128)), B: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): with Ts.sblock("C"): @@ -95,7 +95,7 @@ def write_after_read(A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buff B[vi, vj] = A[vi, vj] * 2.0 @Ts.prim_func -def loop_carried_dependency(A: T.Buffer((128,)), B: T.Buffer((128,)), C: T.Buffer((128,))) -> None: +def loop_carried_dependency(A: T.Tensor((128,)), B: T.Tensor((128,)), C: T.Tensor((128,))) -> None: for i in range(0, 128): with Ts.sblock("B"): @@ -106,7 +106,7 @@ def loop_carried_dependency(A: T.Buffer((128,)), B: T.Buffer((128,)), C: T.Buffe C[vi] = T.if_then_else(vi >= 1, B[vi - 1] + 1.0, 0.0) @Ts.prim_func -def concatenate_multi_producer(A: T.Buffer((128,)), B: T.Buffer((128,))) -> None: +def concatenate_multi_producer(A: T.Tensor((128,)), B: T.Tensor((128,))) -> None: for i in range(0, 64): with Ts.sblock("A_0"): @@ -122,7 +122,7 @@ def concatenate_multi_producer(A: T.Buffer((128,)), B: T.Buffer((128,))) -> None B[vi] = A[vi] * 2.0 @Ts.prim_func -def concatenate_multi_producer_uncovered(A: T.Buffer((128,)), B: T.Buffer((128,))) -> None: +def concatenate_multi_producer_uncovered(A: T.Tensor((128,)), B: T.Tensor((128,))) -> None: for i in range(0, 63): with Ts.sblock("A_0"): @@ -138,7 +138,7 @@ def concatenate_multi_producer_uncovered(A: T.Buffer((128,)), B: T.Buffer((128,) B[vi] = A[vi] * 2.0 @Ts.prim_func -def lca_at_loop(A: T.Buffer((128,)), B: T.Buffer((128,)), C: T.Buffer((128,))) -> None: +def lca_at_loop(A: T.Tensor((128,)), B: T.Tensor((128,)), C: T.Tensor((128,))) -> None: for i in range(0, 128): with Ts.sblock("B"): @@ -149,7 +149,7 @@ def lca_at_loop(A: T.Buffer((128,)), B: T.Buffer((128,)), C: T.Buffer((128,))) - C[vi] = B[vi] + 1.0 @Ts.prim_func -def multi_producer_consumer(A: T.Buffer((128,)), B: T.Buffer((128,))) -> None: +def multi_producer_consumer(A: T.Tensor((128,)), B: T.Tensor((128,))) -> None: for i in range(0, 64): with Ts.sblock("A_0"): @@ -169,7 +169,7 @@ def multi_producer_consumer(A: T.Buffer((128,)), B: T.Buffer((128,))) -> None: B[vi] = A[vi] + 3.0 @Ts.prim_func -def elementwise_affine_producer(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), 'float32')) -> None: +def elementwise_affine_producer(A: T.Tensor((128, 128), 'float32'), C: T.Tensor((128, 128), 'float32')) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j, k, l in T.grid(16, 2, 32, 16): @@ -183,7 +183,7 @@ def elementwise_affine_producer(A: T.Buffer((128, 128), 'float32'), C: T.Buffer( C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def elementwise_subblock(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), 'float32')) -> None: +def elementwise_subblock(A: T.Tensor((128, 128), 'float32'), C: T.Tensor((128, 128), 'float32')) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(32, 32): @@ -201,7 +201,7 @@ def elementwise_subblock(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 1 C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def elementwise_subblock_uncovered(A: T.Buffer((128, 128), 'float32'), C: T.Buffer((128, 128), 'float32')) -> None: +def elementwise_subblock_uncovered(A: T.Tensor((128, 128), 'float32'), C: T.Tensor((128, 128), 'float32')) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(32, 32): @@ -219,7 +219,7 @@ def elementwise_subblock_uncovered(A: T.Buffer((128, 128), 'float32'), C: T.Buff C[vi, vj] = B[vi, vj] + 1.0 @Ts.prim_func -def bound_to_thread(A: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def bound_to_thread(A: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: B = Ts.sblock_alloc_buffer([128, 128], scope="shared") for i in T.thread_binding(0, 128, thread="threadIdx.x"): @@ -233,7 +233,7 @@ def bound_to_thread(A: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: C[vj, vi] = B[vj, vi] + 1.0 @Ts.prim_func -def equal_ranked_threads(A: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def equal_ranked_threads(A: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: B = Ts.sblock_alloc_buffer([128, 128], scope="shared") for i_o in T.thread_binding(0, 16, thread="threadIdx.x"): @@ -250,7 +250,7 @@ def equal_ranked_threads(A: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> No C[vj, vi] = B[vj, vi] + 1.0 @Ts.prim_func -def warp_memory(A: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def warp_memory(A: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: B = Ts.sblock_alloc_buffer([128, 4, 32], scope="warp") for i_o in T.thread_binding(0, 4, thread="threadIdx.y"): @@ -265,7 +265,7 @@ def warp_memory(A: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: C[warp_id * 32 + lane_id, vj] = B[vj, warp_id, lane_id] + 1.0 @Ts.prim_func -def warp_memory_negative(A: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def warp_memory_negative(A: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: B = Ts.sblock_alloc_buffer([128, 4, 32], scope="warp") for i_o in T.thread_binding(0, 4, thread="threadIdx.y"): @@ -283,7 +283,7 @@ def warp_memory_negative(A: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> No C[warp_id * 32 + lane_id, vj] = B[vj, warp_id, lane_id] + 1.0 @Ts.prim_func -def non_perfect_tiling_cache(X: T.Buffer([224, 224], dtype='float32'), Y: T.Buffer([224, 224], dtype='float32')) -> None: +def non_perfect_tiling_cache(X: T.Tensor([224, 224], dtype='float32'), Y: T.Tensor([224, 224], dtype='float32')) -> None: cache = Ts.sblock_alloc_buffer([224, 224], dtype="float32") for hh_0, ww_0 in T.grid(28, 28): @@ -320,7 +320,7 @@ def non_perfect_tiling_cache(X: T.Buffer([224, 224], dtype='float32'), Y: T.Buff ) @Ts.prim_func -def uncovered_producer_region(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): +def uncovered_producer_region(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): for i in range(120): with Ts.sblock("producer"): vi = Ts.axis.S((0, 120), i) @@ -331,7 +331,7 @@ def uncovered_producer_region(A: T.Buffer((128,), "float32"), B: T.Buffer((128,) B[vi] = A[vi] @Ts.prim_func -def matmul_relu_padding(A: T.Buffer((127, 127), "float16"), B: T.Buffer((127, 127), "float16"), compute: T.Buffer((127, 127), "float32")) -> None: +def matmul_relu_padding(A: T.Tensor((127, 127), "float16"), B: T.Tensor((127, 127), "float16"), compute: T.Tensor((127, 127), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -403,7 +403,7 @@ def matmul_relu_padding(A: T.Buffer((127, 127), "float16"), B: T.Buffer((127, 12 @Ts.prim_func def splitted_square_sum_with_predicate( - A: T.Buffer((1, 7, 7, 512), "float32"), B: T.Buffer((1, 1, 1, 512), "float32") + A: T.Tensor((1, 7, 7, 512), "float32"), B: T.Tensor((1, 1, 1, 512), "float32") ) -> None: for i0_i1_i2_i3_0_fused, ax0, ax1, ax2, ax3 in T.grid(2, 1, 1, 1, 256): for ax4_ax5_fused_0, ax4_ax5_fused_1 in T.grid(1, 256): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_storage_align.py b/tests/python/s_tir/schedule/test_tir_schedule_storage_align.py index cb75da3be950..96b95c48e244 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_storage_align.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_storage_align.py @@ -33,8 +33,8 @@ @Ts.prim_func def element_wise( - A: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), - C: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), + A: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), + C: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), ) -> None: # body @@ -58,8 +58,8 @@ def element_wise( @Ts.prim_func def element_wise_storage_align( - A: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), - C: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), + A: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), + C: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), ) -> None: # body @@ -84,8 +84,8 @@ def element_wise_storage_align( @Ts.prim_func def element_wise_invalid_annotation( - A: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), - C: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), + A: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), + C: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), ) -> None: # body 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 37b703740cd1..1339ce7a5afd 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_tensorize.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_tensorize.py @@ -44,7 +44,7 @@ # pylint: disable=no-member,invalid-name,unused-variable,line-too-long,redefined-outer-name,unexpected-keyword-arg,too-many-nested-blocks @Ts.prim_func -def mma_desc(A: T.Buffer((16, 16), align=64, offset_factor=1), B: T.Buffer((16, 16), align=64, offset_factor=1), C: T.Buffer((16, 16), align=64, offset_factor=1)) -> None: +def mma_desc(A: T.Tensor((16, 16), align=64, offset_factor=1), B: T.Tensor((16, 16), align=64, offset_factor=1), C: T.Tensor((16, 16), align=64, offset_factor=1)) -> None: with Ts.sblock("root"): Ts.reads(C[0 : 16, 0 : 16], A[0 : 16, 0 : 16], B[0 : 16, 0 : 16]) @@ -55,7 +55,7 @@ def mma_desc(A: T.Buffer((16, 16), align=64, offset_factor=1), B: T.Buffer((16, C[vii, vjj] = C[vii, vjj] + A[vii, vkk] * B[vjj, vkk] @Ts.prim_func -def mma_intrin(A: T.Buffer((16, 16), align=64, offset_factor=1), B: T.Buffer((16, 16), align=64, offset_factor=1), C: T.Buffer((16, 16), align=64, offset_factor=1)) -> None: +def mma_intrin(A: T.Tensor((16, 16), align=64, offset_factor=1), B: T.Tensor((16, 16), align=64, offset_factor=1), C: T.Tensor((16, 16), align=64, offset_factor=1)) -> None: with Ts.sblock("root"): Ts.reads(C[0 : 16, 0 : 16], A[0 : 16, 0 : 16], B[0 : 16, 0 : 16]) @@ -75,7 +75,7 @@ def mma_intrin(A: T.Buffer((16, 16), align=64, offset_factor=1), B: T.Buffer((16 ) @Ts.prim_func -def dot_product_desc(A: T.Buffer((4,)), B: T.Buffer((4,)), C: T.Buffer(())) -> None: +def dot_product_desc(A: T.Tensor((4,)), B: T.Tensor((4,)), C: T.Tensor(())) -> None: with Ts.sblock("root"): Ts.reads(C[()], A[0 : 4], B[0 : 4]) @@ -86,7 +86,7 @@ def dot_product_desc(A: T.Buffer((4,)), B: T.Buffer((4,)), C: T.Buffer(())) -> N C[()] = C[()] + A[vi] * B[vi] @Ts.prim_func -def dot_product_intrin(A: T.Buffer((4,), offset_factor=1), B: T.Buffer((4,), offset_factor=1), C: T.Buffer((), offset_factor=1)) -> None: +def dot_product_intrin(A: T.Tensor((4,), offset_factor=1), B: T.Tensor((4,), offset_factor=1), C: T.Tensor((), offset_factor=1)) -> None: with Ts.sblock("root"): Ts.reads(C[()], A[0 : 4], B[0 : 4]) @@ -105,7 +105,7 @@ def dot_product_intrin(A: T.Buffer((4,), offset_factor=1), B: T.Buffer((4,), off ) @Ts.prim_func -def dot_product_intrin_annotated(A: T.Buffer((4,), offset_factor=1), B: T.Buffer((4,), offset_factor=1), C: T.Buffer((), offset_factor=1)) -> None: +def dot_product_intrin_annotated(A: T.Tensor((4,), offset_factor=1), B: T.Tensor((4,), offset_factor=1), C: T.Tensor((), offset_factor=1)) -> None: with Ts.sblock("root"): Ts.reads(C[()], A[0 : 4], B[0 : 4]) @@ -125,7 +125,7 @@ def dot_product_intrin_annotated(A: T.Buffer((4,), offset_factor=1), B: T.Buffer ) @Ts.prim_func -def outer_product_desc(A: T.Buffer((16, 1), offset_factor=1), B: T.Buffer((16, 1), offset_factor=1), C: T.Buffer((16, 16), offset_factor=1)) -> None: +def outer_product_desc(A: T.Tensor((16, 1), offset_factor=1), B: T.Tensor((16, 1), offset_factor=1), C: T.Tensor((16, 16), offset_factor=1)) -> None: with Ts.sblock("root"): Ts.reads( @@ -140,7 +140,7 @@ def outer_product_desc(A: T.Buffer((16, 1), offset_factor=1), B: T.Buffer((16, 1 C[vii, vjj] = C[vii, vjj] + A[vii, 0] * B[vjj, 0] @Ts.prim_func -def outer_product_intrin(A: T.Buffer((16, 1), offset_factor=1), B: T.Buffer((16, 1), offset_factor=1), C: T.Buffer((16, 16), offset_factor=1)) -> None: +def outer_product_intrin(A: T.Tensor((16, 1), offset_factor=1), B: T.Tensor((16, 1), offset_factor=1), C: T.Tensor((16, 16), offset_factor=1)) -> None: with Ts.sblock("root"): Ts.reads( @@ -164,9 +164,9 @@ def outer_product_intrin(A: T.Buffer((16, 1), offset_factor=1), B: T.Buffer((16, @Ts.prim_func def matmul( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): @@ -180,7 +180,7 @@ def matmul( C_elem_offset = T.dynamic("C_elem_offset", "int32") @Ts.prim_func -def tensorized_matmul(A: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), B: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), C: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1)) -> None: +def tensorized_matmul(A: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), B: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), C: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1)) -> None: for i_outer, j_outer in T.grid(8, 8): for i_inner_init, j_inner_init in T.grid(16, 16): @@ -230,9 +230,9 @@ def tensorized_matmul(A: T.Buffer([128, 128], elem_offset=0, align=64, offset_fa @Ts.prim_func def batch_matmul( - A: T.Buffer((16, 128, 128), "float32"), - B: T.Buffer((16, 128, 128), "float32"), - C: T.Buffer((16, 128, 128), "float32"), + A: T.Tensor((16, 128, 128), "float32"), + B: T.Tensor((16, 128, 128), "float32"), + C: T.Tensor((16, 128, 128), "float32"), ) -> None: for n, i, j in T.grid(16, 128, 128): with Ts.sblock("init"): @@ -250,9 +250,9 @@ def batch_matmul( @Ts.prim_func def tensorized_batch_matmul_mma( - A: T.Buffer((16, 128, 128), "float32"), - B: T.Buffer((16, 128, 128), "float32"), - C: T.Buffer((16, 128, 128), "float32"), + A: T.Tensor((16, 128, 128), "float32"), + B: T.Tensor((16, 128, 128), "float32"), + C: T.Tensor((16, 128, 128), "float32"), ) -> None: for n, i, j in T.grid(16, 128, 128): with Ts.sblock("init"): @@ -301,9 +301,9 @@ def tensorized_batch_matmul_mma( @Ts.prim_func def tensorized_batch_matmul_dot_product( - A: T.Buffer((16, 128, 128), "float32"), - B: T.Buffer((16, 128, 128), "float32"), - C: T.Buffer((16, 128, 128), "float32"), + A: T.Tensor((16, 128, 128), "float32"), + B: T.Tensor((16, 128, 128), "float32"), + C: T.Tensor((16, 128, 128), "float32"), ) -> None: for n, i, j in T.grid(16, 128, 128): with Ts.sblock("init"): @@ -340,9 +340,9 @@ def tensorized_batch_matmul_dot_product( @Ts.prim_func def tensorized_batch_matmul_outer_product( - A: T.Buffer((16, 128, 128), "float32"), - B: T.Buffer((16, 128, 128), "float32"), - C: T.Buffer((16, 128, 128), "float32"), + A: T.Tensor((16, 128, 128), "float32"), + B: T.Tensor((16, 128, 128), "float32"), + C: T.Tensor((16, 128, 128), "float32"), ) -> None: for n, i, j in T.grid(16, 128, 128): with Ts.sblock("init"): @@ -372,7 +372,7 @@ def tensorized_batch_matmul_outer_product( ) @Ts.prim_func -def annotated_mma_desc(A: T.Buffer((16, 16), align=64, offset_factor=1), B: T.Buffer((16, 16), align=64, offset_factor=1), C: T.Buffer((16, 16), align=64, offset_factor=1)) -> None: +def annotated_mma_desc(A: T.Tensor((16, 16), align=64, offset_factor=1), B: T.Tensor((16, 16), align=64, offset_factor=1), C: T.Tensor((16, 16), align=64, offset_factor=1)) -> None: with Ts.sblock("root"): Ts.reads(C[0 : 16, 0 : 16], A[0 : 16, 0 : 16], B[0 : 16, 0 : 16]) @@ -385,9 +385,9 @@ def annotated_mma_desc(A: T.Buffer((16, 16), align=64, offset_factor=1), B: T.Bu @Ts.prim_func def annotated_matmul( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): @@ -402,7 +402,7 @@ def annotated_matmul( C_elem_offset = T.dynamic("C_elem_offset", "int32") @Ts.prim_func -def annotated_tensorized_matmul(A: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), B: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), C: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1)) -> None: +def annotated_tensorized_matmul(A: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), B: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), C: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1)) -> None: for i_outer, j_outer in T.grid(8, 8): for i_inner_init, j_inner_init in T.grid(16, 16): @@ -700,9 +700,9 @@ def test_tensorize_matmul_mixed_dtype(): # fmt: off @Ts.prim_func def matmul_int64_shape( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(128)), "float32"), - C: T.Buffer((T.int64(128), T.int64(128)), "float32") + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(128)), "float32"), + C: T.Tensor((T.int64(128), T.int64(128)), "float32") ) -> None: for i_0, j_0 in T.grid(T.int64(8), T.int64(8)): for i_1_init, j_1_init in T.grid(T.int64(16), T.int64(16)): @@ -723,9 +723,9 @@ def matmul_int64_shape( @Ts.prim_func def tensorized_matmul_int64_shape( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(128)), "float32"), - C: T.Buffer((T.int64(128), T.int64(128)), "float32") + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(128)), "float32"), + C: T.Tensor((T.int64(128), T.int64(128)), "float32") ) -> None: for i_outer, j_outer in T.grid(T.int64(8), T.int64(8)): for i_inner_init, j_inner_init in T.grid(T.int64(16), T.int64(16)): @@ -793,7 +793,7 @@ def f_convert(nbit: int, val: tirx.Expr, pos: tirx.Expr, dtype: str): return f_convert @Ts.prim_func -def decode_i4s_to_f16_desc(Compressed: T.Buffer([1], dtype='int32', scope='local'), Decompressed: T.Buffer([8], dtype='float16', scope='local')) -> None: +def decode_i4s_to_f16_desc(Compressed: T.Tensor([1], dtype='int32', scope='local'), Decompressed: T.Tensor([8], dtype='float16', scope='local')) -> None: with Ts.sblock("root"): Ts.reads(Compressed[0:1]) @@ -809,7 +809,7 @@ def decode_i4s_to_f16_desc(Compressed: T.Buffer([1], dtype='int32', scope='local ) @Ts.prim_func -def decode_i4s_to_f16_impl(Compressed: T.Buffer([1], dtype='int32', scope='local'), Decompressed: T.Buffer([8], dtype='float16', scope='local')) -> None: +def decode_i4s_to_f16_impl(Compressed: T.Tensor([1], dtype='int32', scope='local'), Decompressed: T.Tensor([8], dtype='float16', scope='local')) -> None: with Ts.sblock("root"): Ts.reads(Compressed[0:1]) diff --git a/tests/python/s_tir/schedule/test_tir_schedule_trace.py b/tests/python/s_tir/schedule/test_tir_schedule_trace.py index 2ec7aed5f5b6..0f054e757537 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_trace.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_trace.py @@ -33,7 +33,7 @@ @Ts.prim_func -def elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): @@ -47,7 +47,7 @@ def elementwise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: @Ts.prim_func -def elementwise_inlined(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def elementwise_inlined(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: for i, j in T.grid(128, 128): with Ts.sblock("C"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -362,7 +362,7 @@ def _test_apply_annotation_trace_from_json(annotation: str): Trace.apply_json_to_schedule(json_obj, sch) @Ts.prim_func - def elementwise_expected(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: + def elementwise_expected(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: B = Ts.sblock_alloc_buffer((128, 128)) for i, j in T.grid(128, 128): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_transform.py b/tests/python/s_tir/schedule/test_tir_schedule_transform.py index ad6979bb0a4d..43971e9fe3d2 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_transform.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_transform.py @@ -26,9 +26,9 @@ class DenseTIRModule: @Ts.prim_func def main( - placeholder: T.Buffer((1024, 1024), "uint8"), - placeholder_1: T.Buffer((64, 256, 16, 4), "int8"), - compute: T.Buffer((1024, 1024), "int32"), + placeholder: T.Tensor((1024, 1024), "uint8"), + placeholder_1: T.Tensor((64, 256, 16, 4), "int8"), + compute: T.Tensor((1024, 1024), "int32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): @@ -50,9 +50,9 @@ def main( class DenseTIRModuleTiled: @Ts.prim_func def main( - placeholder: T.Buffer((1024, 1024), "uint8"), - placeholder_1: T.Buffer((64, 256, 16, 4), "int8"), - compute: T.Buffer((1024, 1024), "int32"), + placeholder: T.Tensor((1024, 1024), "uint8"), + placeholder_1: T.Tensor((64, 256, 16, 4), "int8"), + compute: T.Tensor((1024, 1024), "int32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -76,9 +76,9 @@ def main( class Conv2dNCHWcTIRModule: @Ts.prim_func def main( - placeholder: T.Buffer((1, 4, 56, 56, 16), "uint8"), - placeholder_1: T.Buffer((16, 4, 1, 1, 4, 16, 4), "int8"), - conv2d_NCHWc_int8: T.Buffer((1, 16, 56, 56, 16), "int32"), + placeholder: T.Tensor((1, 4, 56, 56, 16), "uint8"), + placeholder_1: T.Tensor((16, 4, 1, 1, 4, 16, 4), "int8"), + conv2d_NCHWc_int8: T.Tensor((1, 16, 56, 56, 16), "int32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0, i1, i2, i3, i4, i5, i6, i7, i8, i9 in T.grid(1, 16, 56, 56, 16, 1, 1, 4, 4, 4): @@ -117,9 +117,9 @@ def main( class Conv2dNCHWcTIRModuleTiled: @Ts.prim_func def main( - placeholder: T.Buffer((1, 4, 56, 56, 16), "uint8"), - placeholder_1: T.Buffer((16, 4, 1, 1, 4, 16, 4), "int8"), - conv2d_NCHWc_int8: T.Buffer((1, 16, 56, 56, 16), "int32"), + placeholder: T.Tensor((1, 4, 56, 56, 16), "uint8"), + placeholder_1: T.Tensor((16, 4, 1, 1, 4, 16, 4), "int8"), + conv2d_NCHWc_int8: T.Tensor((1, 16, 56, 56, 16), "int32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) 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 5060d47dd62e..0c0beb96cbdb 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 @@ -39,7 +39,7 @@ def packed_index_map_func(m, n): return m // 16, n // 16, m % 16, n % 16 @Ts.prim_func -def two_elementwise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32")) -> None: +def two_elementwise(A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32")) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): with Ts.sblock("B"): @@ -52,7 +52,7 @@ def two_elementwise(A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), @Ts.prim_func def two_elementwise_transformed_intermediate_buffer( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((8, 8, 16, 16), "float32") for i, j in T.grid(128, 128): @@ -66,7 +66,7 @@ def two_elementwise_transformed_intermediate_buffer( @Ts.prim_func def two_elementwise_transformed_input_buffer( - A: T.Buffer((8, 8, 16, 16), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((8, 8, 16, 16), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): @@ -80,7 +80,7 @@ def two_elementwise_transformed_input_buffer( @Ts.prim_func def two_elementwise_transformed_output_buffer( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((8, 8, 16, 16), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((8, 8, 16, 16), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") for i, j in T.grid(128, 128): @@ -93,14 +93,14 @@ def two_elementwise_transformed_output_buffer( C[vi // 16, vj // 16, vi % 16, vj % 16] = B[vi, vj] + 1.0 @Ts.prim_func -def elementwise(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32")) -> None: +def elementwise(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32")) -> None: for i, j in T.grid(128, 128): with Ts.sblock("B"): vi, vj = Ts.axis.remap("SS", [i, j]) B[vi, vj] = A[vi, vj] * 2.0 @Ts.prim_func -def elementwise_transformed(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32")) -> None: +def elementwise_transformed(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32")) -> None: for i in range(16384): with Ts.sblock("B"): vi = Ts.axis.remap("S", [i]) @@ -108,9 +108,9 @@ def elementwise_transformed(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128 @Ts.prim_func def conv2d_nhwc( - Input: T.Buffer((1, 224, 224, 3), "float32"), - Weight: T.Buffer((7, 7, 3, 64), "float32"), - Conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32"), + Input: T.Tensor((1, 224, 224, 3), "float32"), + Weight: T.Tensor((7, 7, 3, 64), "float32"), + Conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32"), ) -> None: PadInput = Ts.sblock_alloc_buffer([1, 230, 230, 3], dtype="float32") for i0, i1, i2, i3 in T.grid(1, 230, 230, 3): @@ -134,9 +134,9 @@ def conv2d_nhwc( @Ts.prim_func def conv2d_nhwc_transformed( - Input: T.Buffer((1, 224, 224, 3), "float32"), - Weight: T.Buffer((7, 7, 3, 64), "float32"), - Conv2d_nhwc: T.Buffer((1, 112, 112, 64), "float32"), + Input: T.Tensor((1, 224, 224, 3), "float32"), + Weight: T.Tensor((7, 7, 3, 64), "float32"), + Conv2d_nhwc: T.Tensor((1, 112, 112, 64), "float32"), ) -> None: PadInput = Ts.sblock_alloc_buffer([1, 230, 230, 3], dtype="float32") for i0, i1, i2, i3 in T.grid(1, 230, 230, 3): @@ -158,7 +158,7 @@ def conv2d_nhwc_transformed( Conv2d_nhwc[0, v0 // 112, v0 % 112, v1] = Conv2d_nhwc[0, v0 // 112, v0 % 112, v1] + PadInput[0, v0 // 112 * 2 + v2 // 21, v0 % 112 * 2 + v2 % 21 // 3, v2 % 3] * Weight[v2 // 21, v2 % 21 // 3, v2 % 3, v1] @Ts.prim_func -def two_elementwise_unit_dim(A: T.Buffer((1, 128), "float32"), C: T.Buffer((1, 128), "float32")) -> None: +def two_elementwise_unit_dim(A: T.Tensor((1, 128), "float32"), C: T.Tensor((1, 128), "float32")) -> None: B = Ts.sblock_alloc_buffer((1, 128), "float32") for i, j in T.grid(1, 128): with Ts.sblock("B"): @@ -279,7 +279,7 @@ def test_simplify(): sch.transform_layout(B, ("write", 0), lambda i, j: (i // 16, j // 16, i % 16, j % 16)) @Ts.prim_func - def ref(B: T.Buffer((8, 8, 16, 16), "float32"), C: T.Buffer((128, 128), "float32")): + def ref(B: T.Tensor((8, 8, 16, 16), "float32"), C: T.Tensor((128, 128), "float32")): for i_0, j_0 in T.grid(8, 8): with Ts.sblock("C_o"): vi_o, vj_o = Ts.axis.remap("SS", [i_0, j_0]) @@ -310,7 +310,7 @@ def ref(B: T.Buffer((8, 8, 16, 16), "float32"), C: T.Buffer((128, 128), "float32 def test_var_args_sugar(): @Ts.prim_func def summation_3d( - A: T.Buffer((1024, 1024, 32), "float32"), B: T.Buffer((1,), "float32") + A: T.Tensor((1024, 1024, 32), "float32"), B: T.Tensor((1,), "float32") ) -> None: B[0] = 0 for i, j, k in T.grid(1024, 1024, 32): @@ -320,7 +320,7 @@ def summation_3d( @Ts.prim_func def summation_3d_split( - A: T.Buffer((1024, 1024, 8, 4), "float32"), B: T.Buffer((1,), "float32") + A: T.Tensor((1024, 1024, 8, 4), "float32"), B: T.Tensor((1,), "float32") ) -> None: B[0] = 0 for i, j, k in T.grid(1024, 1024, 32): @@ -361,7 +361,7 @@ def test_transform_block_layout_unit_dim(use_block_name): @Ts.prim_func def two_elementwise_unit_dim_transformed( - A: T.Buffer((1, 128), "float32"), C: T.Buffer((1, 128), "float32") + A: T.Tensor((1, 128), "float32"), C: T.Tensor((1, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((1, 128), "float32") for j, i in T.grid(128, 1): @@ -423,8 +423,8 @@ def detect(var): def test_transform_block_layout_int64_extent(use_block_name): @Ts.prim_func def elementwise_int64_extent( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(128)), "float32"), ) -> None: for i, j in T.grid(T.int64(128), T.int64(128)): with Ts.sblock("B"): @@ -433,8 +433,8 @@ def elementwise_int64_extent( @Ts.prim_func def elementwise_int64_extent_transformed( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(128)), "float32"), ) -> None: # T.serial with explicit int64 min so the iter_var dom is all-int64 # (`range(T.int64(...))` would emit an int32 min). @@ -696,7 +696,7 @@ def test_padded_transform_if_then_else(dtype): """ @Ts.prim_func(private=True) - def before_func(A: T.Buffer(14, dtype)): + def before_func(A: T.Tensor(14, dtype)): B = Ts.sblock_alloc_buffer(14, dtype) for i in T.serial(14): with Ts.sblock("block"): @@ -706,7 +706,7 @@ def before_func(A: T.Buffer(14, dtype)): pad_value_imm = tirx.IntImm(dtype, 0) @Ts.prim_func(private=True) - def expected_func(A: T.Buffer(14, dtype)): + def expected_func(A: T.Tensor(14, dtype)): B = Ts.sblock_alloc_buffer([4, 4], dtype) for i, j in T.grid(4, 4): with Ts.sblock("block"): @@ -737,7 +737,7 @@ def test_padded_transform_without_loop(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(14, "int32")): + def main(A: T.Tensor(14, "int32")): with Ts.sblock("root"): Ts.reads() Ts.writes() @@ -747,7 +747,7 @@ def main(A: T.Buffer(14, "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((4, 4), "int32")): + def main(A: T.Tensor((4, 4), "int32")): with Ts.sblock("block"): A[0, 0] = 0 @@ -774,7 +774,7 @@ def test_padded_transform_if_then_else_reduction(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((14, 32), "int32")): + def main(A: T.Tensor((14, 32), "int32")): B = Ts.sblock_alloc_buffer(14, "int32") for i, k in T.grid(14, 32): with Ts.sblock("block"): @@ -786,7 +786,7 @@ def main(A: T.Buffer((14, 32), "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((14, 32), "int32")): + def main(A: T.Tensor((14, 32), "int32")): B = Ts.sblock_alloc_buffer([4, 4], "int32") for i, j, k in T.grid(4, 4, 32): with Ts.sblock("block"): @@ -814,7 +814,7 @@ def test_padded_transform_if_then_else_reduction_opaque(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((14, 32), "int32")): + def main(A: T.Tensor((14, 32), "int32")): B = Ts.sblock_alloc_buffer(14, "int32") for i in T.serial(14): B[i] = 0 @@ -825,7 +825,7 @@ def main(A: T.Buffer((14, 32), "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((14, 32), "int32")): + def main(A: T.Tensor((14, 32), "int32")): B = Ts.sblock_alloc_buffer([4, 4], "int32") for i, j in T.grid(4, 4): B[i, j] = T.if_then_else(i == 3 and 2 <= j, 0, 0) @@ -854,7 +854,7 @@ def test_padded_transform_post_proc_if_required_due_to_side_effects(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(14, "int32")): + def main(A: T.Tensor(14, "int32")): B = Ts.sblock_alloc_buffer(14, "int32") C = Ts.sblock_alloc_buffer(14, "int32") for i in T.serial(14): @@ -866,7 +866,7 @@ def main(A: T.Buffer(14, "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer(14, "int32")): + def main(A: T.Tensor(14, "int32")): B = Ts.sblock_alloc_buffer([4, 4], "int32") C = Ts.sblock_alloc_buffer(14, "int32") for i in T.serial(14): @@ -899,7 +899,7 @@ def test_padded_transform_of_input_creates_assumption(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(14, "int32"), B: T.Buffer(14, "int32")): + def main(A: T.Tensor(14, "int32"), B: T.Tensor(14, "int32")): for i in T.serial(14): with Ts.sblock("block"): vi = Ts.axis.remap("S", [i]) @@ -908,7 +908,7 @@ def main(A: T.Buffer(14, "int32"), B: T.Buffer(14, "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((4, 4), "int32"), B: T.Buffer(14, "int32")): + def main(A: T.Tensor((4, 4), "int32"), B: T.Tensor(14, "int32")): for i, j in T.grid(4, 4): with Ts.sblock("buffer_A_assumption"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -940,7 +940,7 @@ def test_padded_transform_non_constant_value(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(14, "int32")): + def main(A: T.Tensor(14, "int32")): B = Ts.sblock_alloc_buffer(14, "int32") for i in T.serial(14): with Ts.sblock("block"): @@ -950,7 +950,7 @@ def main(A: T.Buffer(14, "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer(14, "int32")): + def main(A: T.Tensor(14, "int32")): B = Ts.sblock_alloc_buffer([4, 4], "int32") for i, j in T.grid(4, 4): with Ts.sblock("block"): @@ -981,7 +981,7 @@ def test_padded_transform_repeated_buffer_element(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(14, "int32")): + def main(A: T.Tensor(14, "int32")): B = Ts.sblock_alloc_buffer(14, "int32") for i in T.serial(14): with Ts.sblock("block"): @@ -991,7 +991,7 @@ def main(A: T.Buffer(14, "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((4, 4), "int32")): + def main(A: T.Tensor((4, 4), "int32")): for i, j in T.grid(4, 4): with Ts.sblock("buffer_A_assumption"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -1030,7 +1030,7 @@ def test_pad_value_may_not_reference_other_buffer(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(14, "int32")): + def main(A: T.Tensor(14, "int32")): B = Ts.sblock_alloc_buffer(14, "int32") for i in T.serial(14): with Ts.sblock("block"): @@ -1039,7 +1039,7 @@ def main(A: T.Buffer(14, "int32")): sch = tvm.s_tir.Schedule(Before) A = sch.get(sch.get_sblock("block")).reads[0].source - other = tirx.decl_buffer(1, A.ty.dtype, name="other") + other = tirx.decl_tensor(1, A.ty.dtype, name="other") with pytest.raises(tvm.s_tir.schedule.schedule.ScheduleError): sch.transform_layout( "block", @@ -1055,7 +1055,7 @@ def test_transform_layout_with_var(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(16, "int32"), n: T.int32): + def main(A: T.Tensor(16, "int32"), n: T.int32): B = Ts.sblock_alloc_buffer(16, "int32") for i in T.serial(16): with Ts.sblock("block"): @@ -1065,7 +1065,7 @@ def main(A: T.Buffer(16, "int32"), n: T.int32): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer(16, "int32"), n: T.int32): + def main(A: T.Tensor(16, "int32"), n: T.int32): B = Ts.sblock_alloc_buffer([(-16 % n + 16) // n, n], dtype="int32") for i, j in T.grid((-16 % n + 16) // n, n): with Ts.sblock("block"): @@ -1098,7 +1098,7 @@ def test_index_map_dtype_legalize(): """Test dtype legalization of the index map indices.""" @Ts.prim_func - def func(A: T.Buffer(T.int64(58), "int32")): + def func(A: T.Tensor(T.int64(58), "int32")): for i in T.serial(T.int64(58)): with Ts.sblock("block"): vi = Ts.axis.remap("S", [i]) @@ -1122,7 +1122,7 @@ def test_index_map_dtype_legalize_with_constant(): """ @Ts.prim_func - def func(A: T.Buffer(T.int64(16), "int32")): + def func(A: T.Tensor(T.int64(16), "int32")): for i in T.grid(T.int64(16)): with Ts.sblock("block"): vi = Ts.axis.remap("S", [i]) @@ -1157,7 +1157,7 @@ def test_transform_layout_with_symbolic_bound(): n = T.dynamic("n") @Ts.prim_func - def before(A: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), 'float16'), B: T.Buffer((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), C: T.Buffer((T.int64(1), T.int64(32), T.int64(1), n), 'float16')): + def before(A: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), 'float16'), B: T.Tensor((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), C: T.Tensor((T.int64(1), T.int64(32), T.int64(1), n), 'float16')): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0, i1, i2, i3, k in T.grid(T.int64(1), T.int64(32), T.int64(1), n, T.int64(128)): @@ -1172,7 +1172,7 @@ def before(A: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), 'flo n = T.dynamic("n") @Ts.prim_func - def after(A: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), 'float16'), B: T.Buffer((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), C: T.Buffer((n * T.int64(32),), 'float16')): + def after(A: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), 'float16'), B: T.Tensor((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), C: T.Tensor((n * T.int64(32),), 'float16')): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0, i1, i2, i3, k in T.grid(T.int64(1), T.int64(32), T.int64(1), n, T.int64(128)): @@ -1205,7 +1205,7 @@ def test_transform_block_layout_with_symbolic_bound(): n = T.dynamic("n") @Ts.prim_func - def before(A: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), 'float16'), B: T.Buffer((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), C: T.Buffer((n * T.int64(32),), 'float16')): + def before(A: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), 'float16'), B: T.Tensor((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), C: T.Tensor((n * T.int64(32),), 'float16')): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0, i1, i2, i3, k in T.grid(T.int64(1), T.int64(32), T.int64(1), n, T.int64(128)): @@ -1220,7 +1220,7 @@ def before(A: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), 'flo n = T.dynamic("n") @Ts.prim_func - def after(A: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), 'float16'), B: T.Buffer((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), C: T.Buffer((n * T.int64(32),), 'float16')): + def after(A: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), 'float16'), B: T.Tensor((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), C: T.Tensor((n * T.int64(32),), 'float16')): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for ax0, ax1 in T.grid(n * T.int64(32), T.int64(128)): diff --git a/tests/python/s_tir/schedule/test_tir_schedule_utilities.py b/tests/python/s_tir/schedule/test_tir_schedule_utilities.py index 86b200d779ab..9c4e545d732a 100644 --- a/tests/python/s_tir/schedule/test_tir_schedule_utilities.py +++ b/tests/python/s_tir/schedule/test_tir_schedule_utilities.py @@ -35,7 +35,7 @@ @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j in T.grid(128, 128): with Ts.sblock("init"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -48,7 +48,7 @@ def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 1 @Ts.prim_func def matmul_relu( - A: T.Buffer((1024, 1024)), B: T.Buffer((1024, 1024)), D: T.Buffer((1024, 1024)) + A: T.Tensor((1024, 1024)), B: T.Tensor((1024, 1024)), D: T.Tensor((1024, 1024)) ) -> None: C = Ts.sblock_alloc_buffer((1024, 1024)) @@ -66,7 +66,7 @@ def matmul_relu( @Ts.prim_func def matmul_relu_ann1( - A: T.Buffer((1024, 1024)), B: T.Buffer((1024, 1024)), D: T.Buffer((1024, 1024)) + A: T.Tensor((1024, 1024)), B: T.Tensor((1024, 1024)), D: T.Tensor((1024, 1024)) ) -> None: C = Ts.sblock_alloc_buffer((1024, 1024)) @@ -86,7 +86,7 @@ def matmul_relu_ann1( @Ts.prim_func def matmul_relu_ann2( - A: T.Buffer((1024, 1024)), B: T.Buffer((1024, 1024)), D: T.Buffer((1024, 1024)) + A: T.Tensor((1024, 1024)), B: T.Tensor((1024, 1024)), D: T.Tensor((1024, 1024)) ) -> None: C = Ts.sblock_alloc_buffer((1024, 1024)) @@ -108,8 +108,8 @@ def matmul_relu_ann2( class ModuleWithMultipleFuncs: @Ts.prim_func def vector_add( - A: T.Buffer(128, "float32"), - B: T.Buffer(128, "float32"), + A: T.Tensor(128, "float32"), + B: T.Tensor(128, "float32"), ) -> None: for i in range(128): with Ts.sblock("init"): @@ -118,8 +118,8 @@ def vector_add( @Ts.prim_func def vector_add_2( - A: T.Buffer(128, "float32"), - B: T.Buffer(128, "float32"), + A: T.Tensor(128, "float32"), + B: T.Tensor(128, "float32"), ) -> None: for i in range(128): with Ts.sblock("init"): @@ -128,7 +128,7 @@ def vector_add_2( @Ts.prim_func -def tuple_reduction(data: T.Buffer((4, 32), "float32"), T_add: T.Buffer((4,), "float32")) -> None: +def tuple_reduction(data: T.Tensor((4, 32), "float32"), T_add: T.Tensor((4,), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # body @@ -389,8 +389,8 @@ def test_get_output_blocks_multiple_outputs(): def test_get_output_blocks_nested(): @Ts.prim_func def blockized( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), ) -> None: with Ts.sblock("blockized_B"): vio = Ts.axis.spatial(1, 0) diff --git a/tests/python/s_tir/schedule/test_tir_unsafe_hide_buffer_access.py b/tests/python/s_tir/schedule/test_tir_unsafe_hide_buffer_access.py index d4fd8a419436..a81c2701b137 100644 --- a/tests/python/s_tir/schedule/test_tir_unsafe_hide_buffer_access.py +++ b/tests/python/s_tir/schedule/test_tir_unsafe_hide_buffer_access.py @@ -31,10 +31,10 @@ @Ts.prim_func def indirect_mem_access( - A: T.Buffer([128], dtype="float32"), - IA: T.Buffer([10], dtype="int32"), - B: T.Buffer([128], dtype="float32"), - IB: T.Buffer([10], dtype="int32"), + A: T.Tensor([128], dtype="float32"), + IA: T.Tensor([10], dtype="int32"), + B: T.Tensor([128], dtype="float32"), + IB: T.Tensor([10], dtype="int32"), ) -> None: for i in range(10): with Ts.sblock("B"): @@ -46,10 +46,10 @@ def indirect_mem_access( @Ts.prim_func def indirect_mem_access_hide_ia( - A: T.Buffer([128], dtype="float32"), - IA: T.Buffer([10], dtype="int32"), - B: T.Buffer([128], dtype="float32"), - IB: T.Buffer([10], dtype="int32"), + A: T.Tensor([128], dtype="float32"), + IA: T.Tensor([10], dtype="int32"), + B: T.Tensor([128], dtype="float32"), + IB: T.Tensor([10], dtype="int32"), ) -> None: for i in range(10): with Ts.sblock("B"): @@ -61,10 +61,10 @@ def indirect_mem_access_hide_ia( @Ts.prim_func def indirect_mem_access_hide_ib( - A: T.Buffer([128], dtype="float32"), - IA: T.Buffer([10], dtype="int32"), - B: T.Buffer([128], dtype="float32"), - IB: T.Buffer([10], dtype="int32"), + A: T.Tensor([128], dtype="float32"), + IA: T.Tensor([10], dtype="int32"), + B: T.Tensor([128], dtype="float32"), + IB: T.Tensor([10], dtype="int32"), ) -> None: for i in range(10): with Ts.sblock("B"): diff --git a/tests/python/s_tir/script/test_s_tir_script_basic_usage.py b/tests/python/s_tir/script/test_s_tir_script_basic_usage.py index 69594de37063..9e23b2791e1a 100644 --- a/tests/python/s_tir/script/test_s_tir_script_basic_usage.py +++ b/tests/python/s_tir/script/test_s_tir_script_basic_usage.py @@ -34,10 +34,10 @@ @Ts.prim_func def get_valid_counts( - data_buf: T.Buffer((1, 2500, 6), "float32"), - valid_count_buf: T.Buffer((1,), "int32"), - out_buf: T.Buffer((1, 2500, 6), "float32"), - out_indices_buf: T.Buffer((1, 2500), "int32"), + data_buf: T.Tensor((1, 2500, 6), "float32"), + valid_count_buf: T.Tensor((1,), "int32"), + out_buf: T.Tensor((1, 2500, 6), "float32"), + out_indices_buf: T.Tensor((1, 2500), "int32"), score_threshold: T.float32, id_index: T.int32, score_index: T.int32, @@ -110,7 +110,7 @@ def test_get_valid_counts_script_func(): @Ts.prim_func -def ceildiv_test(A: T.Buffer(16, "int32")): +def ceildiv_test(A: T.Tensor(16, "int32")): for i in range(16): A[i] = T.ceildiv(A[i], 4) @@ -126,7 +126,7 @@ def test_ceildiv(): def test_tir_func_name(): @Ts.prim_func - def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: + def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) @@ -138,7 +138,7 @@ def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 1 def test_tir_func_private_attrs(): @Ts.prim_func(private=True) - def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: + def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: T.func_attr({"attr": "value"}) for i, j, k in T.grid(128, 128, 128): @@ -154,7 +154,7 @@ def test_tir_loop_steps(): @Ts.prim_func(private=True) def loop_with_steps( - A: T.Buffer((N,)), B: T.Buffer((N,)), C: T.Buffer((N,)), tid: T.int32, v: T.int32 + A: T.Tensor((N,)), B: T.Tensor((N,)), C: T.Tensor((N,)), tid: T.int32, v: T.int32 ): for i in T.serial(tid, N, step=2): C[i] = A[i] + B[i] @@ -181,11 +181,11 @@ def bar(val): T.evaluate(val) @Ts.prim_func(private=True) - def func_with_empty_tuple(A: T.Buffer((), "int32"), B: T.Buffer((), "int32")): + def func_with_empty_tuple(A: T.Tensor((), "int32"), B: T.Tensor((), "int32")): bar(val=A[()]) @Ts.prim_func(private=True) - def expected(A: T.Buffer((), "int32"), B: T.Buffer((), "int32")): + def expected(A: T.Tensor((), "int32"), B: T.Tensor((), "int32")): T.evaluate(A[()]) tvm.ir.assert_structural_equal(func_with_empty_tuple, expected) @@ -193,7 +193,7 @@ def expected(A: T.Buffer((), "int32"), B: T.Buffer((), "int32")): def test_thread_binding_dtype(): @Ts.prim_func(private=True) - def func(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))): + def func(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))): for i in T.thread_binding(T.int64(128), "threadIdx.x"): for j in T.thread_binding(128, "threadIdx.y"): B[i, j] = A[i, j] @@ -228,7 +228,7 @@ def test_inferred_ty_with_buffer_args(): """PrimFunc buffer arguments are inferred as R.Tensor""" @Ts.prim_func - def func(A: T.Buffer([16, 16], "float32"), B: T.Buffer([256], "int32")) -> T.float32: + def func(A: T.Tensor([16, 16], "float32"), B: T.Tensor([256], "int32")) -> T.float32: return T.float32(42.0) expected = tvm.relax.FuncType( @@ -250,8 +250,8 @@ def test_inferred_ty_with_internal_allocation(): """ @Ts.prim_func - def func(A: T.Buffer([16, 16], "float32")) -> T.float32: - Sum = T.decl_buffer([], "float32") + def func(A: T.Tensor([16, 16], "float32")) -> T.float32: + Sum = T.decl_tensor([], "float32") Sum[()] = 0.0 for i, j in T.grid(16, 16): Sum[()] = Sum[()] + A[i, j] @@ -275,7 +275,7 @@ def test_inferred_ty_with_output_buffer(): """ @Ts.prim_func - def func(A: T.Buffer(16, "float32"), B: T.Buffer(16, "float32")): + def func(A: T.Tensor(16, "float32"), B: T.Tensor(16, "float32")): for i in range(16): B[i] = A[i] @@ -294,7 +294,7 @@ def test_reinterpret_nop(): """Test builtin reinterpret op""" @Ts.prim_func - def func(A: T.Buffer((32,), "float32"), B: T.Buffer((32,), "float32")) -> None: + def func(A: T.Tensor((32,), "float32"), B: T.Tensor((32,), "float32")) -> None: T.func_attr({"global_symbol": "main"}) for i in T.serial(0, 32): with Ts.sblock(): @@ -302,7 +302,7 @@ def func(A: T.Buffer((32,), "float32"), B: T.Buffer((32,), "float32")) -> None: B[vi] = T.reinterpret("float32", A[vi]) @Ts.prim_func - def expected(A: T.Buffer((32,), "float32"), B: T.Buffer((32,), "float32")) -> None: + def expected(A: T.Tensor((32,), "float32"), B: T.Tensor((32,), "float32")) -> None: T.func_attr({"global_symbol": "main"}) for i in T.serial(0, 32): with Ts.sblock(): @@ -410,7 +410,7 @@ def expected(A: T.buffer((128, 128), "float32")): def test_sequence_compare(): @Ts.prim_func(private=True) - def tir_func(A: T.Buffer((128, 128), "float32")): + def tir_func(A: T.Tensor((128, 128), "float32")): for i, j in T.grid(128, 128): if 0 < i < 128 and 0 < j < 128: A[i, j] = 1 @@ -430,7 +430,7 @@ def expected(A: T.buffer((128, 128), "float32")): def launch_env_thread(): @Ts.prim_func - def main(inputs: T.Buffer((64, 2, 4), "float32")) -> None: + def main(inputs: T.Tensor((64, 2, 4), "float32")) -> None: bx = T.launch_thread("blockIdx.x", 64) for i, j in T.grid(2, 4): T.evaluate(inputs[bx, i, j]) @@ -440,7 +440,7 @@ def main(inputs: T.Buffer((64, 2, 4), "float32")) -> None: def vthread_func(): @Ts.prim_func - def vthread_func(A: T.Buffer([256], "float32"), C: T.Buffer([256], "float32")) -> None: + def vthread_func(A: T.Tensor([256], "float32"), C: T.Tensor([256], "float32")) -> None: i0 = T.env_thread("blockIdx.x") i1 = T.env_thread("threadIdx.x") i2 = T.env_thread("vthread") @@ -448,7 +448,7 @@ def vthread_func(A: T.Buffer([256], "float32"), C: T.Buffer([256], "float32")) - T.launch_thread(i0, 4) T.launch_thread(i1, 2) T.launch_thread(i2, 2) - B = T.alloc_buffer((16,), scope="local") + B = T.alloc_tensor((16,), scope="local") for j in range(16): B[j] = A[i0 * 64 + i1 * 32 + i2 * 16 + j] + T.float32(1) for j in range(16): @@ -460,7 +460,7 @@ def vthread_func(A: T.Buffer([256], "float32"), C: T.Buffer([256], "float32")) - def for_thread_binding(): @Ts.prim_func def for_thread_binding( - A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), B: T.Tensor((16, 16), "float32") ) -> None: for i in T.thread_binding(0, 16, thread="threadIdx.x"): for j in T.thread_binding( @@ -495,7 +495,7 @@ def test_for_thread_binding(): def while_loop(): @Ts.prim_func - def while_loop(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")) -> None: + def while_loop(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")) -> None: i = Ts.sblock_alloc_buffer((), "int32", scope="local") for ii in range(16): with Ts.sblock(): @@ -543,7 +543,7 @@ def func(): def implicit_evaluate(): @Ts.prim_func - def func(A: T.Buffer(1, "int32")): + def func(A: T.Tensor(1, "int32")): T.evaluate(T.assume(A[0] == 5)) A[0] = 10 @@ -619,7 +619,7 @@ class Module: I.module_attrs({"attr": 10}) @Ts.prim_func - def tir_func(A: T.Buffer(16, "int32"), B: T.Buffer(16, "int32")): + def tir_func(A: T.Tensor(16, "int32"), B: T.Tensor(16, "int32")): for i in range(16): B[i] = A[i] @@ -632,7 +632,7 @@ def subroutine_call(): @I.ir_module class mod: @Ts.prim_func - def main(A: T.Buffer(16, "float32")): + def main(A: T.Tensor(16, "float32")): mod.subroutine(A.data, T.int32(16)) @Ts.prim_func @@ -648,7 +648,7 @@ def subroutine_call_returning_int(): @I.ir_module class mod: @Ts.prim_func - def main(A: T.Buffer(2, "float32")): + def main(A: T.Tensor(2, "float32")): mod.subroutine(A[0]) + mod.subroutine(A[1]) @Ts.prim_func @@ -699,7 +699,7 @@ def func() -> T.int32: def func_with_loop_jumps(): @Ts.prim_func - def func(In: T.Buffer((1,), "int32"), Out: T.Buffer((2,), "int32")): + def func(In: T.Tensor((1,), "int32"), Out: T.Tensor((2,), "int32")): Out[0] = 0 Out[1] = 0 for i in range(1000): @@ -716,7 +716,7 @@ def func(In: T.Buffer((1,), "int32"), Out: T.Buffer((2,), "int32")): def func_with_loop_steps(): @Ts.prim_func def func( - A: T.Buffer((1024,)), B: T.Buffer((1024,)), C: T.Buffer((1024,)), tid: T.int32, v: T.int32 + A: T.Tensor((1024,)), B: T.Tensor((1024,)), C: T.Tensor((1024,)), tid: T.int32, v: T.int32 ): for i in T.serial(tid, 1024, step=2): C[i] = A[i] + B[i] @@ -769,7 +769,7 @@ def func(x: T.int32): @Ts.prim_func -def loop_no_syntax_sugar(A: T.Buffer((128, 128, 128, 128))) -> None: +def loop_no_syntax_sugar(A: T.Tensor((128, 128, 128, 128))) -> None: for i in T.serial(0, 128): for j in T.parallel(0, 128): for k in T.vectorized(0, 128): @@ -780,7 +780,7 @@ def loop_no_syntax_sugar(A: T.Buffer((128, 128, 128, 128))) -> None: @Ts.prim_func -def loop_syntax_sugar(A: T.Buffer((128, 128, 128, 128))) -> None: +def loop_syntax_sugar(A: T.Tensor((128, 128, 128, 128))) -> None: for i in T.serial(128): for j in T.parallel(128): for k in T.vectorized(128): @@ -796,8 +796,8 @@ def test_loop_syntax_sugar(): @Ts.prim_func def elementwise_buffer_default_dtype( - A: T.Buffer((128, 128, 128, 128)), - B: T.Buffer((128, 128, 128, 128)), + A: T.Tensor((128, 128, 128, 128)), + B: T.Tensor((128, 128, 128, 128)), ) -> None: for i, j, k, l in T.grid(128, 128, 128, 128): # noqa: E741 with Ts.sblock("B"): @@ -807,8 +807,8 @@ def elementwise_buffer_default_dtype( @Ts.prim_func def elementwise_buffer_kwargs( - a: T.Buffer(shape=(128, 128, 128, 128), dtype="float32"), - b: T.Buffer(shape=(128, 128, 128, 128), dtype="float32"), + a: T.Tensor(shape=(128, 128, 128, 128), dtype="float32"), + b: T.Tensor(shape=(128, 128, 128, 128), dtype="float32"), ) -> None: for i, j, k, l in T.grid(128, 128, 128, 128): # noqa: E741 with Ts.sblock("B"): @@ -818,8 +818,8 @@ def elementwise_buffer_kwargs( @Ts.prim_func def elementwise_buffer_no_kwargs( - a: T.Buffer((128, 128, 128, 128), "float32"), - b: T.Buffer((128, 128, 128, 128), "float32"), + a: T.Tensor((128, 128, 128, 128), "float32"), + b: T.Tensor((128, 128, 128, 128), "float32"), ) -> None: for i, j, k, l in T.grid(128, 128, 128, 128): # noqa: E741 with Ts.sblock("B"): @@ -840,12 +840,12 @@ def test_buffer_signature_syntax_sugar(): def test_buffer_1d(): @Ts.prim_func - def func_no_sugar(A: T.Buffer(shape=(16,))): + def func_no_sugar(A: T.Tensor(shape=(16,))): for i in T.serial(16): A[i] = 0.0 @Ts.prim_func - def func_with_sugar(A: T.Buffer(16, "float32")): + def func_with_sugar(A: T.Tensor(16, "float32")): for i in T.serial(16): A[i] = 0.0 @@ -861,19 +861,19 @@ def test_bind_bufferload_without_type_annotation(): # Failure occurred during parsing of the tvmscript. @Ts.prim_func - def func_without_type_annotation(A: T.Buffer((1,), "int32")): + def func_without_type_annotation(A: T.Tensor((1,), "int32")): x = A[0] T.evaluate(x) def test_implicit_evaluate_assume(): @Ts.prim_func - def explicit(A: T.Buffer(1, "int32")): + def explicit(A: T.Tensor(1, "int32")): T.evaluate(T.assume(A[0] == 5)) A[0] = 10 @Ts.prim_func - def implicit(A: T.Buffer(1, "int32")): + def implicit(A: T.Tensor(1, "int32")): T.assume(A[0] == 5) A[0] = 10 @@ -882,11 +882,11 @@ def implicit(A: T.Buffer(1, "int32")): def test_implicit_evaluate_call_extern(): @Ts.prim_func - def explicit(A: T.Buffer(1, "int32")): + def explicit(A: T.Tensor(1, "int32")): T.evaluate(T.call_extern("extern_func", A.data, dtype="int32")) @Ts.prim_func - def implicit(A: T.Buffer(1, "int32")): + def implicit(A: T.Tensor(1, "int32")): T.call_extern("extern_func", A.data, dtype="int32") assert_structural_equal_ignore_global_symbol(implicit, explicit) @@ -895,7 +895,7 @@ def implicit(A: T.Buffer(1, "int32")): def test_preserve_trivial_let_binding(): """Trivial `T.let[...]` annotations survive the parser as LetStmt and are not inlined. - In fork, bare `j = i` lowers to a local_scalar (AllocBuffer + BufferStore); the + In fork, bare `j = i` lowers to a local_scalar (AllocTensor + BufferStore); the LetStmt form is opt-in via `T.let[T.dtype]`. Both the explicit `T.bind(..., var=j)` builder API and the `j: T.let[T.dtype]` annotation produce the same LetStmt IR. """ @@ -945,7 +945,7 @@ def func(i: T.int32): @pytest.mark.parametrize("mutable", [False, True]) def test_preserve_variable_name(mutable): - """Use variable name when generating tirx::Bind / AllocBuffer""" + """Use variable name when generating tirx::Bind / AllocTensor""" # Bare bindings name the immutable Var; explicit declarations name scalar storage. annotation = ": T.int32" if mutable else "" @@ -961,7 +961,7 @@ def func(): binding = func.body.body.seq[0] if mutable: assert isinstance(binding, tvm.tirx.Bind) and isinstance(binding.value, tvm.ir.Call) - assert binding.value.op.name == "tirx.alloc_buffer" + assert binding.value.op.name == "tirx.alloc_tensor" var_name = binding.var.name else: assert isinstance(binding, tvm.tirx.Bind) @@ -1084,7 +1084,7 @@ def test_roundtrip_basic_usage(ir_generator): # Import-time construction also checks the annotated S-TIR API. @Ts.prim_func def element_wise_env_thread_x( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: j1_0 = T.env_thread("threadIdx.x") j0_0 = T.env_thread("threadIdx.x") diff --git a/tests/python/s_tir/script/test_s_tir_script_blocks.py b/tests/python/s_tir/script/test_s_tir_script_blocks.py index 5e5570c6c854..7f6c5f66bae8 100644 --- a/tests/python/s_tir/script/test_s_tir_script_blocks.py +++ b/tests/python/s_tir/script/test_s_tir_script_blocks.py @@ -30,7 +30,7 @@ @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) @@ -41,7 +41,7 @@ def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 1 @Ts.prim_func def matmul_original( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i, j in T.grid(32, 32): with Ts.sblock("init"): @@ -61,7 +61,7 @@ def matmul_original( @Ts.prim_func def elementwise_with_root( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: with Ts.sblock(): for i, j in T.grid(128, 128): @@ -76,7 +76,7 @@ def elementwise_with_root( @Ts.prim_func def func_with_part_access_region( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: with Ts.sblock(): for i, j in T.grid(128, 128): @@ -184,7 +184,7 @@ def test_complete_part_region(): @Ts.prim_func def func_with_bufferslice_indices( - data_buf: T.Buffer((16, 16), "float32"), index_buf: T.Buffer((1,), "int32") + data_buf: T.Tensor((16, 16), "float32"), index_buf: T.Tensor((1,), "int32") ) -> None: out_buf = Ts.sblock_alloc_buffer((16, 16), "float32") @@ -196,8 +196,8 @@ def func_with_bufferslice_indices( @Ts.prim_func def expected_bufferslice_indices( - data_buf: T.Buffer([16, 16], elem_offset=0, align=64, offset_factor=1), - index_buf: T.Buffer([1], dtype="int32", elem_offset=0, align=64, offset_factor=1), + data_buf: T.Tensor([16, 16], elem_offset=0, align=64, offset_factor=1), + index_buf: T.Tensor([1], dtype="int32", elem_offset=0, align=64, offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads([]) @@ -213,7 +213,7 @@ def expected_bufferslice_indices( @Ts.prim_func def func_with_recursive_bufferslice_indices( - data_buf: T.Buffer((16, 16), "float32"), index_buf: T.Buffer((1,), "int32") + data_buf: T.Tensor((16, 16), "float32"), index_buf: T.Tensor((1,), "int32") ) -> None: out_buf = Ts.sblock_alloc_buffer((16, 16), "float32") @@ -225,8 +225,8 @@ def func_with_recursive_bufferslice_indices( @Ts.prim_func def expected_recursive_bufferslice_indices( - data_buf: T.Buffer([16, 16], elem_offset=0, align=64, offset_factor=1), - index_buf: T.Buffer([1], dtype="int32", elem_offset=0, align=64, offset_factor=1), + data_buf: T.Tensor([16, 16], elem_offset=0, align=64, offset_factor=1), + index_buf: T.Tensor([1], dtype="int32", elem_offset=0, align=64, offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads([]) @@ -263,7 +263,7 @@ def test_complete_buffer_indices(): @Ts.prim_func -def match_buffer_func(A: T.Buffer((16, 16))) -> None: +def match_buffer_func(A: T.Tensor((16, 16))) -> None: for i in range(0, 16): with Ts.sblock(): A0 = Ts.match_buffer(A[i, 0:16], (16)) @@ -275,7 +275,7 @@ def match_buffer_func(A: T.Buffer((16, 16))) -> None: @Ts.prim_func -def expected_match_buffer_func(A: T.Buffer((16, 16))) -> None: +def expected_match_buffer_func(A: T.Tensor((16, 16))) -> None: for i in range(0, 16): with Ts.sblock(): Ts.reads([]) @@ -301,7 +301,7 @@ def test_complete_match_buffer(): @Ts.prim_func def alloc_buffer_func( - A: T.Buffer([2, 2], dtype="float32"), B: T.Buffer([2, 2], dtype="float32") + A: T.Tensor([2, 2], dtype="float32"), B: T.Tensor([2, 2], dtype="float32") ) -> None: C = Ts.sblock_alloc_buffer([2, 2], dtype="float32") A[(0, 0)] = T.float32(2) @@ -311,8 +311,8 @@ def alloc_buffer_func( @Ts.prim_func def expect_alloc_buffer_func( - A: T.Buffer([2, 2], dtype="float32", elem_offset=0, align=64, offset_factor=1), - B: T.Buffer([2, 2], dtype="float32", elem_offset=0, align=64, offset_factor=1), + A: T.Tensor([2, 2], dtype="float32", elem_offset=0, align=64, offset_factor=1), + B: T.Tensor([2, 2], dtype="float32", elem_offset=0, align=64, offset_factor=1), ) -> None: with Ts.sblock("root"): Ts.reads([]) @@ -337,7 +337,7 @@ def test_complete_alloc_buffer(): @Ts.prim_func def alloc_zero_dim_buffer( - A: T.Buffer([], dtype="float32"), B: T.Buffer([], dtype="float32") + A: T.Tensor([], dtype="float32"), B: T.Tensor([], dtype="float32") ) -> None: # body # tirx.with block("root") @@ -348,7 +348,7 @@ def alloc_zero_dim_buffer( @Ts.prim_func -def alloc_zero_dim_buffer_block(A: T.Buffer((), "float32"), B: T.Buffer((), "float32")) -> None: +def alloc_zero_dim_buffer_block(A: T.Tensor((), "float32"), B: T.Tensor((), "float32")) -> None: with Ts.sblock("root"): Ts.reads([]) Ts.writes([]) @@ -405,7 +405,7 @@ def test_alloc_zero_dim_buffer_round_trip(): @Ts.prim_func def slice_op_test( - A: T.Buffer((10,), "float32"), B: T.Buffer((10,), "float32"), C: T.Buffer((10,), "uint32") + A: T.Tensor((10,), "float32"), B: T.Tensor((10,), "float32"), C: T.Tensor((10,), "uint32") ): B[0:5] = A[0:5] + B[0:5] B[0:5] = A[0:5] - B[0:5] @@ -436,7 +436,7 @@ def slice_op_test( @Ts.prim_func def slice_op_test_ref( - A: T.Buffer((10,), "float32"), B: T.Buffer((10,), "float32"), C: T.Buffer((10,), "uint32") + A: T.Tensor((10,), "float32"), B: T.Tensor((10,), "float32"), C: T.Tensor((10,), "uint32") ): B[0:5] = A[0:5] + B[0:5] B[0:5] = A[0:5] - B[0:5] @@ -496,7 +496,7 @@ def func_ref(): def roundtrip_matmul(): @Ts.prim_func def roundtrip_matmul( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): @@ -511,7 +511,7 @@ def roundtrip_matmul( def roundtrip_matmul_original(): @Ts.prim_func def roundtrip_matmul_original( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i, j in T.grid(128, 128): with Ts.sblock("init"): @@ -529,7 +529,7 @@ def roundtrip_matmul_original( def element_wise(): @Ts.prim_func def element_wise( - A: T.Buffer((128, 128), "float32"), C: T.Buffer((128, 128), "float32") + A: T.Tensor((128, 128), "float32"), C: T.Tensor((128, 128), "float32") ) -> None: B = Ts.sblock_alloc_buffer((128, 128), "float32") @@ -547,7 +547,7 @@ def element_wise( def predicate(): @Ts.prim_func - def predicate(B: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: + def predicate(B: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i, jo, ji in T.grid(16, 4, 5): with Ts.sblock("update"): vi = Ts.axis.S(16, i) @@ -638,7 +638,7 @@ def test_predicate(): def match_buffer_region(): @Ts.prim_func def match_buffer_region( - A: T.Buffer((16, 16, 16), "float32"), B: T.Buffer(1, "float32") + A: T.Tensor((16, 16, 16), "float32"), B: T.Tensor(1, "float32") ) -> None: for i, j in T.grid(16, 4): with Ts.sblock(): @@ -688,7 +688,7 @@ def test_match_buffer_region(): def block_elements(): @Ts.prim_func - def block_elements(A: T.Buffer((16, 16), "float32"), B: T.Buffer((1, 1), "float32")) -> None: + def block_elements(A: T.Tensor((16, 16), "float32"), B: T.Tensor((1, 1), "float32")) -> None: with Ts.sblock("update"): vi = Ts.axis.S(1, 0) Ts.where(True) @@ -729,7 +729,7 @@ def test_block_elements(): def opaque_block(): @Ts.prim_func - def opaque_block(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")) -> None: + def opaque_block(A: T.Tensor((16, 16), "float32"), B: T.Tensor((16, 16), "float32")) -> None: for i in range(16): for j in range(16): with Ts.sblock(): @@ -772,7 +772,7 @@ def test_opaque_block(): def rank0(): @Ts.prim_func - def rank0(A: T.Buffer((), "float32")) -> None: + def rank0(A: T.Tensor((), "float32")) -> None: B = Ts.sblock_alloc_buffer((), "float32") A[()] = 2 B[()] = A[()] @@ -782,7 +782,7 @@ def rank0(A: T.Buffer((), "float32")) -> None: def rank0_block(): @Ts.prim_func - def rank0_block(A: T.Buffer((), "float32")) -> None: + def rank0_block(A: T.Tensor((), "float32")) -> None: B = Ts.sblock_alloc_buffer((), "float32") B[()] = A[()] @@ -797,7 +797,7 @@ def rank0_block(A: T.Buffer((), "float32")) -> None: def nontrivial_range_axis(): @Ts.prim_func - def nontrivial_range_axis(A: T.Buffer(10, "float32")) -> None: + def nontrivial_range_axis(A: T.Tensor(10, "float32")) -> None: for i in range(10): with Ts.sblock("block"): vi = Ts.axis.spatial((1, 11), i + 1) @@ -818,7 +818,7 @@ def func_root_attr(): def func_trivial_root_block(): @Ts.prim_func - def func(A: T.Buffer(1, "int32")): + def func(A: T.Tensor(1, "int32")): with Ts.sblock("root"): A[0] = 0 @@ -827,7 +827,7 @@ def func(A: T.Buffer(1, "int32")): def func_nested_root_block(): @Ts.prim_func - def func(A: T.Buffer(1, "int32")): + def func(A: T.Tensor(1, "int32")): with Ts.sblock("root"): with Ts.sblock("block"): A[0] = 0 @@ -838,8 +838,8 @@ def func(A: T.Buffer(1, "int32")): def int64_support(): @Ts.prim_func def elementwise_shape_int64( - A: T.Buffer((T.int64(128), T.int64(128)), dtype="float32"), - C: T.Buffer((T.int64(128), T.int64(128)), dtype="float32"), + A: T.Tensor((T.int64(128), T.int64(128)), dtype="float32"), + C: T.Tensor((T.int64(128), T.int64(128)), dtype="float32"), ) -> None: B = Ts.sblock_alloc_buffer((T.int64(128), T.int64(128)), dtype="float32") @@ -858,9 +858,9 @@ def elementwise_shape_int64( def func_attr_with_list(): @Ts.prim_func def func( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - D: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + D: T.Tensor((128, 128), "float32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True, "layout_free_buffers": [1]}) C = Ts.sblock_alloc_buffer([128, 128], dtype="float32") @@ -881,7 +881,7 @@ def func( @Ts.prim_func def transformed_matmul_no_syntax_sugar( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i0, i1, i2_outer, i2_inner_outer, i2_inner_inner in T.grid(128, 128, 4, 8, 4): with Ts.sblock("update"): @@ -897,7 +897,7 @@ def transformed_matmul_no_syntax_sugar( @Ts.prim_func def transformed_matmul_syntax_sugar( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i0, i1, i2_outer, i2_inner_outer, i2_inner_inner in T.grid(128, 128, 4, 8, 4): with Ts.sblock("update"): @@ -919,13 +919,13 @@ def test_reads_writes_syntax_sugar(): def test_match_buffer_region_has_implicit_shape_dtype(): @Ts.prim_func - def explicit_shape_dtype(A: T.Buffer((16, 64), "int32")): + def explicit_shape_dtype(A: T.Tensor((16, 64), "int32")): with Ts.sblock(): B = Ts.match_buffer(A[8:16, 32:64], shape=(8, 32), dtype="int32") T.evaluate(0) @Ts.prim_func - def implicit_shape_dtype(A: T.Buffer((16, 64), "int32")): + def implicit_shape_dtype(A: T.Tensor((16, 64), "int32")): with Ts.sblock(): B = Ts.match_buffer(A[8:16, 32:64]) T.evaluate(0) @@ -966,8 +966,8 @@ def test_roundtrip_blocks(ir_generator): # Import-time construction also checks the annotated S-TIR API. @Ts.prim_func def element_wise_storage_align( - A: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), - C: T.Buffer([128, 128], elem_offset=0, align=64, offset_factor=1), + A: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), + C: T.Tensor([128, 128], elem_offset=0, align=64, offset_factor=1), ) -> None: # body with Ts.sblock("root"): @@ -994,7 +994,7 @@ def element_wise_storage_align( # Import-time construction also checks the annotated S-TIR API. @Ts.prim_func def loop_split( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: for i, ko in T.grid(128, 4): for ki in T.thread_binding(0, 32, thread="threadIdx.x"): @@ -1011,7 +1011,7 @@ def loop_split( # Import-time construction also checks the annotated S-TIR API. @Ts.prim_func def different_access_indices( - A: T.Buffer([128, 128, 128], dtype="float32"), B: T.Buffer([128, 128], dtype="float32") + A: T.Tensor([128, 128, 128], dtype="float32"), B: T.Tensor([128, 128], dtype="float32") ) -> None: for i, j in T.grid(128, 128): for k in T.thread_binding(0, 128, thread="threadIdx.x"): diff --git a/tests/python/s_tir/script/test_s_tir_script_dynamic_shape.py b/tests/python/s_tir/script/test_s_tir_script_dynamic_shape.py index 1230499a107e..d82771511022 100644 --- a/tests/python/s_tir/script/test_s_tir_script_dynamic_shape.py +++ b/tests/python/s_tir/script/test_s_tir_script_dynamic_shape.py @@ -33,12 +33,12 @@ def test_tir_starred_shape_expression(): dims = (128, 128) @Ts.prim_func(private=True) - def starred(A: T.Buffer([128, *dims], "int32")) -> None: + def starred(A: T.Tensor([128, *dims], "int32")) -> None: for i, j, k in T.grid(*A.shape): A[i, j, k] = T.int32(1) @Ts.prim_func(private=True) - def non_starred(A: T.Buffer([128, 128, 128], "int32")) -> None: + def non_starred(A: T.Tensor([128, 128, 128], "int32")) -> None: for i, j, k in T.grid(128, 128, 128): A[i, j, k] = T.int32(1) @@ -52,7 +52,7 @@ def test_inferred_ty_with_dynamic_buffer(): N = T.dynamic("N", "int64") @Ts.prim_func - def func(A: T.Buffer([M, N], "float32"), B: T.Buffer([M * N], "float32")): + def func(A: T.Tensor([M, N], "float32"), B: T.Tensor([M * N], "float32")): for i, j in T.grid(M, N): B[i * N + j] = A[i, j] @@ -71,7 +71,7 @@ def func(A: T.Buffer([M, N], "float32"), B: T.Buffer([M * N], "float32")): def test_tir_buffer_region_extent_correct_dtype(): @Ts.prim_func - def func(A: T.Buffer((T.int64(16), T.int64(1)), "float32")): + def func(A: T.Tensor((T.int64(16), T.int64(1)), "float32")): for i in T.grid(T.int64(16)): with Ts.sblock("block"): vi = Ts.axis.remap("S", [i]) @@ -92,7 +92,7 @@ def func(A: T.Buffer((T.int64(16), T.int64(1)), "float32")): @Ts.prim_func def gemm_dyn_shape( - A: T.Buffer((N, K), "float32"), B: T.Buffer((K, M), "float32"), C: T.Buffer((N, M), "float32") + A: T.Tensor((N, K), "float32"), B: T.Tensor((K, M), "float32"), C: T.Tensor((N, M), "float32") ): for i, j, k in T.grid(N, M, K): with Ts.sblock("gemm"): @@ -113,8 +113,8 @@ def test_dynamic_shape_gemm(): @Ts.prim_func def buffer_int64( - A: T.Buffer((T.int64(128), T.int64(128)), dtype="float32"), - C: T.Buffer((T.int64(128), T.int64(128)), dtype="float32"), + A: T.Tensor((T.int64(128), T.int64(128)), dtype="float32"), + C: T.Tensor((T.int64(128), T.int64(128)), dtype="float32"), ) -> None: B = Ts.sblock_alloc_buffer((T.int64(128), T.int64(128)), dtype="float32") @@ -130,8 +130,8 @@ def buffer_int64( @Ts.prim_func def buffer_int64_after_roundtrip( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - C: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + C: T.Tensor((T.int64(128), T.int64(128)), "float32"), ) -> None: B = Ts.sblock_alloc_buffer((T.int64(128), T.int64(128)), dtype="float32") for i, j in T.grid(128, 128): @@ -153,8 +153,8 @@ def test_buffer_int64(): def test_int64_loop(): @Ts.prim_func def int64_grid( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(128)), "float32"), ) -> None: for i, j in T.grid(T.int64(128), T.int64(128)): with Ts.sblock("C"): @@ -163,8 +163,8 @@ def int64_grid( @Ts.prim_func def int64_grid_expanded( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(128)), "float32"), ) -> None: for i in range(T.int64(0), T.int64(128)): for j in range(T.int64(0), T.int64(128)): @@ -178,7 +178,7 @@ def int64_grid_expanded( def loop_extent_dependent(): @Ts.prim_func - def loop_extent_dependent(A: T.Buffer([], dtype="int32")) -> None: + def loop_extent_dependent(A: T.Tensor([], dtype="int32")) -> None: for i in T.serial(0, 128): for j in T.serial(0, i): A[()] = A[()] + j @@ -190,9 +190,9 @@ def parse_bufferslice_as_range_bound(): # apparently the use of i in the "outer" block when it is defined outside of a block is wrong @Ts.prim_func(check_well_formed=False) def segment_sum( - A: T.Buffer([m], dtype="float32"), # noqa: F821 - B: T.Buffer([n], dtype="float32"), # noqa: F821 - indptr: T.Buffer([n + 1], dtype="int32"), # noqa: F821 + A: T.Tensor([m], dtype="float32"), # noqa: F821 + B: T.Tensor([n], dtype="float32"), # noqa: F821 + indptr: T.Tensor([n + 1], dtype="int32"), # noqa: F821 n: T.int32, m: T.int32, ) -> None: @@ -219,7 +219,7 @@ def undefined_shape_in_decl_buffer(): @Ts.prim_func(check_well_formed=False) def func(): - buf = T.decl_buffer(shape=[size], dtype="float32") + buf = T.decl_tensor(shape=[size], dtype="float32") T.evaluate(buf[0]) return func @@ -232,7 +232,7 @@ def undefined_stride_in_decl_buffer(): @Ts.prim_func(check_well_formed=False) def func(): data_ptr = T.handle("float32") - buf = T.decl_buffer(shape=[1], dtype="float32", data=data_ptr, strides=[stride]) + buf = T.decl_tensor(shape=[1], dtype="float32", data=data_ptr, strides=[stride]) T.evaluate(buf[0]) return func @@ -245,7 +245,7 @@ def undefined_elem_offset_in_decl_buffer(): @Ts.prim_func(check_well_formed=False) def func(): data_ptr = T.handle("float32") - buf = T.decl_buffer(shape=[1], dtype="float32", data=data_ptr, elem_offset=elem_offset) + buf = T.decl_tensor(shape=[1], dtype="float32", data=data_ptr, elem_offset=elem_offset) T.evaluate(buf[0]) return func diff --git a/tests/python/s_tir/script/test_s_tir_script_error_handling.py b/tests/python/s_tir/script/test_s_tir_script_error_handling.py index ae83dc488b55..7762703f3bf8 100644 --- a/tests/python/s_tir/script/test_s_tir_script_error_handling.py +++ b/tests/python/s_tir/script/test_s_tir_script_error_handling.py @@ -65,14 +65,14 @@ def check_error(func, rel_lineno, error_type): def test_buffer_bind(): - def buffer_bind_missing_args(A: T.Buffer(dtype="float32")) -> None: # error + def buffer_bind_missing_args(A: T.Tensor(dtype="float32")) -> None: # error T.evaluate(0) check_error(buffer_bind_missing_args, 1, TypeError) def test_undefined_buffer(): - def undefined_buffer(A: T.Buffer((16, 16), "float32")) -> None: + def undefined_buffer(A: T.Tensor((16, 16), "float32")) -> None: for i in T.serial(16): for j in T.serial(0, 16): C[i, j] = 0.0 # error # noqa: F821 @@ -81,7 +81,7 @@ def undefined_buffer(A: T.Buffer((16, 16), "float32")) -> None: def test_unsupported_function_call(): - def unsupported_function_call(A: T.Buffer((16, 16), "float32")) -> None: + def unsupported_function_call(A: T.Tensor((16, 16), "float32")) -> None: for i in T.const_range(16): # error for j in T.serial(0, 16): A[i, j] = 0.0 @@ -90,7 +90,7 @@ def unsupported_function_call(A: T.Buffer((16, 16), "float32")) -> None: def test_invalid_for_function(): - def invalid_for_function(A: T.Buffer((16, 16), "float32")) -> None: + def invalid_for_function(A: T.Tensor((16, 16), "float32")) -> None: for i in T.evaluate(0.0): # error for j in T.serial(0, 16): A[i, j] = 0.0 @@ -99,7 +99,7 @@ def invalid_for_function(A: T.Buffer((16, 16), "float32")) -> None: def test_invalid_block_function(): - def invalid_block_function(A: T.Buffer((16, 16), "float32")) -> None: + def invalid_block_function(A: T.Tensor((16, 16), "float32")) -> None: with T.evaluate(0.0): # error T.evaluate(1.0) @@ -116,7 +116,7 @@ def return_not_allowed(a: T.handle) -> None: def test_no_body(): - def no_body(A: T.Buffer((16, 16), "float32")) -> None: + def no_body(A: T.Tensor((16, 16), "float32")) -> None: T.realize(A, "") # error check_error(no_body, 2, AttributeError) @@ -155,7 +155,7 @@ def error_remap_value() -> None: def test_invalid_block_axes(): - def invalid_block_axes(A: T.Buffer((16, 16), "float32")) -> None: + def invalid_block_axes(A: T.Tensor((16, 16), "float32")) -> None: for i, j in T.grid(16, 16): with Ts.sblock(): vi = Ts.axis.S(i, A) # error @@ -208,7 +208,7 @@ def invalid_loop_var() -> None: def test_inconsistent_grid(): - def inconsistent_grid(A: T.Buffer(16)) -> None: + def inconsistent_grid(A: T.Tensor(16)) -> None: for (i,) in T.grid(16, 16): # error: one explicit target cannot unpack two variables T.evaluate(A[i]) @@ -307,7 +307,7 @@ def duplicate_sblock_attrs_with_same_key_diff_value() -> None: def test_opaque_access_during_complete(): - def opaque_access_during_complete(A: T.Buffer((16, 16), "float32")) -> None: # error + def opaque_access_during_complete(A: T.Tensor((16, 16), "float32")) -> None: # error for i, j in T.grid(16, 16): with Ts.sblock(): T.evaluate(T.call_extern("dummy_extern_function", A.data, dtype="int32")) @@ -343,21 +343,21 @@ def scope_handler_except() -> None: def test_tvm_exception_catch_from_bare_intrin(): - def intrin_except_unassign(A: T.Buffer((16, 16), "float32")) -> None: + def intrin_except_unassign(A: T.Tensor((16, 16), "float32")) -> None: T.evaluate(A) # error check_error(intrin_except_unassign, 2, tvm.error.InternalError) def test_tvm_exception_catch_from_assigned_intrin(): - def intrin_except_assign(A: T.Buffer((16, 16), "float32")) -> None: + def intrin_except_assign(A: T.Tensor((16, 16), "float32")) -> None: A[0, 0] = A[A] # error check_error(intrin_except_assign, 2, tvm.error.InternalError) def test_match_buffer_shape_mismatch(): - def buffer_shape_mismatch(A: T.Buffer((8, 8))) -> None: + def buffer_shape_mismatch(A: T.Tensor((8, 8))) -> None: for i, j in T.grid(8, 2): with Ts.sblock(): Ts.reads([]) @@ -374,7 +374,7 @@ def buffer_shape_mismatch(A: T.Buffer((8, 8))) -> None: def test_high_dim_store(): def high_dim_store() -> None: with Ts.sblock("root"): - B = T.alloc_buffer((256,), "float32") + B = T.alloc_tensor((256,), "float32") for i, j in T.grid(16, 16): B[i, j] = 1.0 # error: Store is only allowed with one index @@ -419,7 +419,7 @@ def implicit_root_has_axes(): @Ts.prim_func def elementwise_not_affine( - A: T.Buffer((128, 128, 128, 128)), B: T.Buffer((128, 128, 128, 128)) + A: T.Tensor((128, 128, 128, 128)), B: T.Tensor((128, 128, 128, 128)) ) -> None: for i, j, k, l in T.grid(128, 128, 128, 8): # noqa: E741 with Ts.sblock("B"): @@ -430,7 +430,7 @@ def elementwise_not_affine( @Ts.prim_func def elementwise_non_single_branch( - A: T.Buffer((128, 128, 128)), B: T.Buffer((128, 128, 128)) + A: T.Tensor((128, 128, 128)), B: T.Tensor((128, 128, 128)) ) -> None: C = Ts.sblock_alloc_buffer((128, 128, 128)) @@ -519,28 +519,28 @@ def store_var_multiple() -> None: def test_load_handle(): - def load_handle(h: T.handle, h_: T.Buffer([1])) -> None: + def load_handle(h: T.handle, h_: T.Tensor([1])) -> None: h_[0] = h[0] # error cannot load from handle check_error(load_handle, 2, TypeError) def test_store_handle(): - def store_handle(h: T.handle, h_: T.Buffer([1])) -> None: + def store_handle(h: T.handle, h_: T.Tensor([1])) -> None: h[0] = h_[0] # error cannot store to handle check_error(store_handle, 2, TypeError) def test_binop_bad_ast_type(): - def binop_bad_ast_type(h: T.handle, h_: T.Buffer([1])): + def binop_bad_ast_type(h: T.handle, h_: T.Tensor([1])): h_[0] = h + [2] # error rhs should be a primexpr # noqa: RUF005 check_error(binop_bad_ast_type, 2, TypeError) def test_binop_bad_type(): - def binop_bad_type(h: T.handle, h_: T.Buffer([1])): + def binop_bad_type(h: T.handle, h_: T.Tensor([1])): h_[0] = h + 2 # error lhs and rhs should be the same type check_error(binop_bad_type, 2, TypeError) @@ -555,7 +555,7 @@ def non_integer_typed_block_iter(): def test_illegal_buffer_slice(): - def strided_buffer_region(A: T.Buffer((128, 128), "int32")): + def strided_buffer_region(A: T.Tensor((128, 128), "int32")): # do not allow stride in buffer region with Ts.sblock("block"): @@ -563,12 +563,12 @@ def strided_buffer_region(A: T.Buffer((128, 128), "int32")): Ts.writes([A[0:128:2, 0:128:3]]) # error T.evaluate(T.call_extern("strided_compute", dtype="")) - def access_reversed_slice(A: T.Buffer((128,), "int32")): + def access_reversed_slice(A: T.Tensor((128,), "int32")): # do not allow reversed slice step A[0:128:-1] = T.broadcast(1, 128) # error - def access_non_const_slice_length(A: T.Buffer((128,), "int32")): + def access_non_const_slice_length(A: T.Tensor((128,), "int32")): # do not allow non-constant slice length for i in range(4): @@ -580,7 +580,7 @@ def access_non_const_slice_length(A: T.Buffer((128,), "int32")): def test_syntax_sugar_fail(): - def loop_syntax_sugar_fail(A: T.Buffer((128,))) -> None: + def loop_syntax_sugar_fail(A: T.Tensor((128,))) -> None: for i in T.thread_binding(128, 128): A[i] = A[i] * 2.0 @@ -632,7 +632,7 @@ def test_tir_func_private_manual_global_symbol_fail(): @Ts.prim_func(private=True) def matmul( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: T.func_attr({"global_symbol": "matmul"}) @@ -649,7 +649,7 @@ def test_buffer_input_requires_shape_arg(): with pytest.raises(TypeError): @Ts.prim_func - def func(A: T.Buffer(dtype="int32")): + def func(A: T.Tensor(dtype="int32")): T.evaluate(0) diff --git a/tests/python/s_tir/script/test_s_tir_script_ir_builder.py b/tests/python/s_tir/script/test_s_tir_script_ir_builder.py index 43195e6c98ea..3fb4c4c5137f 100644 --- a/tests/python/s_tir/script/test_s_tir_script_ir_builder.py +++ b/tests/python/s_tir/script/test_s_tir_script_ir_builder.py @@ -74,9 +74,9 @@ def test_ir_builder_tir_primfunc_complete(): with build_prim_func(): T.arg_("a", T.handle()) T.arg_("b", T.int64()) - T.arg_("c", T.Buffer((128, 128), "float32")) - buffer_d = T.arg_("d", T.Buffer((64, 64), "int64")) - e = T.arg_("e", T.Buffer((1024,), "int8")) + T.arg_("c", T.Tensor((128, 128), "float32")) + buffer_d = T.arg_("d", T.Tensor((64, 64), "int64")) + e = T.arg_("e", T.Tensor((1024,), "int8")) T.func_attr({"key": "value"}) T.func_ret(tvm.ir.PrimType("int64")) @@ -86,9 +86,9 @@ def test_ir_builder_tir_primfunc_complete(): prim_func_actual = ib.get() # the expected prim_func - c_buffer = tirx.decl_buffer((128, 128), "float32", name="c", layout=None) - d_buffer = tirx.decl_buffer((64, 64), "int64", name="d", layout=None) - e_buffer = tirx.decl_buffer((1024,), "int8", name="e", layout=None) + c_buffer = tirx.decl_tensor((128, 128), "float32", name="c", layout=None) + d_buffer = tirx.decl_tensor((64, 64), "int64", name="d", layout=None) + e_buffer = tirx.decl_tensor((1024,), "int8", name="e", layout=None) prim_func_expected = tirx.PrimFunc( params=[ tirx.Var("a", tvm.ir.PointerType(tvm.ir.PrimType("void"))), @@ -158,10 +158,10 @@ def test_ir_builder_tir_block_complete(): # the expected block var_a = tirx.Var("a", "int64") - buffer_b = tirx.decl_buffer((128, 128), "float32", name="b") - buffer_c = tirx.decl_buffer((128, 128), "float32", name="c") + buffer_b = tirx.decl_tensor((128, 128), "float32", name="b") + buffer_c = tirx.decl_tensor((128, 128), "float32", name="c") var_d = tirx.Var("d", "int32") - buffer_e = tirx.decl_buffer((128, 128), "float32", name="c") + buffer_e = tirx.decl_tensor((128, 128), "float32", name="c") var_f = tirx.Var("f", "int32") block_expected = s_tir.SBlock( iter_vars=[tirx.IterVar((0, 128), tirx.Var("", "int32"), iter_type=tirx.IterVar.DataPar)], @@ -169,9 +169,9 @@ def test_ir_builder_tir_block_complete(): writes=[buffer_c[var_d:128, var_d:128]], name_hint="block", body=tirx.Evaluate(0), - alloc_buffers=[tirx.decl_buffer((128, 128), "float32")], + alloc_buffers=[tirx.decl_tensor((128, 128), "float32")], match_buffers=[ - s_tir.MatchBufferRegion(tirx.decl_buffer((32, 32), "float32"), buffer_e[0:32, 0:32]) + s_tir.MatchBufferRegion(tirx.decl_tensor((32, 32), "float32"), buffer_e[0:32, 0:32]) ], annotations={"key": "value"}, ) @@ -252,17 +252,17 @@ def test_ir_builder_tir_allocate(): with IRBuilder() as ib: with build_prim_func(): T.func_name_("test") - buf = T.alloc_buffer([10], "float32", scope="local") + buf = T.alloc_tensor([10], "float32", scope="local") T.evaluate(1) # the allocate generated by IRBuilder ir_actual = ib.get() body = ir_actual.body - # AllocBuffer is flat: body should be a SeqStmt with [AllocBuffer, Evaluate(1)] + # AllocTensor is flat: body should be a SeqStmt with [AllocTensor, Evaluate(1)] assert isinstance(body, tirx.SeqStmt), f"Expected SeqStmt but got {type(body)}" assert len(body) == 2 - assert _is_buffer_binding(body[0], "tirx.alloc_buffer") + assert _is_buffer_binding(body[0], "tirx.alloc_tensor") assert isinstance(body[1], tirx.Evaluate) assert body[1].value.value == 1 @@ -271,17 +271,17 @@ def test_ir_builder_tir_decl_buffer(): with IRBuilder() as ib: with build_prim_func(): T.func_name_("test") - buf = T.decl_buffer([128, 128], "float32") + buf = T.decl_tensor([128, 128], "float32") T.evaluate(1) - # the decl_buffer generated by IRBuilder + # the decl_tensor generated by IRBuilder ir_actual = ib.get() body = ir_actual.body - # decl_buffer without data emits AllocBuffer (flat): body should be SeqStmt + # decl_tensor without data emits AllocTensor (flat): body should be SeqStmt assert isinstance(body, tirx.SeqStmt), f"Expected SeqStmt but got {type(body)}" assert len(body) == 2 - assert _is_buffer_binding(body[0], "tirx.alloc_buffer") + assert _is_buffer_binding(body[0], "tirx.alloc_tensor") assert isinstance(body[1], tirx.Evaluate) assert body[1].value.value == 1 diff --git a/tests/python/s_tir/script/test_s_tir_script_meta_programming.py b/tests/python/s_tir/script/test_s_tir_script_meta_programming.py index a275458af406..d624de5e90ce 100644 --- a/tests/python/s_tir/script/test_s_tir_script_meta_programming.py +++ b/tests/python/s_tir/script/test_s_tir_script_meta_programming.py @@ -34,9 +34,9 @@ def test_meta_programming_matmul(): def matmul_generator(M: int, N: int, K: int, dtype: str): @Ts.prim_func def matmul( - A: T.Buffer([M, K], dtype=dtype), - B: T.Buffer([N, K], dtype=dtype), - C: T.Buffer([M, N], dtype=dtype), + A: T.Tensor([M, K], dtype=dtype), + B: T.Tensor([N, K], dtype=dtype), + C: T.Tensor([M, N], dtype=dtype), ) -> None: for i, j, k in T.grid(M, N, K): with Ts.sblock(): @@ -49,9 +49,9 @@ def matmul( @Ts.prim_func def matmul_128_128_128_fp16( - A: T.Buffer([128, 128], dtype="float16"), - B: T.Buffer([128, 128], dtype="float16"), - C: T.Buffer([128, 128], dtype="float16"), + A: T.Tensor([128, 128], dtype="float16"), + B: T.Tensor([128, 128], dtype="float16"), + C: T.Tensor([128, 128], dtype="float16"), ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock(): @@ -67,7 +67,7 @@ def matmul_128_128_128_fp16( def test_meta_programming_uncaptured_var(): def generate_erf(dtype): @Ts.prim_func - def main(A: T.Buffer((1,), dtype), C: T.Buffer((1,), dtype)): + def main(A: T.Tensor((1,), dtype), C: T.Tensor((1,), dtype)): for i in range(1): with Ts.sblock("C"): C[i] = T.erf(A[i]) @@ -75,13 +75,13 @@ def main(A: T.Buffer((1,), dtype), C: T.Buffer((1,), dtype)): return main @Ts.prim_func - def fp32(A: T.Buffer((1,), "float32"), C: T.Buffer((1,), "float32")): + def fp32(A: T.Tensor((1,), "float32"), C: T.Tensor((1,), "float32")): for i in range(1): with Ts.sblock("C"): C[i] = T.erf(A[i]) @Ts.prim_func - def fp16(A: T.Buffer((1,), "float16"), C: T.Buffer((1,), "float16")): + def fp16(A: T.Tensor((1,), "float16"), C: T.Tensor((1,), "float16")): for i in range(1): with Ts.sblock("C"): C[i] = T.erf(A[i]) @@ -134,7 +134,7 @@ def assign(i, *args, t1, **kwargs): @Ts.prim_func(private=True) def matmul_w_macro( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): @@ -142,7 +142,7 @@ def matmul_w_macro( @Ts.prim_func(private=True) def matmul_no_macro( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): @@ -160,12 +160,12 @@ def static_capture(A, B): B[()] = A[x_value] @Ts.prim_func(private=True) - def use_hygienic(A: T.Buffer((1024,), "int32"), B: T.Buffer((), "int32")) -> None: + def use_hygienic(A: T.Tensor((1024,), "int32"), B: T.Tensor((), "int32")) -> None: for x_value in T.serial(10): static_capture(A, B) @Ts.prim_func(private=True) - def expected_hygienic(A: T.Buffer((1024,), "int32"), B: T.Buffer((), "int32")) -> None: + def expected_hygienic(A: T.Tensor((1024,), "int32"), B: T.Tensor((), "int32")) -> None: for x_value in range(10): B[()] = A[128] @@ -177,7 +177,7 @@ def test_tir_inline_late_binding(): it sees the current value of variables from its enclosing scope at call time.""" @Ts.prim_func(private=True) - def use_late_binding(A: T.Buffer((1024,), "int32"), B: T.Buffer((), "int32")) -> None: + def use_late_binding(A: T.Tensor((1024,), "int32"), B: T.Tensor((), "int32")) -> None: for x_value in T.serial(10): @T.inline @@ -187,7 +187,7 @@ def capture(A, B): capture(A, B) @Ts.prim_func(private=True) - def expected(A: T.Buffer((1024,), "int32"), B: T.Buffer((), "int32")) -> None: + def expected(A: T.Tensor((1024,), "int32"), B: T.Tensor((), "int32")) -> None: for x_value in range(10): B[()] = A[x_value] @@ -196,11 +196,11 @@ def expected(A: T.Buffer((1024,), "int32"), B: T.Buffer((), "int32")) -> None: def test_tir_macro_in_class(): class Object: - def __init__(self, x: T.Buffer): + def __init__(self, x: T.Tensor): self.local_x = Ts.sblock_alloc_buffer(x.shape, x.dtype) @T.inline - def load(self, x: T.Buffer): + def load(self, x: T.Tensor): N, M = T.meta_var(self.local_x.shape) for i, j in T.grid(N, M): with Ts.sblock("update"): @@ -208,14 +208,14 @@ def load(self, x: T.Buffer): self.local_x[vi, vj] = x[vi, vj] @Ts.prim_func(private=True) - def func_w_macro(A: T.Buffer([128, 128])): + def func_w_macro(A: T.Tensor([128, 128])): o1 = T.meta_var(Object(A)) o1.load(A) o2 = T.meta_var(Object(A)) o2.load(o1.local_x) @Ts.prim_func(private=True) - def func_no_macro(A: T.Buffer([128, 128])): + def func_no_macro(A: T.Tensor([128, 128])): local_a = Ts.sblock_alloc_buffer([128, 128]) N, M = local_a.shape for i, j in T.grid(N, M): @@ -236,12 +236,12 @@ def test_tir_starred_expression(): dims = (128, 128) @Ts.prim_func(private=True) - def starred(A: T.Buffer([128, *dims], "int32")) -> None: + def starred(A: T.Tensor([128, *dims], "int32")) -> None: for i, j, k in T.grid(128, *dims): A[i, j, k] = T.int32(1) @Ts.prim_func(private=True) - def non_starred(A: T.Buffer([128, 128, 128], "int32")) -> None: + def non_starred(A: T.Tensor([128, 128, 128], "int32")) -> None: for i, j, k in T.grid(128, 128, 128): A[i, j, k] = T.int32(1) @@ -252,12 +252,12 @@ def test_tir_dynamic_for_loop(): dims = (128, 128) @Ts.prim_func(private=True) - def starred(A: T.Buffer([128, *dims], "int32")) -> None: + def starred(A: T.Tensor([128, *dims], "int32")) -> None: for (*iters,) in T.grid(*A.shape): A[iters] = T.int32(1) @Ts.prim_func(private=True) - def non_starred(A: T.Buffer([128, 128, 128], "int32")) -> None: + def non_starred(A: T.Tensor([128, 128, 128], "int32")) -> None: for i, j, k in T.grid(128, 128, 128): A[i, j, k] = T.int32(1) @@ -268,7 +268,7 @@ def test_tir_starred_for_loop(): dims = (128, 128) @Ts.prim_func(private=True) - def starred(A: T.Buffer([*dims, 128], "int32"), B: T.Buffer(dims, "int32")): + def starred(A: T.Tensor([*dims, 128], "int32"), B: T.Tensor(dims, "int32")): for *spatial, reduction in T.grid(*A.shape): with Ts.sblock("reduce"): with Ts.init(): @@ -276,7 +276,7 @@ def starred(A: T.Buffer([*dims, 128], "int32"), B: T.Buffer(dims, "int32")): B[spatial] = B[spatial] + A[(*spatial, reduction)] @Ts.prim_func(private=True) - def non_starred(A: T.Buffer([128, 128, 128], "int32"), B: T.Buffer([128, 128], "int32")): + def non_starred(A: T.Tensor([128, 128, 128], "int32"), B: T.Tensor([128, 128], "int32")): for i, j, k in T.grid(128, 128, 128): with Ts.sblock("reduce"): with Ts.init(): @@ -290,12 +290,12 @@ def test_tir_builtin_expression(): dims = (128, 128) @Ts.prim_func(private=True) - def with_builtin(A: T.Buffer([len(dims), *dims], "int32")) -> None: + def with_builtin(A: T.Tensor([len(dims), *dims], "int32")) -> None: for i, j, k in T.grid(*A.shape): A[i, j, k] = T.int32(1 + len(A.shape)) @Ts.prim_func(private=True) - def evaluated(A: T.Buffer((2, 128, 128), "int32")): + def evaluated(A: T.Tensor((2, 128, 128), "int32")): for i, j, k in T.grid(2, 128, 128): A[i, j, k] = 4 @@ -334,14 +334,14 @@ def operation(A, idx): A[v] = A[v] * T.float32(2) @Ts.prim_func(private=True) - def func_w_macro(A: T.Buffer([10])) -> None: + def func_w_macro(A: T.Tensor([10])) -> None: for i in T.serial(0, 10): operation(A, i) operation(A, i) operation(A, i) @Ts.prim_func(private=True) - def expected(A: T.Buffer([10])) -> None: + def expected(A: T.Tensor([10])) -> None: for i in T.serial(0, 10): with Ts.sblock("op"): v = Ts.axis.remap("S", [i]) @@ -366,17 +366,17 @@ def test_prim_func_closure_shape(): def f(M=16): @Ts.prim_func - def func(A: T.Buffer((M,), "float32")): + def func(A: T.Tensor((M,), "float32")): T.evaluate(0) return func @Ts.prim_func - def expected_16(A: T.Buffer((16,), "float32")): + def expected_16(A: T.Tensor((16,), "float32")): T.evaluate(0) @Ts.prim_func - def expected_32(A: T.Buffer((32,), "float32")): + def expected_32(A: T.Tensor((32,), "float32")): T.evaluate(0) tvm.ir.assert_structural_equal(_normalize(f(16)), _normalize(expected_16)) @@ -388,17 +388,17 @@ def test_prim_func_closure_dtype(): def f(dtype="float32"): @Ts.prim_func - def func(A: T.Buffer((16,), dtype)): + def func(A: T.Tensor((16,), dtype)): T.evaluate(0) return func @Ts.prim_func - def expected_f32(A: T.Buffer((16,), "float32")): + def expected_f32(A: T.Tensor((16,), "float32")): T.evaluate(0) @Ts.prim_func - def expected_f16(A: T.Buffer((16,), "float16")): + def expected_f16(A: T.Tensor((16,), "float16")): T.evaluate(0) tvm.ir.assert_structural_equal(_normalize(f("float32")), _normalize(expected_f32)) @@ -417,7 +417,7 @@ def test_prim_func_nested_closure(): def outer(M=16): def middle(N=8): @Ts.prim_func - def func(A: T.Buffer((M, N), "float32")): + def func(A: T.Tensor((M, N), "float32")): T.evaluate(0) return func @@ -425,11 +425,11 @@ def func(A: T.Buffer((M, N), "float32")): return middle() @Ts.prim_func - def expected_16_8(A: T.Buffer((16, 8), "float32")): + def expected_16_8(A: T.Tensor((16, 8), "float32")): T.evaluate(0) @Ts.prim_func - def expected_32_8(A: T.Buffer((32, 8), "float32")): + def expected_32_8(A: T.Tensor((32, 8), "float32")): T.evaluate(0) tvm.ir.assert_structural_equal(_normalize(outer(16)), _normalize(expected_16_8)) @@ -443,17 +443,17 @@ def f(M=16): @I.ir_module class Mod: @Ts.prim_func - def main(A: T.Buffer((M,), "float32")): + def main(A: T.Tensor((M,), "float32")): T.evaluate(0) return Mod @Ts.prim_func - def expected_16(A: T.Buffer((16,), "float32")): + def expected_16(A: T.Tensor((16,), "float32")): T.evaluate(0) @Ts.prim_func - def expected_32(A: T.Buffer((32,), "float32")): + def expected_32(A: T.Tensor((32,), "float32")): T.evaluate(0) tvm.ir.assert_structural_equal(_normalize(f(16)["main"]), _normalize(expected_16)) @@ -465,17 +465,17 @@ def test_mixed_closure_usage(): def f(M=16): @Ts.prim_func - def func(A: T.Buffer((M,), "float32")): + def func(A: T.Tensor((M,), "float32")): T.evaluate(M) return func @Ts.prim_func - def expected_16(A: T.Buffer((16,), "float32")): + def expected_16(A: T.Tensor((16,), "float32")): T.evaluate(16) @Ts.prim_func - def expected_32(A: T.Buffer((32,), "float32")): + def expected_32(A: T.Tensor((32,), "float32")): T.evaluate(32) tvm.ir.assert_structural_equal(_normalize(f(16)), _normalize(expected_16)) @@ -486,7 +486,7 @@ def expected_32(A: T.Buffer((32,), "float32")): @Ts.prim_func -def matmul(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j, k in T.grid(128, 128, 128): with Ts.sblock("update"): vi, vj, vk = Ts.axis.remap("SSR", [i, j, k]) @@ -543,9 +543,9 @@ def shared_16x16_to_ldmatrix_32x8_layout(i, j): @Ts.prim_func def mma_sync_m16n16k16_desc( - A: T.Buffer((32, 8), "float16", align=64, offset_factor=16, scope="warp"), - B: T.Buffer((32, 8), "float16", align=64, offset_factor=16, scope="warp"), - C: T.Buffer((32, 8), "float16", align=64, offset_factor=16, scope="warp"), + A: T.Tensor((32, 8), "float16", align=64, offset_factor=16, scope="warp"), + B: T.Tensor((32, 8), "float16", align=64, offset_factor=16, scope="warp"), + C: T.Tensor((32, 8), "float16", align=64, offset_factor=16, scope="warp"), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:32, 0:8], A[0:32, 0:8], B[0:32, 0:8]) @@ -570,9 +570,9 @@ def mma_sync_m16n16k16_desc( @Ts.prim_func def mma_sync_m16n16k16_desc_manual( - A: T.Buffer((32, 8), "float16", align=64, offset_factor=16, scope="warp"), - B: T.Buffer((32, 8), "float16", align=64, offset_factor=16, scope="warp"), - C: T.Buffer((32, 8), "float16", align=64, offset_factor=16, scope="warp"), + A: T.Tensor((32, 8), "float16", align=64, offset_factor=16, scope="warp"), + B: T.Tensor((32, 8), "float16", align=64, offset_factor=16, scope="warp"), + C: T.Tensor((32, 8), "float16", align=64, offset_factor=16, scope="warp"), ) -> None: with Ts.sblock("root"): Ts.reads(C[0:32, 0:8], A[0:32, 0:8], B[0:32, 0:8]) diff --git a/tests/python/s_tir/script/test_s_tir_script_namespace.py b/tests/python/s_tir/script/test_s_tir_script_namespace.py index 182b93aaf795..18c0d6b94b02 100644 --- a/tests/python/s_tir/script/test_s_tir_script_namespace.py +++ b/tests/python/s_tir/script/test_s_tir_script_namespace.py @@ -37,7 +37,7 @@ def test_shared_operations_and_aliases(): n = S.dynamic("n") @S.prim_func - def shared(A: S.Buffer((n,), "float32")): + def shared(A: S.Tensor((n,), "float32")): for i in T.serial(A.shape[0]): with S.sblock("copy"): v = Ts.axis.spatial(A.shape[0], i) @@ -45,7 +45,7 @@ def shared(A: S.Buffer((n,), "float32")): assert shared.attrs["s_tir"] assert shared.params[0].ty.layout is None - assert Ts.Buffer is T.Buffer + assert Ts.Tensor is T.Tensor assert Ts.serial is T.serial assert Ts.bind is T.bind assert Ts.tile is T.tile @@ -62,14 +62,14 @@ def test_mixed_module_roundtrip(): @I.ir_module class Mixed: @Ts.prim_func - def scheduled(A: Ts.Buffer((4,), "float32")): + def scheduled(A: Ts.Tensor((4,), "float32")): for i in Ts.serial(4): with Ts.sblock("copy"): v = Ts.axis.spatial(4, i) A[v] = 1.0 @T.prim_func - def direct(A: T.Buffer((4,), "float32")): + def direct(A: T.Tensor((4,), "float32")): for i in T.serial(4): A[i] = 2.0 @@ -100,7 +100,7 @@ def reject_s_tir_analysis(*args, **kwargs): @I.ir_module class Direct: @T.prim_func - def main(A: T.Buffer((4,), "int32")): + def main(A: T.Tensor((4,), "int32")): for i in T.serial(4): A[i] = A[i] + 3 @@ -202,7 +202,7 @@ def fill(A): A[i] = Ts.float32(1) @Ts.prim_func(private=True) - def scheduled(A: Ts.Buffer((4,), "float32")): + def scheduled(A: Ts.Tensor((4,), "float32")): fill(A) assert scheduled.attrs["s_tir"] @@ -236,7 +236,7 @@ def test_import_order(first): "from tvm.script import s_tir as Ts, tirx as T\n" "from tvm.script.parser import s_tir as parser\n" "assert Ts is direct is parser\n" - "assert Ts.Buffer is T.Buffer\n" + "assert Ts.Tensor is T.Tensor\n" "from tvm.s_tir.script import ir_builder as owned\n" "from tvm.script.ir_builder import s_tir as alias\n" "import tvm.s_tir.script.ir_builder.frame as owned_frame\n" diff --git a/tests/python/s_tir/script/test_s_tir_script_printer.py b/tests/python/s_tir/script/test_s_tir_script_printer.py index 3554a2ea61d1..82fecb2eb383 100644 --- a/tests/python/s_tir/script/test_s_tir_script_printer.py +++ b/tests/python/s_tir/script/test_s_tir_script_printer.py @@ -42,7 +42,7 @@ def _assert_print(obj, expected): @pytest.mark.skipif(sys.version_info < (3, 12), reason="PEP 695 requires Python 3.12") def test_prim_func_symbolic_buffer_param_roundtrip(): n = tirx.Var("n", "int32") - A = tirx.decl_buffer(shape=[n + 1, n], dtype="float32", name="A", layout=None) + A = tirx.decl_tensor(shape=[n + 1, n], dtype="float32", name="A", layout=None) func = ( tirx.PrimFunc(params=[A], body=tirx.Evaluate(n)) .with_attr("global_symbol", "main") @@ -50,7 +50,7 @@ def test_prim_func_symbolic_buffer_param_roundtrip(): ) source = func.script() - assert "T.Buffer((n + 1, n)" in source + assert "T.Tensor((n + 1, n)" in source assert source.index("def main[n: T.int32](") < source.index("T.evaluate(n)") tvm.ir.assert_structural_equal( tvm.script.from_source( @@ -63,7 +63,7 @@ def test_prim_func_symbolic_buffer_param_roundtrip(): @pytest.mark.skipif(sys.version_info < (3, 12), reason="PEP 695 requires Python 3.12") def test_prim_func_compound_buffer_shape_first_use_roundtrip(): n = tirx.Var("n", "int32") - A = tirx.decl_buffer(shape=[tirx.max(n, 1)], dtype="float32", name="A", layout=None) + A = tirx.decl_tensor(shape=[tirx.max(n, 1)], dtype="float32", name="A", layout=None) func = ( tirx.PrimFunc(params=[A], body=tirx.Evaluate(n)) .with_attr("global_symbol", "main") @@ -71,7 +71,7 @@ def test_prim_func_compound_buffer_shape_first_use_roundtrip(): ) source = func.script() - assert "T.Buffer((T.max(n, 1),)" in source + assert "T.Tensor((T.max(n, 1),)" in source assert source.index("def main[n: T.int32](") < source.index("T.evaluate(n)") tvm.ir.assert_structural_equal( tvm.script.from_source( @@ -83,7 +83,7 @@ def test_prim_func_compound_buffer_shape_first_use_roundtrip(): def test_prim_func_symbolic_alloc_buffer_roundtrip(): size = tirx.Var("size", "int32") - buf = tirx.decl_buffer(shape=[size], dtype="float32", name="buf", layout=None) + buf = tirx.decl_tensor(shape=[size], dtype="float32", name="buf", layout=None) func = tirx.PrimFunc( params=[], body=tirx.SeqStmt( @@ -91,7 +91,7 @@ def test_prim_func_symbolic_alloc_buffer_roundtrip(): tvm.tirx.Bind( buf, tvm.ir.Call( - "tirx.alloc_buffer", + "tirx.alloc_tensor", [ tvm.ir.Tuple(buf.shape), tvm.ir.DataTypeImm(tvm.DataType(buf.dtype)), @@ -107,7 +107,7 @@ def test_prim_func_symbolic_alloc_buffer_roundtrip(): ).with_attr("s_tir", True) source = func.script() - assert 'T.alloc_buffer((size,), "float32")' in source + assert 'T.alloc_tensor((size,), "float32")' in source tvm.ir.assert_structural_equal( tvm.script.from_source( source, @@ -119,8 +119,8 @@ def test_prim_func_symbolic_alloc_buffer_roundtrip(): def test_prim_func(): - A = tirx.decl_buffer(shape=[128, 128], dtype="float32", name="A") - B = tirx.decl_buffer(shape=[256, 256], dtype="float32", name="B") + A = tirx.decl_tensor(shape=[128, 128], dtype="float32", name="A") + B = tirx.decl_tensor(shape=[256, 256], dtype="float32", name="B") func = ( tirx.PrimFunc( params=[A, B], @@ -139,14 +139,14 @@ def test_prim_func(): # from tvm.script import tirx as T @Ts.prim_func -def main(A: T.Buffer((128, 128), "float32", layout="default"), B: T.Buffer((256, 256), "float32", layout="default")): +def main(A: T.Tensor((128, 128), "float32", layout="default"), B: T.Tensor((256, 256), "float32", layout="default")): T.evaluate(0)""", ) def test_prim_func_buffer_data_use(): - A = tirx.decl_buffer(shape=[128, 128], dtype="float32", name="A") - B = tirx.decl_buffer(shape=[256, 256], dtype="float32", name="B") + A = tirx.decl_tensor(shape=[128, 128], dtype="float32", name="A") + B = tirx.decl_tensor(shape=[256, 256], dtype="float32", name="B") func = ( tirx.PrimFunc( params=[A, B], @@ -165,16 +165,16 @@ def test_prim_func_buffer_data_use(): # from tvm.script import tirx as T @Ts.prim_func -def main(A: T.Buffer((128, 128), "float32", layout="default"), B: T.Buffer((256, 256), "float32", layout="default")): +def main(A: T.Tensor((128, 128), "float32", layout="default"), B: T.Tensor((256, 256), "float32", layout="default")): T.evaluate(A.data) """, ) def test_prim_func_buffer_data_argument_is_scope_hint(): - buffer_data = tirx.decl_buffer(shape=[128, 128], dtype="float32", name="A").data - A = tirx.decl_buffer(shape=[128, 128], dtype="float32", name="A", data=buffer_data) - B = tirx.decl_buffer(shape=[256, 256], dtype="float32", name="B", data=buffer_data) + buffer_data = tirx.decl_tensor(shape=[128, 128], dtype="float32", name="A").data + A = tirx.decl_tensor(shape=[128, 128], dtype="float32", name="A", data=buffer_data) + B = tirx.decl_tensor(shape=[256, 256], dtype="float32", name="B", data=buffer_data) func = ( tirx.PrimFunc( params=[A, B], @@ -193,7 +193,7 @@ def test_prim_func_buffer_data_argument_is_scope_hint(): # from tvm.script import tirx as T @Ts.prim_func -def main(A: T.Buffer((128, 128), "float32", layout="default"), B: T.Buffer((256, 256), "float32", layout="default")): +def main(A: T.Tensor((128, 128), "float32", layout="default"), B: T.Tensor((256, 256), "float32", layout="default")): T.evaluate(0) """, ) @@ -255,8 +255,8 @@ def test_block(): def test_match_buffer_region(): - src = tirx.decl_buffer((128, 128), "float32", name="src") - tgt = tirx.decl_buffer((64, 64), "float32", name="tgt") + src = tirx.decl_tensor((128, 128), "float32", name="src") + tgt = tirx.decl_tensor((64, 64), "float32", name="tgt") obj = s_tir.MatchBufferRegion( tgt, tirx.BufferRegion( @@ -270,7 +270,7 @@ def test_match_buffer_region(): _assert_print( obj, """ -src = T.Var("src", T.Buffer((128, 128), "float32", layout="default")) +src = T.Var("src", T.Tensor((128, 128), "float32", layout="default")) tgt = Ts.match_buffer(src[64:128, 64:128], (64, 64), "float32", layout="default") """, ) @@ -388,8 +388,8 @@ def main(): def test_private_primfunc(): - A = tirx.decl_buffer(shape=[128, 128], dtype="float32", name="A") - B = tirx.decl_buffer(shape=[256, 256], dtype="float32", name="B") + A = tirx.decl_tensor(shape=[128, 128], dtype="float32", name="A") + B = tirx.decl_tensor(shape=[256, 256], dtype="float32", name="B") func = tirx.PrimFunc( params=[A, B], ret_type=None, @@ -404,7 +404,7 @@ def test_private_primfunc(): # from tvm.script import tirx as T @Ts.prim_func(private=True) -def main(A: T.Buffer((128, 128), "float32", layout="default"), B: T.Buffer((256, 256), "float32", layout="default")): +def main(A: T.Tensor((128, 128), "float32", layout="default"), B: T.Tensor((256, 256), "float32", layout="default")): T.evaluate(0)""", ) @@ -413,7 +413,7 @@ def test_prim_func_different_symbol(): from tvm.script import tirx as TB @Ts.prim_func - def main(A: TB.Buffer((128, 128), "float32"), B: TB.Buffer((256, 256), "float32")): + def main(A: TB.Tensor((128, 128), "float32"), B: TB.Tensor((256, 256), "float32")): TB.func_attr({"global_symbol": "func"}) TB.evaluate(0) @@ -424,7 +424,7 @@ def main(A: TB.Buffer((128, 128), "float32"), B: TB.Buffer((256, 256), "float32" # from tvm.script import tirx as T @Ts.prim_func -def func(A: T.Buffer((128, 128), "float32"), B: T.Buffer((256, 256), "float32")): +def func(A: T.Tensor((128, 128), "float32"), B: T.Tensor((256, 256), "float32")): T.evaluate(0)""" _assert_print(main, expected_output) @@ -488,7 +488,7 @@ def test_predicated_load_store(): from tvm.script import tirx as TB @Ts.prim_func - def main(A: TB.Buffer((128, 128), "float32"), B: TB.Buffer((256, 256), "float32")): + def main(A: TB.Tensor((128, 128), "float32"), B: TB.Tensor((256, 256), "float32")): TB.func_attr({"global_symbol": "func"}) a_load = TB.meta_var( TB.call_intrin( @@ -519,7 +519,7 @@ def main(A: TB.Buffer((128, 128), "float32"), B: TB.Buffer((256, 256), "float32" # from tvm.script import tirx as T @Ts.prim_func -def func(A: T.Buffer((128, 128), "float32"), B: T.Buffer((256, 256), "float32")): +def func(A: T.Tensor((128, 128), "float32"), B: T.Tensor((256, 256), "float32")): a_load: T.let[T.float32x4] = T.masked_load("float32x4", A, 0, T.Ramp(0, 4, 4), T.Broadcast(T.bool(False), 4)) T.masked_store(A, a_load, 0, T.Ramp(0, 2, 4), T.Broadcast(T.bool(False), 4))""" _assert_print(main, expected_output) @@ -529,8 +529,8 @@ def test_predicated_buffer_load_store(): a = tirx.Var("a", "handle") b = tirx.Var("b", "handle") buffers = { - a: tirx.decl_buffer(shape=[128, 128], dtype="float32", name="A"), - b: tirx.decl_buffer(shape=[256, 256], dtype="float32", name="B"), + a: tirx.decl_tensor(shape=[128, 128], dtype="float32", name="A"), + b: tirx.decl_tensor(shape=[256, 256], dtype="float32", name="B"), } buffer_load = tirx.call_intrin( "float32x4", @@ -564,7 +564,7 @@ def test_predicated_buffer_load_store(): # from tvm.script import tirx as T @Ts.prim_func(private=True) -def main(A: T.Buffer((128, 128), "float32", layout="default"), B: T.Buffer((256, 256), "float32", layout="default")): +def main(A: T.Tensor((128, 128), "float32", layout="default"), B: T.Tensor((256, 256), "float32", layout="default")): T.masked_store(A, T.masked_load("float32x4", B, 0, T.Ramp(0, 4, 4), T.Broadcast(T.bool(False), 4)), 0, T.Ramp(0, 2, 4), T.Broadcast(T.bool(False), 4))""" _assert_print(func, expected_output) @@ -573,7 +573,7 @@ def test_predicated_scalable_load_store(): from tvm.script import tirx as TB @Ts.prim_func - def main(A: TB.Buffer((128, 128), "float32"), B: TB.Buffer((256, 256), "float32")): + def main(A: TB.Tensor((128, 128), "float32"), B: TB.Tensor((256, 256), "float32")): TB.func_attr({"global_symbol": "func"}) mask = TB.meta_var(TB.get_active_lane_mask("uint1xvscalex4", 0, 13)) a_load = TB.meta_var( @@ -594,7 +594,7 @@ def main(A: TB.Buffer((128, 128), "float32"), B: TB.Buffer((256, 256), "float32" # from tvm.script import tirx as T @Ts.prim_func -def func(A: T.Buffer((128, 128), "float32"), B: T.Buffer((256, 256), "float32")): +def func(A: T.Tensor((128, 128), "float32"), B: T.Tensor((256, 256), "float32")): mask: T.let["uint1xvscalex4"] = T.get_active_lane_mask("uint1xvscalex4", 0, 13) a_load: T.let["float32xvscalex4"] = T.masked_load("float32xvscalex4", A, 0, T.Ramp(0, 4, T.vscale() * 4), mask) T.masked_store(A, a_load, 0, T.Ramp(0, 2, T.vscale() * 4), mask) @@ -607,11 +607,11 @@ def test_masked_load_prevents_scalar_allocation_init_fusion(): @Ts.prim_func def main(): - A = TB.alloc_buffer((1,), "float32x4") + A = TB.alloc_tensor((1,), "float32x4") A[0] = TB.masked_load("float32x4", A, 0, TB.Broadcast(TB.bool(True), 4)) source = main.script() - assert "A = T.alloc_buffer" in source + assert "A = T.alloc_tensor" in source assert "A[0] = T.masked_load" in source tvm.ir.assert_structural_equal( tvm.script.from_source( @@ -625,7 +625,7 @@ def test_vload_with_explicit_scalable_data_type(): from tvm.script import tirx as TB @Ts.prim_func - def main(A: TB.Buffer((128,), "float32"), B: TB.Buffer((128,), "float32")): + def main(A: TB.Tensor((128,), "float32"), B: TB.Tensor((128,), "float32")): B[0 : TB.vscale() * 4] = A.vload([TB.Ramp(0, 1, TB.vscale() * 4)], dtype="float32xvscalex4") expected_output = """ @@ -635,7 +635,7 @@ def main(A: TB.Buffer((128,), "float32"), B: TB.Buffer((128,), "float32")): # from tvm.script import tirx as T @Ts.prim_func -def main(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): +def main(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): B[0:T.vscale() * 4] = A[T.Ramp(0, 1, T.vscale() * 4)]""" _assert_print(main, expected_output) @@ -644,7 +644,7 @@ def test_vectorize_llvm_pure_intrin(): from tvm.script import tirx as TB @Ts.prim_func - def main(A: TB.Buffer((4,), "float32"), B: TB.Buffer((4,), "float32")): + def main(A: TB.Tensor((4,), "float32"), B: TB.Tensor((4,), "float32")): A[TB.Ramp(0, 1, 4)] = TB.call_llvm_pure_intrin( "float32x4", "llvm.sqrt", B[TB.Ramp(0, 1, 4)] ) @@ -656,7 +656,7 @@ def main(A: TB.Buffer((4,), "float32"), B: TB.Buffer((4,), "float32")): # from tvm.script import tirx as T @Ts.prim_func -def main(A: T.Buffer((4,), "float32"), B: T.Buffer((4,), "float32")): +def main(A: T.Tensor((4,), "float32"), B: T.Tensor((4,), "float32")): A[0:4] = T.call_llvm_pure_intrin("float32x4", "llvm.sqrt", B[T.Ramp(0, 1, 4)])""" _assert_print(main, expected_output) @@ -665,7 +665,7 @@ def test_func_with_loop_jumps(): from tvm.script import tirx as TB @Ts.prim_func - def main(A: TB.Buffer((4,), "float32"), B: TB.Buffer((4,), "float32")): + def main(A: TB.Tensor((4,), "float32"), B: TB.Tensor((4,), "float32")): for i in range(1000): if i % 13 == 0: A[1] = A[1] + 1 @@ -680,7 +680,7 @@ def main(A: TB.Buffer((4,), "float32"), B: TB.Buffer((4,), "float32")): # from tvm.script import tirx as T @Ts.prim_func -def main(A: T.Buffer((4,), "float32"), B: T.Buffer((4,), "float32")): +def main(A: T.Tensor((4,), "float32"), B: T.Tensor((4,), "float32")): for i in range(1000): if i % 13 == 0: A[1] = A[1] + T.float32(1.0) @@ -697,19 +697,19 @@ def opt_gemm_lower(): class Module: @Ts.prim_func def mmult( - A_1: T.Buffer([16384], elem_offset=0, align=64, offset_factor=1), - B_1: T.Buffer([1024, 1024], elem_offset=0, align=64, offset_factor=1), - C_1: T.Buffer([16384], elem_offset=0, align=64, offset_factor=1), + A_1: T.Tensor([16384], elem_offset=0, align=64, offset_factor=1), + B_1: T.Tensor([1024, 1024], elem_offset=0, align=64, offset_factor=1), + C_1: T.Tensor([16384], elem_offset=0, align=64, offset_factor=1), ) -> None: # function attr dict T.func_attr({"tirx.noalias": True}) # body - packedB = T.alloc_buffer((32768,)) + packedB = T.alloc_tensor((32768,)) for x in T.parallel(0, 32): for y in T.serial(0, 1024): packedB[T.ramp(((x * 32768) + (y * 32)), 1, 32)] = B_1[y, T.ramp(x * 32, 1, 32)] for x_outer in T.parallel(0, 32): - C_global = T.alloc_buffer((1024,)) + C_global = T.alloc_tensor((1024,)) for y_outer in T.serial(0, 32): for x_c_init in T.serial(0, 32): C_global[T.ramp((x_c_init * 32), 1, 32)] = T.broadcast(T.float32(0), 32) @@ -754,16 +754,16 @@ def opt_conv_tensorcore_lower(): @Ts.prim_func def func( - A: T.Buffer((16, 14, 14, 16, 16, 16), "float16"), - W: T.Buffer((3, 3, 16, 32, 16, 16), "float16"), - Conv: T.Buffer((16, 14, 14, 32, 16, 16), "float32"), + A: T.Tensor((16, 14, 14, 16, 16, 16), "float16"), + W: T.Tensor((3, 3, 16, 32, 16, 16), "float16"), + Conv: T.Tensor((16, 14, 14, 32, 16, 16), "float32"), ) -> None: # function attr dict T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) # body - A_1 = T.decl_buffer([12845056], dtype="float16", data=A.data) - W_1 = T.decl_buffer([1179648], dtype="float16", data=W.data) - Conv_1 = T.decl_buffer([25690112], data=Conv.data) + A_1 = T.decl_tensor([12845056], dtype="float16", data=A.data) + W_1 = T.decl_tensor([1179648], dtype="float16", data=W.data) + Conv_1 = T.decl_tensor([25690112], data=Conv.data) bx = T.env_thread("blockIdx.x") by = T.env_thread("blockIdx.y") bz = T.env_thread("blockIdx.z") @@ -771,11 +771,11 @@ def func( ty = T.env_thread("threadIdx.y") tz = T.env_thread("threadIdx.z") T.launch_thread(bz, 196) - Conv_wmma_accumulator = T.alloc_buffer((2048,), scope="wmma.accumulator") - Apad_shared = T.alloc_buffer((12288,), "float16", scope="shared") - W_shared = T.alloc_buffer((12288,), "float16", scope="shared") - Apad_shared_wmma_matrix_a = T.alloc_buffer((512,), "float16", scope="wmma.matrix_a") - W_shared_wmma_matrix_b = T.alloc_buffer((1024,), "float16", scope="wmma.matrix_b") + Conv_wmma_accumulator = T.alloc_tensor((2048,), scope="wmma.accumulator") + Apad_shared = T.alloc_tensor((12288,), "float16", scope="shared") + W_shared = T.alloc_tensor((12288,), "float16", scope="shared") + Apad_shared_wmma_matrix_a = T.alloc_tensor((512,), "float16", scope="wmma.matrix_a") + W_shared_wmma_matrix_b = T.alloc_tensor((1024,), "float16", scope="wmma.matrix_b") T.launch_thread(bx, 2) T.launch_thread(by, 4) T.launch_thread(ty, 4) @@ -1055,7 +1055,7 @@ def opt_conv_tensorcore_mod_host(): @Ts.prim_func def opt_conv_tensorcore_mod_host( args: T.handle, - arg_type_ids: T.Buffer((3,), "int32"), + arg_type_ids: T.Tensor((3,), "int32"), num_args: T.int32, out_ret_value: T.handle, out_ret_tcode: T.handle, @@ -1072,7 +1072,7 @@ def opt_conv_tensorcore_mod_host( ) # body stack_tcode_data: T.let[T.handle("int32")] = T.tvm_stack_alloca("arg_tcode", 10) - stack_tcode = T.decl_buffer([9], "int32", data=stack_tcode_data) + stack_tcode = T.decl_tensor([9], "int32", data=stack_tcode_data) stack_value: T.let[T.handle] = T.tvm_stack_alloca("arg_value", 10) assert num_args == 3, "default_function: num_args should be 3" arg0: T.let[T.handle] = T.tvm_struct_get(args, 0, 12, dtype="handle") @@ -1085,11 +1085,11 @@ def opt_conv_tensorcore_mod_host( arg0_shape_data: T.let[T.handle("int64")] = T.tvm_struct_get( arg0, 0, 2, dtype=T.handle("int64").ty ) - arg0_shape = T.decl_buffer([6], "int64", data=arg0_shape_data) + arg0_shape = T.decl_tensor([6], "int64", data=arg0_shape_data) arg0_strides_data: T.let[T.handle("int64")] = T.tvm_struct_get( arg0, 0, 3, dtype=T.handle("int64").ty ) - arg0_strides = T.decl_buffer([6], "int64", data=arg0_strides_data) + arg0_strides = T.decl_tensor([6], "int64", data=arg0_strides_data) dev_id: T.let[T.int32] = T.tvm_struct_get(arg0, 0, 9, dtype="int32") @@ -1178,7 +1178,7 @@ def opt_conv_tensorcore_mod_host( def select(): @Ts.prim_func - def select(A: T.Buffer((), "float32")) -> None: + def select(A: T.Tensor((), "float32")) -> None: A[()] = T.Select(True, 1, 2) return select @@ -1186,7 +1186,7 @@ def select(A: T.Buffer((), "float32")) -> None: def minmax(): @Ts.prim_func - def minmax(A: T.Buffer((), "float32")) -> None: + def minmax(A: T.Tensor((), "float32")) -> None: A[()] = T.min(1, 2) A[()] = T.max(1, 2) @@ -1195,7 +1195,7 @@ def minmax(A: T.Buffer((), "float32")) -> None: def abs(): @Ts.prim_func - def abs(A: T.Buffer((128, 128), "float32")) -> None: + def abs(A: T.Tensor((128, 128), "float32")) -> None: for i, j in T.grid(128, 128): with Ts.sblock("A"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -1206,7 +1206,7 @@ def abs(A: T.Buffer((128, 128), "float32")) -> None: def constant_folding(): @Ts.prim_func - def constant_folding(A: T.Buffer((), "float32")) -> None: + def constant_folding(A: T.Tensor((), "float32")) -> None: A[()] = T.min(2.2, 5.2) A[()] = T.max(T.float32(2.2), T.float32(T.float32(5.2))) A[()] = T.min(2.2, 5.0) @@ -1230,7 +1230,7 @@ def simplify_bracket() -> None: def var_with_same_name(): @Ts.prim_func - def var_with_same_name(A: T.Buffer((16, 16), "float32")) -> None: + def var_with_same_name(A: T.Tensor((16, 16), "float32")) -> None: for i, j in T.grid(16, 16): with Ts.sblock(): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -1265,8 +1265,8 @@ def test_same_name_var(): def primfunc_with_allocate_annotations(): @Ts.prim_func def primfunc_with_allocate_annotations( - placeholder_29: T.Buffer([802816], dtype="uint8", elem_offset=0, align=64, offset_factor=1), - T_cast_7: T.Buffer([200704], dtype="int16", elem_offset=0, align=64, offset_factor=1), + placeholder_29: T.Tensor([802816], dtype="uint8", elem_offset=0, align=64, offset_factor=1), + T_cast_7: T.Tensor([200704], dtype="int16", elem_offset=0, align=64, offset_factor=1), ) -> None: # function attr dict T.func_attr( @@ -1274,7 +1274,7 @@ def primfunc_with_allocate_annotations( ) # body - tensor_2 = T.alloc_buffer((200704,), "uint8", annotations={"attr1_key": "attr1_value"}) + tensor_2 = T.alloc_tensor((200704,), "uint8", annotations={"attr1_key": "attr1_value"}) for ax0_ax1_fused_4 in T.serial(0, 56): for ax2_4 in T.serial(0, 56): for ax3_init in T.serial(0, 64): @@ -1317,14 +1317,14 @@ def primfunc_with_allocate_annotations( def comm_reducer_single_reduce_group(): @Ts.prim_func def comm_reducer_single_reduce_group( - A: T.Buffer([16384], dtype="float32"), b: T.handle + A: T.Tensor([16384], dtype="float32"), b: T.handle ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) threadIdx_x = T.env_thread("threadIdx.x") for i in T.serial(0, 128): T.launch_thread(threadIdx_x, 128) - reduce_temp0 = T.alloc_buffer((1,), scope="local") + reduce_temp0 = T.alloc_tensor((1,), scope="local") with T.attr( T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), "reduce_scope", T.int32(0) ): @@ -1344,14 +1344,14 @@ def comm_reducer_single_reduce_group( def comm_reducer_multiple_reduce_groups(): @Ts.prim_func def comm_reducer_multiple_reduce_groups( - A: T.Buffer([16384], dtype="float32"), b: T.handle + A: T.Tensor([16384], dtype="float32"), b: T.handle ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) threadIdx_x = T.env_thread("threadIdx.x") for i in T.serial(0, 128): T.launch_thread(threadIdx_x, 128) - reduce_temp0 = T.alloc_buffer((1,), scope="local") + reduce_temp0 = T.alloc_tensor((1,), scope="local") with T.attr( T.comm_reducer( lambda x0, x1, y0, y1: ( @@ -1499,7 +1499,7 @@ def func_T_ptr_let_statement( ) -> None: # The T.Ptr declaration in the parameter list should parse # correctly, and should be usable as the data pointer in a buffer. - arg_type_ids = T.decl_buffer([2], dtype="int32", data=arg_type_ids_handle) + arg_type_ids = T.decl_tensor([2], dtype="int32", data=arg_type_ids_handle) arg0: T.let[T.handle] = T.tvm_struct_get(args, 0, 12, dtype="handle") arg1: T.let[T.handle] = T.tvm_struct_get(args, 1, 12, dtype="handle") @@ -1515,12 +1515,12 @@ def func_T_ptr_let_statement( # this function. It should only be defined after the data pointer # has been defined, and should not be hoisted into the header of # the function as other buffer_decl statements can be. - A = T.decl_buffer([1024], dtype="float32", data=A_data) + A = T.decl_tensor([1024], dtype="float32", data=A_data) B_data: T.let[T.handle("float32")] = T.reinterpret( T.handle("float32").ty, T.tvm_struct_get(arg1, 0, 1, dtype="handle"), ) - B = T.decl_buffer([1024], dtype="float32", data=B_data) + B = T.decl_tensor([1024], dtype="float32", data=B_data) B[0] = A[0] @@ -1530,7 +1530,7 @@ def func_T_ptr_let_statement( def func_T_ptr_allocate(): @Ts.prim_func def func_T_ptr_allocate() -> None: - A = T.alloc_buffer((1024,)) + A = T.alloc_tensor((1024,)) A[0] = 0.0 return func_T_ptr_allocate @@ -1538,7 +1538,7 @@ def func_T_ptr_allocate() -> None: def llvm_intrin_call(): @Ts.prim_func - def ctpop(A: T.Buffer((16,), "uint8"), B: T.Buffer((16,), "uint8")) -> None: + def ctpop(A: T.Tensor((16,), "uint8"), B: T.Tensor((16,), "uint8")) -> None: for i in range(0, 16): with Ts.sblock("A"): vi = Ts.axis.remap( @@ -1577,8 +1577,8 @@ def string_annotation_of_special_chars(): def pointer_type(): @Ts.prim_func def func_with_ptr_type_annotations(x: T.handle("int32"), y: T.handle("int32", "shared")): - xx = T.alloc_buffer((16,), "int32") - yy = T.alloc_buffer((16,), "int32", scope="shared") + xx = T.alloc_tensor((16,), "int32") + yy = T.alloc_tensor((16,), "int32", scope="shared") a: T.let[T.handle("int32")] = T.address_of(xx[0]) b: T.let[T.handle("int32", "shared")] = T.address_of(yy[0]) T.evaluate(T.call_extern("copy", a, b, dtype="")) @@ -1589,9 +1589,9 @@ def func_with_ptr_type_annotations(x: T.handle("int32"), y: T.handle("int32", "s def buffer_ramp_access_as_slice_index(): @Ts.prim_func def buffer_ramp_access( - A: T.Buffer((128,), "float32"), - B: T.Buffer((128,), "float32"), - C: T.Buffer((128,), "float32"), + A: T.Tensor((128,), "float32"), + B: T.Tensor((128,), "float32"), + C: T.Tensor((128,), "float32"), ) -> None: for i in range(128): A[i : i + 1 : 1] = i @@ -1613,7 +1613,7 @@ def func() -> None: def scalable_vectors(): @Ts.prim_func - def func(A: T.Buffer((200,), "float32")): + def func(A: T.Tensor((200,), "float32")): A[T.Ramp(11, 2, 4 * tirx.vscale())] = T.Broadcast(125, 4 * tirx.vscale()) return func @@ -1621,7 +1621,7 @@ def func(A: T.Buffer((200,), "float32")): def predicated_buffer_load_store(): @Ts.prim_func - def func(A: T.Buffer((4,), "float32"), B: T.Buffer((8,), "float32")): + def func(A: T.Tensor((4,), "float32"), B: T.Tensor((8,), "float32")): for i_0 in range(4): load_a = T.meta_var( T.call_intrin( @@ -1715,12 +1715,12 @@ def func(out_ret_value: T.handle("void")): return func -def decl_buffer(): +def decl_tensor(): @Ts.prim_func - def func(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")) -> None: - A_flattened = T.decl_buffer(data=A.data, shape=(256,), dtype="float32") - B_flattened = T.decl_buffer(data=B.data, shape=(256,), dtype="float32") - C_alias = T.decl_buffer(data=A_flattened.data, shape=(256,), dtype="float32") + def func(A: T.Tensor((16, 16), "float32"), B: T.Tensor((16, 16), "float32")) -> None: + A_flattened = T.decl_tensor(data=A.data, shape=(256,), dtype="float32") + B_flattened = T.decl_tensor(data=B.data, shape=(256,), dtype="float32") + C_alias = T.decl_tensor(data=A_flattened.data, shape=(256,), dtype="float32") for i in range(256): B_flattened[i] = A_flattened[i] + C_alias[i] + T.float32(1.0) @@ -1729,10 +1729,10 @@ def func(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")) -> def allocate_and_decl_buffer(): @Ts.prim_func - def func(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")) -> None: - D = T.alloc_buffer((16,)) + def func(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")) -> None: + D = T.alloc_tensor((16,)) for i in range(4): - C = T.alloc_buffer((4,)) + C = T.alloc_tensor((4,)) for j in range(4): C[j] = A[i * 4 + j] + T.float32(1.0) for j in range(4): @@ -1745,8 +1745,8 @@ def func(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")) -> None: def alloc_buffer_example(): @Ts.prim_func - def func(A: T.Buffer((128,), "float32"), C: T.Buffer((128,), "float32")): - B = T.alloc_buffer((128,), "float32") + def func(A: T.Tensor((128,), "float32"), C: T.Tensor((128,), "float32")): + B = T.alloc_tensor((128,), "float32") for i in range(128): B[i] = A[i] * T.float32(2) for i in range(128): @@ -1758,7 +1758,7 @@ def func(A: T.Buffer((128,), "float32"), C: T.Buffer((128,), "float32")): def float_infinity(): @Ts.prim_func def func( - placeholder: T.Buffer((1, 512, 768), "float32"), T_isinf: T.Buffer((1, 512, 768), "bool") + placeholder: T.Tensor((1, 512, 768), "float32"), T_isinf: T.Tensor((1, 512, 768), "bool") ) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -1826,7 +1826,7 @@ def nested_boolean_expressions(): def make_ir_generator(name, expression): def inner(): @Ts.prim_func - def func(A: T.Buffer(1, "bool"), i: T.bool, j: T.bool, k: T.bool): + def func(A: T.Tensor(1, "bool"), i: T.bool, j: T.bool, k: T.bool): A[0] = expression(i, j, k) return func @@ -1842,7 +1842,7 @@ def func(A: T.Buffer(1, "bool"), i: T.bool, j: T.bool, k: T.bool): def multi_env_threads(): @Ts.prim_func - def func(A: T.Buffer(128, "float32"), C: T.Buffer(128, "float32")): + def func(A: T.Tensor(128, "float32"), C: T.Tensor(128, "float32")): B = Ts.sblock_alloc_buffer([128], dtype="float32") for i in T.thread_binding(128, thread="threadIdx.x"): B[i] = A[i] + 1.0 @@ -1872,29 +1872,29 @@ def func( ): blockIdx_x = T.launch_thread("blockIdx.x", 1) threadIdx_x = T.launch_thread("threadIdx.x", 32) - A_warp = T.alloc_buffer((1,), scope="local") - B_warp = T.alloc_buffer((1,), scope="local") - red_buf0 = T.alloc_buffer((1,), scope="local") - A_warp_1 = T.decl_buffer((32,), data=A_warp.data, scope="local") - A_1 = T.decl_buffer((32,), data=A) # A is a handle param + A_warp = T.alloc_tensor((1,), scope="local") + B_warp = T.alloc_tensor((1,), scope="local") + red_buf0 = T.alloc_tensor((1,), scope="local") + A_warp_1 = T.decl_tensor((32,), data=A_warp.data, scope="local") + A_1 = T.decl_tensor((32,), data=A) # A is a handle param A_warp_1[0] = A_1[threadIdx_x] - B_warp_1 = T.decl_buffer((32,), data=B_warp.data, scope="local") + B_warp_1 = T.decl_tensor((32,), data=B_warp.data, scope="local") T.tvm_storage_sync("warp") B_warp_1[0] = T.tvm_warp_shuffle( T.tvm_warp_activemask(), A_warp_1[0], threadIdx_x % 4 * 8 + threadIdx_x // 4, 32, 32 ) + T.float32(1) - red_buf0_1 = T.decl_buffer((1,), data=red_buf0.data, scope="local") + red_buf0_1 = T.decl_tensor((1,), data=red_buf0.data, scope="local") with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0), ): - mask = T.alloc_buffer((1,), "uint32", scope="local") - t0 = T.alloc_buffer((1,), scope="local") + mask = T.alloc_tensor((1,), "uint32", scope="local") + t0 = T.alloc_tensor((1,), scope="local") red_buf0_1[0] = A_warp_1[0] - mask_1 = T.decl_buffer((1,), "uint32", data=mask.data, scope="local") + mask_1 = T.decl_tensor((1,), "uint32", data=mask.data, scope="local") mask_1[0] = T.tvm_warp_activemask() - t0_1 = T.decl_buffer((1,), data=t0.data, scope="local") + t0_1 = T.decl_tensor((1,), data=t0.data, scope="local") t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 16, 32, 32) red_buf0_1[0] = red_buf0_1[0] + t0_1[0] t0_1[0] = T.tvm_warp_shuffle_down(mask_1[0], red_buf0_1[0], 8, 32, 32) @@ -1909,9 +1909,9 @@ def func( # NOTE(Zihao): test tvm_warp_shuffle_up red_buf0_1[0] = T.tvm_warp_shuffle_up(mask_1[0], red_buf0_1[0], 0, 32, 32) if threadIdx_x == 0: - C_1 = T.decl_buffer((1,), data=C) + C_1 = T.decl_tensor((1,), data=C) C_1[0] = red_buf0_1[0] - B_1 = T.decl_buffer((32,), data=B) + B_1 = T.decl_tensor((32,), data=B) B_1[threadIdx_x] = B_warp_1[0] return func @@ -1919,7 +1919,7 @@ def func( def make_packed_api_result(): @Ts.prim_func - def func(A: T.Buffer(64, "float32")): + def func(A: T.Tensor(64, "float32")): T.func_attr({"global_symbol": "main", "target": T.target("cuda")}) bx = T.launch_thread("blockIdx.x", 64) T.evaluate(A[bx]) @@ -1939,7 +1939,7 @@ def tvm_struct_set_generated_in_cpp(): @I.ir_module class Module: @Ts.prim_func - def tir_packed_call(A: T.Buffer(16)): + def tir_packed_call(A: T.Tensor(16)): T.attr(0, "device_id", 0) T.attr(0, "device_type", 0) T.evaluate( @@ -1960,10 +1960,10 @@ def tir_packed_call(A: T.Buffer(16)): def undefined_data_ptr_in_decl_buffer(): - """The T.decl_buffer syntax should not introduce an Allocate + """The T.decl_tensor syntax should not introduce an Allocate - While T.decl_buffer can be used to represent an - Allocate/DeclBuffer pair, performing a round-trip through + While T.decl_tensor can be used to represent an + Allocate/DeclTensor pair, performing a round-trip through TVMScript should not introduce an Allocate node. """ @@ -1971,7 +1971,7 @@ def undefined_data_ptr_in_decl_buffer(): @Ts.prim_func(check_well_formed=False) def func(): data_ptr = T.handle("float32") - buf = T.decl_buffer(shape=[1], dtype="float32", data=data_ptr) + buf = T.decl_tensor(shape=[1], dtype="float32", data=data_ptr) T.evaluate(buf[0]) return func @@ -2027,7 +2027,7 @@ def func(): def test_address_of_buffer(): @Ts.prim_func - def func(A: T.Buffer((128, 128), "float32")): + def func(A: T.Tensor((128, 128), "float32")): T.evaluate(T.address_of(A)) assert "T.address_of(A[0, 0])" in func.script() @@ -2106,7 +2106,7 @@ def _func(): y: T.int32 = x + 1 T.evaluate(y - 1) - # Explicit scalar declarations lower to AllocBuffer + BufferStore (local_scalar). + # Explicit scalar declarations lower to AllocTensor + BufferStore (local_scalar). # The printer fuses each pair into one line; annotate the allocation for y. result = _func.with_attr("global_symbol", "main").script( obj_to_annotate={ @@ -2255,7 +2255,7 @@ def test_roundtrip_expressions(ir_generator): scalable_vectors, predicated_buffer_load_store, void_ptr, - decl_buffer, + decl_tensor, allocate_and_decl_buffer, alloc_buffer_example, undefined_data_ptr_in_decl_buffer, @@ -2305,7 +2305,7 @@ def test_roundtrip_metadata(ir_generator): # Import-time construction also checks the annotated S-TIR API. @Ts.prim_func def lowered_loop_split( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: reduce_temp0 = Ts.sblock_alloc_buffer([1], dtype="float32", strides=[1], scope="local") normal_reduce_temp0 = Ts.sblock_alloc_buffer([1], dtype="float32", strides=[1], scope="local") diff --git a/tests/python/s_tir/script/test_s_tir_script_source_locations.py b/tests/python/s_tir/script/test_s_tir_script_source_locations.py index f09095cca32a..7d9e5fbf5b30 100644 --- a/tests/python/s_tir/script/test_s_tir_script_source_locations.py +++ b/tests/python/s_tir/script/test_s_tir_script_source_locations.py @@ -32,7 +32,7 @@ class dummy: class Module: @Ts.prim_func def impl( - A: T.Buffer((12, 196, 64), "float32"), + A: T.Tensor((12, 196, 64), "float32"), ) -> None: T.evaluate(0) diff --git a/tests/python/s_tir/test_arith_domain_touched.py b/tests/python/s_tir/test_arith_domain_touched.py index 58f7d3db36bb..5c9561a077bc 100644 --- a/tests/python/s_tir/test_arith_domain_touched.py +++ b/tests/python/s_tir/test_arith_domain_touched.py @@ -27,7 +27,7 @@ @Ts.prim_func -def scalar_func(A: T.Buffer((100, m)), B: T.Buffer((100, m))): +def scalar_func(A: T.Tensor((100, m)), B: T.Tensor((100, m))): for i, j in T.grid(100, m): A[i, j] = B[i - 1, j + 1] + A[i - 1, j - 1] @@ -75,7 +75,7 @@ def test_domain_touched_vector(): m = tvm.runtime.convert(128) @Ts.prim_func - def func(A: T.Buffer((n * m,)), B: T.Buffer((n * m,)), n: T.int32): # noqa: F821 + def func(A: T.Tensor((n * m,)), B: T.Tensor((n * m,)), n: T.int32): # noqa: F821 for i in T.serial(n): A[i * m : (i + 1) * m : 1] = A[i * m : (i + 1) * m : 1] + B[i * m : (i + 1) * m : 1] 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 fc425824bc3f..d6c8fe329e93 100644 --- a/tests/python/s_tir/test_s_tir_renew_defs.py +++ b/tests/python/s_tir/test_s_tir_renew_defs.py @@ -53,7 +53,7 @@ def _check_block_signature_remap(lhs: SBlock, rhs: SBlock): def test_simple(): @Ts.prim_func # Buffer A should be remapped - def elementwise(A: T.Buffer((128, 128), "float32")): + def elementwise(A: T.Tensor((128, 128), "float32")): # Buffer B should be remapped B = Ts.sblock_alloc_buffer((128, 128), "float32") # i, j should be remapped @@ -92,7 +92,7 @@ def test_match_buffer(): @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")): + def func_match_buffer(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32")): with Ts.sblock("root"): # A0 should be remapped A0 = Ts.match_buffer( @@ -139,7 +139,7 @@ def test_undefined_buffer(): @Ts.prim_func def access_alloc(): # Buffer A should be remapped - A = T.alloc_buffer((128,), "float16") + A = T.alloc_tensor((128,), "float16") T.evaluate(A.data) for i in range(128): A[i] = A[i] + T.float16(1.0) @@ -148,11 +148,11 @@ def access_alloc(): f2 = tvm.s_tir.renew_defs(f1) tvm.ir.assert_structural_equal(f1, f2) - # AllocBuffer is now a flat statement in SeqStmt + # AllocTensor is now a flat statement in SeqStmt assert f1.body.seq[0].var.data != f2.body.seq[0].var.data def _get_buffer_store_buffer(f): - # SeqStmt: [AllocBuffer, Evaluate, For]; For body has the BufferStore + # SeqStmt: [AllocTensor, Evaluate, For]; For body has the BufferStore return f.body.seq[2].body.buffer _check_buffer_decl(_get_buffer_store_buffer(f1), _get_buffer_store_buffer(f2)) @@ -162,7 +162,7 @@ def test_symbolic_func(): m = T.dynamic("m", "int32") @Ts.prim_func - def symbolic_func(A: T.Buffer((n, m)), B: T.Buffer((n, m * 2)), n: T.int32): # noqa: F821 + def symbolic_func(A: T.Tensor((n, m)), B: T.Tensor((n, m * 2)), n: T.int32): # noqa: F821 for i, j in T.grid(n, m): B[i, j * 2] = A[i, j] B[i, j * 2 + 1] = A[i, j] @@ -176,7 +176,7 @@ def test_buffer_params(): m = T.dynamic("m") @Ts.prim_func - def main(A: T.Buffer((m * 2,)), B: T.Buffer((m, 2))): + def main(A: T.Tensor((m * 2,)), B: T.Tensor((m, 2))): for i, j in T.grid(m, 2): with Ts.sblock("B"): vi, vj = Ts.axis.remap("SS", [i, j]) @@ -190,7 +190,7 @@ def main(A: T.Buffer((m * 2,)), B: T.Buffer((m, 2))): def test_compound_buffer_param_shape_var(): n = tvm.tirx.Var("n", "int32") - A = tvm.tirx.decl_buffer((tvm.tirx.max(n, 1),), layout=None) + A = tvm.tirx.decl_tensor((tvm.tirx.max(n, 1),), layout=None) f1 = tvm.tirx.PrimFunc([A], tvm.tirx.Evaluate(n)) f2 = tvm.s_tir.renew_defs(f1) @@ -202,9 +202,9 @@ def test_compound_buffer_param_shape_var(): def test_gather(): @Ts.prim_func(private=True) def take( - A: T.Buffer((4096, 4096), "float16"), - B: T.Buffer((1,), "int32"), - T_take: T.Buffer((1, 4096), "float16"), + A: T.Tensor((4096, 4096), "float16"), + B: T.Tensor((1,), "int32"), + T_take: T.Tensor((1, 4096), "float16"), ): for ax0, ax1 in T.grid(1, 4096): with Ts.sblock("T_take"): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_annotate_irregular_loop.py b/tests/python/s_tir/transform/test_s_tir_transform_annotate_irregular_loop.py index 0a8b6e1220d1..ee81cd8d9351 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_annotate_irregular_loop.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_annotate_irregular_loop.py @@ -30,7 +30,7 @@ def test_handle_irrgular_unit_loop(): """Dedicated testcase to check the unitloop with loop jump not simplified""" @Ts.prim_func - def before(A: T.Buffer((10,), "int32")): + def before(A: T.Tensor((10,), "int32")): for i in T.serial(1): if A[i] > 5: break @@ -43,7 +43,7 @@ def before(A: T.Buffer((10,), "int32")): A[k] = A[k] + 1 @Ts.prim_func - def expected(A: T.Buffer((10,), "int32")): + def expected(A: T.Tensor((10,), "int32")): for i in T.serial(1, annotations={"irregular_loop_mark": 1}): if A[i] > 5: break @@ -67,7 +67,7 @@ def test_annotate_loop_with_break(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((10,), "int32")): + def main(A: T.Tensor((10,), "int32")): for i in T.serial(10): if A[i] > 5: break @@ -76,7 +76,7 @@ def main(A: T.Buffer((10,), "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((10,), "int32")): + def main(A: T.Tensor((10,), "int32")): for i in T.serial(10, annotations={"irregular_loop_mark": 1}): if A[i] > 5: break @@ -93,7 +93,7 @@ def test_annotate_loop_with_continue(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((10,), "int32")): + def main(A: T.Tensor((10,), "int32")): for i in T.serial(10): if A[i] < 0: continue @@ -102,7 +102,7 @@ def main(A: T.Buffer((10,), "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((10,), "int32")): + def main(A: T.Tensor((10,), "int32")): for i in T.serial(10, annotations={"irregular_loop_mark": 1}): if A[i] < 0: continue @@ -119,7 +119,7 @@ def test_nested_irregular_both_loops(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((10, 10), "int32")): + def main(A: T.Tensor((10, 10), "int32")): for i in T.serial(10): if i > 7: break @@ -131,7 +131,7 @@ def main(A: T.Buffer((10, 10), "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((10, 10), "int32")): + def main(A: T.Tensor((10, 10), "int32")): for i in T.serial(10, annotations={"irregular_loop_mark": 1}): if i > 7: break @@ -151,7 +151,7 @@ def test_while_loop_with_break(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((10,), "int32")): + def main(A: T.Tensor((10,), "int32")): i = T.int32(0) while i < 10: if A[i] > 5: @@ -162,7 +162,7 @@ def main(A: T.Buffer((10,), "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((10,), "int32")): + def main(A: T.Tensor((10,), "int32")): i = T.int32(0) while i < 10: if A[i] > 5: @@ -181,7 +181,7 @@ def test_break_in_nested_conditional(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((10,), "int32"), flag1: T.int32, flag2: T.int32): + def main(A: T.Tensor((10,), "int32"), flag1: T.int32, flag2: T.int32): for i in T.serial(10): if flag1 > 0: if flag2 > 0: @@ -192,7 +192,7 @@ def main(A: T.Buffer((10,), "int32"), flag1: T.int32, flag2: T.int32): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((10,), "int32"), flag1: T.int32, flag2: T.int32): + def main(A: T.Tensor((10,), "int32"), flag1: T.int32, flag2: T.int32): for i in T.serial(10, annotations={"irregular_loop_mark": 1}): if flag1 > 0: if flag2 > 0: @@ -211,7 +211,7 @@ def test_while_loop_with_break_standalone(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((10,), "int32")): + def main(A: T.Tensor((10,), "int32")): i = T.int32(0) while i < 10: if A[i] > 5: @@ -222,7 +222,7 @@ def main(A: T.Buffer((10,), "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((10,), "int32")): + def main(A: T.Tensor((10,), "int32")): i = T.int32(0) while i < 10: if A[i] > 5: @@ -241,7 +241,7 @@ def test_nested_irregular_loop_standalone(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((5, 5, 5), "int32")): + def main(A: T.Tensor((5, 5, 5), "int32")): for i in T.serial(5): for j in T.serial(5): for k in T.serial(5): @@ -254,7 +254,7 @@ def main(A: T.Buffer((5, 5, 5), "int32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((5, 5, 5), "int32")): + def main(A: T.Tensor((5, 5, 5), "int32")): for i in T.serial(5): for j in T.serial(5): for k in T.serial(5, annotations={"irregular_loop_mark": 1}): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_canonicalize_loop.py b/tests/python/s_tir/transform/test_s_tir_transform_canonicalize_loop.py index c4ca9cec37e3..5de6903fcd8b 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_canonicalize_loop.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_canonicalize_loop.py @@ -25,13 +25,13 @@ def test_canonicalize_loop(): @Ts.prim_func - def before(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def before(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): T.func_attr({"global_symbol": "main"}) for i in range(1, 128, 5): B[i] = A[i] + 1.0 @Ts.prim_func - def expected(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def expected(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): T.func_attr({"global_symbol": "main"}) for i in T.serial(0, 26): B[i * 5 + 1] = A[i * 5 + 1] + 1.0 @@ -43,14 +43,14 @@ def expected(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): def test_canonicalize_nested_loop(): @Ts.prim_func - def before(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32")): + def before(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32")): T.func_attr({"global_symbol": "main"}) for i in range(1, 128, 5): for j in range(2, 128, 3): B[i, j] = A[i, j] + 1.0 @Ts.prim_func - def expected(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float32")): + def expected(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128, 128), "float32")): T.func_attr({"global_symbol": "main"}) for i in T.serial(0, 26): for j in T.serial(0, 42): @@ -63,7 +63,7 @@ def expected(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128, 128), "float3 def test_canonicalize_negative_step(): @Ts.prim_func - def before(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def before(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): T.func_attr({"global_symbol": "main"}) for i in T.serial(0, 127, step=-3): B[i] = A[i] + 1.0 @@ -77,7 +77,7 @@ def test_canonicalize_dynamic_step(): """Currently we report error for dynamic step since we could not prove it is positive""" @Ts.prim_func - def before(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32"), step: T.int32): + def before(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32"), step: T.int32): T.func_attr({"global_symbol": "main"}) for i in T.serial(0, 128, step=step): B[i] = A[i] + 1.0 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 2fd48723f660..a28c75387cc5 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 @@ -80,7 +80,7 @@ def test_compact(self): class TestElemwise(BaseCompactTest): @Ts.prim_func - def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: + def before(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i in range(0, 16): with Ts.sblock(): Ts.reads(A[i, 0:16]) @@ -98,7 +98,7 @@ def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) - C[i, j] = B[i, j] * 2.0 @Ts.prim_func - def expected(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: + def expected(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i in range(0, 16): with Ts.sblock(): Ts.reads(A[i, 0:16]) @@ -118,7 +118,7 @@ def expected(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) class TestUnschedulableFunc(BaseCompactTest): @Ts.prim_func - def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: + def before(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i in range(0, 16): with Ts.sblock(): Ts.reads(A[i, 0:16]) @@ -135,7 +135,7 @@ def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) - class TestParamBufferAccess(BaseCompactTest): @Ts.prim_func - def before(A: T.Buffer((20, 20), "float32"), B: T.Buffer((20, 20), "float32")) -> None: + def before(A: T.Tensor((20, 20), "float32"), B: T.Tensor((20, 20), "float32")) -> None: for i in range(0, 16): with Ts.sblock(): Ts.reads(A[i, 0:16]) @@ -151,7 +151,7 @@ def before(A: T.Buffer((20, 20), "float32"), B: T.Buffer((20, 20), "float32")) - class TestSharedMem(BaseCompactTest): @Ts.prim_func - def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: + def before(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i0 in T.thread_binding(0, 2, thread="blockIdx.x"): for i1 in T.thread_binding(0, 2, thread="vthread"): for i2 in T.thread_binding(0, 4, thread="threadIdx.x"): @@ -171,7 +171,7 @@ def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) - C[i0 * 8 + i1 * 4 + i2, j] = B[i0 * 8 + i1 * 4 + i2, j] * 2.0 @Ts.prim_func - def expected(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: + def expected(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i0 in T.thread_binding(0, 2, thread="blockIdx.x"): for i1 in T.thread_binding(0, 2, thread="vthread"): for i2 in T.thread_binding(0, 4, thread="threadIdx.x"): @@ -193,7 +193,7 @@ def expected(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) class TestWrapMem(BaseCompactTest): @Ts.prim_func - def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: + def before(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i0 in T.thread_binding(0, 2, thread="blockIdx.x"): for i1 in T.thread_binding(0, 2, thread="vthread"): for i2 in T.thread_binding(0, 4, thread="threadIdx.x"): @@ -213,7 +213,7 @@ def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) - C[i0 * 8 + i1 * 4 + i2, j] = B[i0 * 8 + i1 * 4 + i2, j] * 2.0 @Ts.prim_func - def expected(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: + def expected(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i0 in T.thread_binding(0, 2, thread="blockIdx.x"): for i1 in T.thread_binding(0, 2, thread="vthread"): for i2 in T.thread_binding(0, 4, thread="threadIdx.x"): @@ -236,8 +236,8 @@ def expected(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) class TestSymbolic(BaseCompactTest): @Ts.prim_func def before( - A: T.Buffer((n * 8,), "float32"), # noqa: F821 - C: T.Buffer((n * 8,), "float32"), # noqa: F821 + A: T.Tensor((n * 8,), "float32"), # noqa: F821 + C: T.Tensor((n * 8,), "float32"), # noqa: F821 n: T.int32, ) -> None: for i in range(0, n): @@ -258,8 +258,8 @@ def before( @Ts.prim_func def expected( - A: T.Buffer((n * 8,), "float32"), # noqa: F821 - C: T.Buffer((n * 8,), "float32"), # noqa: F821 + A: T.Tensor((n * 8,), "float32"), # noqa: F821 + C: T.Tensor((n * 8,), "float32"), # noqa: F821 n: T.int32, ) -> None: for i in range(0, n): @@ -281,7 +281,7 @@ def expected( class TestComplexFunc(BaseCompactTest): @Ts.prim_func - def before(A: T.Buffer((8, 8), "float32"), C: T.Buffer((8, 8), "float32"), n: T.int32) -> None: + def before(A: T.Tensor((8, 8), "float32"), C: T.Tensor((8, 8), "float32"), n: T.int32) -> None: for i in range(0, 8): with Ts.sblock(): Ts.reads(A[0, 8]) @@ -309,7 +309,7 @@ def before(A: T.Buffer((8, 8), "float32"), C: T.Buffer((8, 8), "float32"), n: T. @Ts.prim_func def expected( - A: T.Buffer((8, 8), "float32"), C: T.Buffer((8, 8), "float32"), n: T.int32 + A: T.Tensor((8, 8), "float32"), C: T.Tensor((8, 8), "float32"), n: T.int32 ) -> None: for i in range(0, 8): with Ts.sblock(): @@ -341,7 +341,7 @@ class TestMatchBuffer(BaseCompactTest): is_lower_order_free = False @Ts.prim_func - def before(A: T.Buffer((16, 16)), C: T.Buffer((16, 16))) -> None: + def before(A: T.Tensor((16, 16)), C: T.Tensor((16, 16))) -> None: for i in range(0, 16): with Ts.sblock(): A0 = Ts.match_buffer(A[i, 0:16], (16)) @@ -361,7 +361,7 @@ def before(A: T.Buffer((16, 16)), C: T.Buffer((16, 16))) -> None: C1[()] = B2[()] * 2.0 @Ts.prim_func - def expected(A: T.Buffer((16, 16)), C: T.Buffer((16, 16))) -> None: + def expected(A: T.Tensor((16, 16)), C: T.Tensor((16, 16))) -> None: for i in range(0, 16): with Ts.sblock(): A0 = Ts.match_buffer(A[i, 0:16], (16)) @@ -383,7 +383,7 @@ def expected(A: T.Buffer((16, 16)), C: T.Buffer((16, 16))) -> None: class TestStorageAlign(BaseCompactTest): @Ts.prim_func - def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: + def before(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i in range(0, 16): with Ts.sblock(): Ts.reads(A[i, 0:16]) @@ -402,7 +402,7 @@ def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) - C[i, j] = B[i, j] * 2.0 @Ts.prim_func - def expected(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: + def expected(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i in range(0, 16): with Ts.sblock(): Ts.reads(A[i, 0:16]) @@ -423,7 +423,7 @@ def expected(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) class TestPaddingPattern(BaseCompactTest): @Ts.prim_func - def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((20, 20), "float32")) -> None: + def before(A: T.Tensor((16, 16), "float32"), C: T.Tensor((20, 20), "float32")) -> None: with Ts.sblock(): B = Ts.sblock_alloc_buffer((20, 20), dtype="float32") for i, j in T.grid(16, 16): @@ -439,7 +439,7 @@ def before(A: T.Buffer((16, 16), "float32"), C: T.Buffer((20, 20), "float32")) - @Ts.prim_func def expected( - A: T.Buffer([16, 16], dtype="float32"), C: T.Buffer([20, 20], dtype="float32") + A: T.Tensor([16, 16], dtype="float32"), C: T.Tensor([20, 20], dtype="float32") ) -> None: with Ts.sblock(): B = Ts.sblock_alloc_buffer([16, 16], dtype="float32") @@ -458,7 +458,7 @@ def expected( class TestPaddingPatternInlined(BaseCompactTest): @Ts.prim_func def before( - X: T.Buffer([224, 224], dtype="float32"), Y: T.Buffer([224, 224], dtype="float32") + X: T.Tensor([224, 224], dtype="float32"), Y: T.Tensor([224, 224], dtype="float32") ) -> None: cache = Ts.sblock_alloc_buffer([224, 224], dtype="float32") for h, w in T.grid(224, 224): @@ -479,7 +479,7 @@ def before( ) @Ts.prim_func - def expected(X: T.Buffer((224, 224), "float32"), Y: T.Buffer((224, 224), "float32")) -> None: + def expected(X: T.Tensor((224, 224), "float32"), Y: T.Tensor((224, 224), "float32")) -> None: cache = Ts.sblock_alloc_buffer([224, 224], dtype="float32") for h, w in T.grid(224, 224): with Ts.sblock("cache"): @@ -501,7 +501,7 @@ def expected(X: T.Buffer((224, 224), "float32"), Y: T.Buffer((224, 224), "float3 class TestMemAccessInBranch(BaseCompactTest): @Ts.prim_func - def before(A: T.Buffer((224, 224), "float32")) -> None: + def before(A: T.Tensor((224, 224), "float32")) -> None: with Ts.sblock(): B1 = Ts.sblock_alloc_buffer((224, 224), dtype="float32") B2 = Ts.sblock_alloc_buffer((224, 224), dtype="float32") @@ -523,7 +523,7 @@ def before(A: T.Buffer((224, 224), "float32")) -> None: B4[i, j] = A[i, j] + 3.0 @Ts.prim_func - def expected(A: T.Buffer([224, 224], dtype="float32")) -> None: + def expected(A: T.Tensor([224, 224], dtype="float32")) -> None: with Ts.sblock(): B1 = Ts.sblock_alloc_buffer([112, 112], dtype="float32") B2 = Ts.sblock_alloc_buffer([224, 224], dtype="float32") @@ -547,7 +547,7 @@ class TestAnnotatedOpaqueAccess(BaseCompactTest): is_lower_order_free = False @Ts.prim_func - def before(A: T.Buffer((1024,), "float32")) -> None: + def before(A: T.Tensor((1024,), "float32")) -> None: with Ts.sblock(): B = Ts.sblock_alloc_buffer((1024,), dtype="float32") C = Ts.sblock_alloc_buffer((1024,), dtype="float32") @@ -571,7 +571,7 @@ def before(A: T.Buffer((1024,), "float32")) -> None: C[i] = B[i] @Ts.prim_func - def expected(A: T.Buffer((1024,), "float32")) -> None: + def expected(A: T.Tensor((1024,), "float32")) -> None: with Ts.sblock(): B = Ts.sblock_alloc_buffer((1024,), dtype="float32") C = Ts.sblock_alloc_buffer((520,), dtype="float32") @@ -598,9 +598,9 @@ def expected(A: T.Buffer((1024,), "float32")) -> None: class TestSparseReadCache(BaseCompactTest): @Ts.prim_func def before( - A_data: T.Buffer((819,), "float32"), - B: T.Buffer((128,), "float32"), - A_indptr: T.Buffer((129,), "int32"), + A_data: T.Tensor((819,), "float32"), + B: T.Tensor((128,), "float32"), + A_indptr: T.Tensor((129,), "int32"), ) -> None: for i in T.serial(128): with Ts.sblock("rowsum_outer"): @@ -629,9 +629,9 @@ def before( @Ts.prim_func def expected( - A_data: T.Buffer((819,), "float32"), - B: T.Buffer((128,), "float32"), - A_indptr: T.Buffer((129,), "int32"), + A_data: T.Tensor((819,), "float32"), + B: T.Tensor((128,), "float32"), + A_indptr: T.Tensor((129,), "int32"), ) -> None: for i in T.serial(128): with Ts.sblock("rowsum_outer"): @@ -665,13 +665,13 @@ class TestDataDependentRegion(BaseCompactTest): @Ts.prim_func def before( - p0: T.Buffer((30,), "float32"), - p1: T.Buffer((1,), "int32"), - hybrid_nms: T.Buffer((30,), "float32"), + p0: T.Tensor((30,), "float32"), + p1: T.Tensor((1,), "int32"), + hybrid_nms: T.Tensor((30,), "float32"), ): - argsort_nms_cpu = T.decl_buffer([5], "int32", scope="global") + argsort_nms_cpu = T.decl_tensor([5], "int32", scope="global") for i in range(1): - nkeep = T.decl_buffer([1], "int32", scope="global") + nkeep = T.decl_tensor([1], "int32", scope="global") if 0 < p1[i]: nkeep[0] = p1[i] if 2 < nkeep[0]: @@ -693,7 +693,7 @@ def before( class TestNarrowShape(BaseCompactTest): @Ts.prim_func - def before(A: T.Buffer((10,), "float32"), B: T.Buffer((10,), "float32")) -> None: + def before(A: T.Tensor((10,), "float32"), B: T.Tensor((10,), "float32")) -> None: B_cache = Ts.sblock_alloc_buffer(10, "float32") for j in T.serial(3): for k in T.serial(4): @@ -704,7 +704,7 @@ def before(A: T.Buffer((10,), "float32"), B: T.Buffer((10,), "float32")) -> None A[i] = B_cache[i] + T.float32(1) @Ts.prim_func - def expected(A: T.Buffer((10,), "float32"), B: T.Buffer((10,), "float32")) -> None: + def expected(A: T.Tensor((10,), "float32"), B: T.Tensor((10,), "float32")) -> None: B_cache = Ts.sblock_alloc_buffer([10], dtype="float32") for j, k in T.grid(3, 4): with Ts.sblock("B_cache"): @@ -752,7 +752,7 @@ def before(): class TestSpatialTiledPadPooling(BaseCompactTest): @Ts.prim_func - def before(X: T.Buffer((64, 112, 112), "int32"), Y: T.Buffer((64, 56, 56), "int32")) -> None: + def before(X: T.Tensor((64, 112, 112), "int32"), Y: T.Tensor((64, 56, 56), "int32")) -> None: for h_o, w_o in T.grid(14, 14): with Ts.sblock(): X_cache = Ts.sblock_alloc_buffer([112, 112, 64], dtype="int32") @@ -789,7 +789,7 @@ def before(X: T.Buffer((64, 112, 112), "int32"), Y: T.Buffer((64, 56, 56), "int3 ) @Ts.prim_func - def expected(X: T.Buffer((64, 112, 112), "int32"), Y: T.Buffer((64, 56, 56), "int32")) -> None: + def expected(X: T.Tensor((64, 112, 112), "int32"), Y: T.Tensor((64, 56, 56), "int32")) -> None: for h_o, w_o in T.grid(14, 14): with Ts.sblock(): Ts.reads(X[0:64, h_o * 8 - 1 : h_o * 8 + 8, w_o * 8 - 1 : w_o * 8 + 8]) @@ -843,7 +843,7 @@ class TestComplexCase1(BaseCompactTest): # fmt: off @Ts.prim_func - def before(A: T.Buffer((960, 770), "float32"), B: T.Buffer((770, 2304), "float32"), C: T.Buffer((960, 2304), "float32")) -> None: + def before(A: T.Tensor((960, 770), "float32"), B: T.Tensor((770, 2304), "float32"), C: T.Tensor((960, 2304), "float32")) -> None: for bx in T.thread_binding(144, thread="blockIdx.x"): for vx in T.thread_binding(2, thread="vthread.x"): for tx_p in T.thread_binding(256, thread="threadIdx.x"): @@ -869,7 +869,7 @@ def before(A: T.Buffer((960, 770), "float32"), B: T.Buffer((770, 2304), "float32 C[(((bx // 18 + 0) * 8 + tx_p // 32) * 8 + i_3) * 2 + i_4, ((bx % 18 * 2 + vx % 2) * 32 + tx_p % 32 + j_3) * 2 + j_4] = C[(((bx // 18 + 0) * 8 + tx_p // 32) * 8 + i_3) * 2 + i_4, ((bx % 18 * 2 + vx % 2) * 32 + tx_p % 32 + j_3) * 2 + j_4] + A_shared[(((bx // 18 + 0) * 8 + tx_p // 32) * 8 + i_3) * 2 + i_4, (k_0 + k_1) * 4 + k_2] * B_shared[(k_0 + k_1) * 4 + k_2, ((bx % 18 * 2 + vx % 2) * 32 + tx_p % 32 + j_3) * 2 + j_4] @Ts.prim_func - def expected(A: T.Buffer((960, 770), "float32"), B: T.Buffer((770, 2304), "float32"), C: T.Buffer((960, 2304), "float32")) -> None: + def expected(A: T.Tensor((960, 770), "float32"), B: T.Tensor((770, 2304), "float32"), C: T.Tensor((960, 2304), "float32")) -> None: for bx in T.thread_binding(144, thread="blockIdx.x"): for vx in T.thread_binding(2, thread="vthread.x"): for tx_p in T.thread_binding(256, thread="threadIdx.x"): @@ -926,9 +926,9 @@ class TestDependentBufferIndicesOfPackedMatmul(BaseCompactTest): @Ts.prim_func def before( - A: T.Buffer((1020, 64), "float32"), - B: T.Buffer((1000, 64), "float32"), - C: T.Buffer((1020, 1000), "float32"), + A: T.Tensor((1020, 64), "float32"), + B: T.Tensor((1000, 64), "float32"), + C: T.Tensor((1020, 1000), "float32"), ): for i0, i1 in T.grid(4, 1): with Ts.sblock(): @@ -965,9 +965,9 @@ def before( @Ts.prim_func def expected( - A: T.Buffer((1020, 64), "float32"), - B: T.Buffer((1000, 64), "float32"), - C: T.Buffer((1020, 1000), "float32"), + A: T.Tensor((1020, 64), "float32"), + B: T.Tensor((1000, 64), "float32"), + C: T.Tensor((1020, 1000), "float32"), ) -> None: for i0, i1 in T.grid(4, 1): with Ts.sblock(): @@ -1007,15 +1007,15 @@ class TestTileAwareCompaction(BaseCompactTest): def before(self): @Ts.prim_func def main( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ): for i_0 in range(5, annotations={"pragma_loop_partition_hint": 1}): for j_0 in range(5, annotations={"pragma_loop_partition_hint": 1}): - A_local = T.decl_buffer((26, 128), scope="local") - B_local = T.decl_buffer((128, 26), scope="local") - C_local = T.decl_buffer((26, 26), scope="local") + A_local = T.decl_tensor((26, 128), scope="local") + B_local = T.decl_tensor((128, 26), scope="local") + C_local = T.decl_tensor((26, 26), scope="local") for ax0, ax1 in T.grid(26, 128): if i_0 * 26 + ax0 < 128: A_local[ax0, ax1] = A[i_0 * 26 + ax0, ax1] @@ -1045,15 +1045,15 @@ def main( @Ts.prim_func def expected( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ): for i_0 in range(4): for j_0 in range(4): - A_local_tile0 = T.decl_buffer((26, 128), scope="local") - B_local_tile0 = T.decl_buffer((128, 26), scope="local") - C_local_tile0 = T.decl_buffer((26, 26), scope="local") + A_local_tile0 = T.decl_tensor((26, 128), scope="local") + B_local_tile0 = T.decl_tensor((128, 26), scope="local") + C_local_tile0 = T.decl_tensor((26, 26), scope="local") for ax0, ax1 in T.grid(26, 128): A_local_tile0[ax0, ax1] = A[i_0 * 26 + ax0, ax1] for ax0, ax1 in T.grid(128, 26): @@ -1067,9 +1067,9 @@ def expected( for ax0, ax1 in T.grid(26, 26): C[i_0 * 26 + ax0, j_0 * 26 + ax1] = C_local_tile0[ax0, ax1] - A_local_tile1 = T.decl_buffer((26, 128), scope="local") - B_local_tile1 = T.decl_buffer((128, 24), scope="local") - C_local_tile1 = T.decl_buffer((26, 24), scope="local") + A_local_tile1 = T.decl_tensor((26, 128), scope="local") + B_local_tile1 = T.decl_tensor((128, 24), scope="local") + C_local_tile1 = T.decl_tensor((26, 24), scope="local") for ax0, ax1 in T.grid(26, 128): A_local_tile1[ax0, ax1] = A[i_0 * 26 + ax0, ax1] for ax0, ax1 in T.grid(128, 26): @@ -1087,9 +1087,9 @@ def expected( C[i_0 * 26 + ax0, ax1 + 104] = C_local_tile1[ax0, ax1] for j_0 in range(4): - A_local_tile2 = T.decl_buffer((24, 128), scope="local") - B_local_tile2 = T.decl_buffer((128, 26), scope="local") - C_local_tile2 = T.decl_buffer((24, 26), scope="local") + A_local_tile2 = T.decl_tensor((24, 128), scope="local") + B_local_tile2 = T.decl_tensor((128, 26), scope="local") + C_local_tile2 = T.decl_tensor((24, 26), scope="local") for ax0, ax1 in T.grid(26, 128): if ax0 < 24: A_local_tile2[ax0, ax1] = A[ax0 + 104, ax1] @@ -1106,9 +1106,9 @@ def expected( if ax0 < 24: C[ax0 + 104, j_0 * 26 + ax1] = C_local_tile2[ax0, ax1] - A_local_tile3 = T.decl_buffer((24, 128), scope="local") - B_local_tile3 = T.decl_buffer((128, 24), scope="local") - C_local_tile3 = T.decl_buffer((24, 24), scope="local") + A_local_tile3 = T.decl_tensor((24, 128), scope="local") + B_local_tile3 = T.decl_tensor((128, 24), scope="local") + C_local_tile3 = T.decl_tensor((24, 24), scope="local") for ax0, ax1 in T.grid(26, 128): if ax0 < 24: A_local_tile3[ax0, ax1] = A[ax0 + 104, ax1] @@ -1132,9 +1132,9 @@ class TestNonStrictCompactionForPaddedMatmul(BaseCompactTest): @Ts.prim_func def before( - A: T.Buffer((127, 127), "float32"), - B: T.Buffer((127, 127), "float32"), - C: T.Buffer((127, 127), "float32"), + A: T.Tensor((127, 127), "float32"), + B: T.Tensor((127, 127), "float32"), + C: T.Tensor((127, 127), "float32"), ): """A mock workload where the intermediate buffer allocation is not enought originally""" for i_0, j_0 in T.grid(4, 4): @@ -1170,9 +1170,9 @@ def before( @Ts.prim_func def expected( - A: T.Buffer((127, 127), "float32"), - B: T.Buffer((127, 127), "float32"), - C: T.Buffer((127, 127), "float32"), + A: T.Tensor((127, 127), "float32"), + B: T.Tensor((127, 127), "float32"), + C: T.Tensor((127, 127), "float32"), ): for i_0, j_0 in T.grid(4, 4): with Ts.sblock(""): @@ -1211,9 +1211,9 @@ class TestNotCompactAliasBuffer(BaseCompactTest): def before(): """Partially accessed buffer, but should not compact because existence of aliasing buffer B.""" - data = T.alloc_buffer((1024,), "int8") - A = T.decl_buffer([1024], "int8", data=data.data) - B = T.decl_buffer([512], "float16", data=data.data) + data = T.alloc_tensor((1024,), "int8") + A = T.decl_tensor([1024], "int8", data=data.data) + B = T.decl_tensor([512], "float16", data=data.data) for i in range(10): A[i] = A[i] + T.int8(1) for i in range(10): @@ -1230,8 +1230,8 @@ class TestNotCompactBufferWithDifferentDtype(BaseCompactTest): def before(): """Partially accessed buffer, but should not compact because existence of aliasing buffer B.""" - data = T.alloc_buffer((1024,), "int8") - A = T.decl_buffer([256], "int32", data=data.data) + data = T.alloc_tensor((1024,), "int8") + A = T.decl_tensor([256], "int32", data=data.data) for i in range(10): A[i] = A[i] + 1 @@ -1244,14 +1244,14 @@ class TestNonBoolCondition(BaseCompactTest): @Ts.prim_func def before(): - A = T.decl_buffer([12], "int32") + A = T.decl_tensor([12], "int32") for i in range(10): if i: A[i] = A[i] + 1 @Ts.prim_func def expected(): - A = T.decl_buffer((9,), "int32") + A = T.decl_tensor((9,), "int32") for i in range(10): if i: A[i - 1] = A[i - 1] + 1 @@ -1261,8 +1261,8 @@ def test_loop_var_does_not_escape_compacted_buffer_extent(): n = T.dynamic("n") @Ts.prim_func(private=True) - def before(A: T.Buffer((n,), "int32")): - tmp = T.alloc_buffer((n,), "int32") + def before(A: T.Tensor((n,), "int32")): + tmp = T.alloc_tensor((n,), "int32") for i in range(n): length: T.let[T.int64] = T.ceildiv(n, T.shift_left(T.int64(1), i + 1)) for j in range(length): @@ -1277,8 +1277,8 @@ class TestCompactSymbolicBound0: @Ts.prim_func def before( - X: T.Buffer((T.int64(8), n * T.int64(32))), # noqa: F821 - Y: T.Buffer((T.int64(8), n * T.int64(32))), # noqa: F821 + X: T.Tensor((T.int64(8), n * T.int64(32))), # noqa: F821 + Y: T.Tensor((T.int64(8), n * T.int64(32))), # noqa: F821 n: T.int64, ): for i, k_0 in T.grid(T.int64(8), n): @@ -1293,8 +1293,8 @@ def before( @Ts.prim_func def expected( - X: T.Buffer((T.int64(8), n * T.int64(32))), # noqa: F821 - Y: T.Buffer((T.int64(8), n * T.int64(32))), # noqa: F821 + X: T.Tensor((T.int64(8), n * T.int64(32))), # noqa: F821 + Y: T.Tensor((T.int64(8), n * T.int64(32))), # noqa: F821 n: T.int64, ): for i, k_0 in T.grid(T.int64(8), n): @@ -1313,8 +1313,8 @@ class TestCompactSymbolicBound1: @Ts.prim_func def before( - X: T.Buffer((T.int64(8), n * T.int64(32))), # noqa: F821 - Y: T.Buffer((T.int64(8), n * T.int64(32))), # noqa: F821 + X: T.Tensor((T.int64(8), n * T.int64(32))), # noqa: F821 + Y: T.Tensor((T.int64(8), n * T.int64(32))), # noqa: F821 n: T.int64, ): for i, k_0 in T.grid(T.int64(8), n): @@ -1329,8 +1329,8 @@ def before( @Ts.prim_func def expected( - X: T.Buffer((T.int64(8), n * T.int64(32))), # noqa: F821 - Y: T.Buffer((T.int64(8), n * T.int64(32))), # noqa: F821 + X: T.Tensor((T.int64(8), n * T.int64(32))), # noqa: F821 + Y: T.Tensor((T.int64(8), n * T.int64(32))), # noqa: F821 n: T.int64, ): # with Ts.sblock("root"): @@ -1349,7 +1349,7 @@ class TestSymbolicDiagMaskCase: """Test symbolic allocation not too complex""" @Ts.prim_func - def before(A: T.Buffer((1, 1, n, n)), n: T.int32): # noqa: F821 + def before(A: T.Tensor((1, 1, n, n)), n: T.int32): # noqa: F821 B = Ts.sblock_alloc_buffer((n, n)) for i in T.thread_binding(256, thread="blockIdx.x"): for j in T.thread_binding(256, thread="threadIdx.x"): @@ -1377,7 +1377,7 @@ def before(A: T.Buffer((1, 1, n, n)), n: T.int32): # noqa: F821 ] @Ts.prim_func - def expected(A: T.Buffer((1, 1, n, n)), n: T.int32): # noqa: F821 + def expected(A: T.Tensor((1, 1, n, n)), n: T.int32): # noqa: F821 B = Ts.sblock_alloc_buffer((n, n)) for i in T.thread_binding(256, thread="blockIdx.x"): for j in T.thread_binding(256, thread="threadIdx.x"): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_convert_blocks_to_opaque.py b/tests/python/s_tir/transform/test_s_tir_transform_convert_blocks_to_opaque.py index 06dd70067083..01d6d55d1457 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_convert_blocks_to_opaque.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_convert_blocks_to_opaque.py @@ -34,7 +34,7 @@ def _check(original, transformed): @Ts.prim_func -def elementwise_func(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: +def elementwise_func(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i in range(0, 16): with Ts.sblock(): Ts.reads(A[i, 0:16]) @@ -54,7 +54,7 @@ def elementwise_func(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "fl @Ts.prim_func def substituted_elementwise_func( - A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32") ) -> None: for i in range(0, 16): with Ts.sblock(): @@ -81,7 +81,7 @@ def test_error_if_predicate_uses_block_variables(): @I.ir_module(check_well_formed=False) class Before: @Ts.prim_func - def main(A: T.Buffer(8, "int32")): + def main(A: T.Tensor(8, "int32")): for i in T.serial(8): with Ts.sblock(): vi = Ts.axis.remap("S", [i]) 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 1285d1e8ef39..56ad5a0446ac 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 @@ -33,8 +33,8 @@ def test_broadcast_to_symbolic(): class Before: @Ts.prim_func def broadcast_to( - rxplaceholder: T.Buffer((T.int64(3), T.int64(1)), "float32"), - T_broadcast_to: T.Buffer((x_0, x_1)), + rxplaceholder: T.Tensor((T.int64(3), T.int64(1)), "float32"), + T_broadcast_to: T.Tensor((x_0, x_1)), ): T.func_attr({"tirx.noalias": True}) @@ -52,7 +52,7 @@ def broadcast_to( @tvm.script.ir_module class Expected: @Ts.prim_func - def broadcast_to(rxplaceholder: T.Buffer((T.int64(3), T.int64(1)), "float32"), T_broadcast_to: T.Buffer((x_0, x_1))): + def broadcast_to(rxplaceholder: T.Tensor((T.int64(3), T.int64(1)), "float32"), T_broadcast_to: T.Tensor((x_0, x_1))): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) for ax0_ax1_fused_1 in T.thread_binding(T.int64(256), thread="blockIdx.x"): @@ -78,9 +78,9 @@ def test_matmul(): class Before: @Ts.prim_func def matmul( - A: T.Buffer((32, 32), "float16"), - B: T.Buffer((32, 32), "float16"), - C: T.Buffer((32, 32), "float16"), + A: T.Tensor((32, 32), "float16"), + B: T.Tensor((32, 32), "float16"), + C: T.Tensor((32, 32), "float16"), ): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): @@ -95,9 +95,9 @@ def matmul( @Ts.prim_func def matmul_gpu( - A: T.Buffer((32, 32), "float16"), - B: T.Buffer((32, 32), "float16"), - C: T.Buffer((32, 32), "float16"), + A: T.Tensor((32, 32), "float16"), + B: T.Tensor((32, 32), "float16"), + C: T.Tensor((32, 32), "float16"), ): T.func_attr({"global_symbol": "main", "target": T.target({"arch": "sm_86", @@ -119,9 +119,9 @@ def matmul_gpu( @Ts.prim_func def matmul_cpu( - A: T.Buffer((32, 32), "float16"), - B: T.Buffer((32, 32), "float16"), - C: T.Buffer((32, 32), "float16"), + A: T.Tensor((32, 32), "float16"), + B: T.Tensor((32, 32), "float16"), + C: T.Tensor((32, 32), "float16"), ): T.func_attr({"global_symbol": "main", "target": T.target({"keys": ["cpu"], "kind": "llvm", "tag": ""}), @@ -140,9 +140,9 @@ def matmul_cpu( class Expected: @Ts.prim_func def matmul( - A: T.Buffer((32, 32), "float16"), - B: T.Buffer((32, 32), "float16"), - C: T.Buffer((32, 32), "float16"), + A: T.Tensor((32, 32), "float16"), + B: T.Tensor((32, 32), "float16"), + C: T.Tensor((32, 32), "float16"), ): T.func_attr({"tirx.is_scheduled": True, "global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): @@ -164,7 +164,7 @@ def matmul( C[v_i, v_j] = C[v_i, v_j] + A[v_i, v_k] * B[v_k, v_j] @Ts.prim_func - def matmul_cpu(A: T.Buffer((32, 32), "float16"), B: T.Buffer((32, 32), "float16"), C: T.Buffer((32, 32), "float16")): + def matmul_cpu(A: T.Tensor((32, 32), "float16"), B: T.Tensor((32, 32), "float16"), C: T.Tensor((32, 32), "float16")): T.func_attr({"global_symbol": "main", "target": T.target({"keys": ["cpu"], "kind": "llvm", "tag": ""}), "tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): for i, j, k in T.grid(32, 32, 32): @@ -177,7 +177,7 @@ def matmul_cpu(A: T.Buffer((32, 32), "float16"), B: T.Buffer((32, 32), "float16" C[v_i, v_j] = C[v_i, v_j] + A[v_i, v_k] * B[v_k, v_j] @Ts.prim_func - def matmul_gpu(A: T.Buffer((32, 32), "float16"), B: T.Buffer((32, 32), "float16"), C: T.Buffer((32, 32), "float16")): + def matmul_gpu(A: T.Tensor((32, 32), "float16"), B: T.Tensor((32, 32), "float16"), C: T.Tensor((32, 32), "float16")): T.func_attr({"global_symbol": "main", "target": T.target({"arch": "sm_86", "keys": ["cuda", "gpu"], "kind": "cuda", "max_num_threads": 1024, "tag": "", "thread_warp_size": 32}), "tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): for i_j_fused_0 in T.thread_binding(1, thread="blockIdx.x"): @@ -206,7 +206,7 @@ def test_add(): @tvm.script.ir_module class Before: @Ts.prim_func - def add(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_add: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def add(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_add: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_add"): @@ -219,11 +219,11 @@ def add(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32") class Expected: @Ts.prim_func def add( - rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), - rxplaceholder_1: T.Buffer( + rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), + rxplaceholder_1: T.Tensor( (T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32" ), - T_add: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32"), + T_add: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32"), ): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -278,7 +278,7 @@ def test_full(): @tvm.script.ir_module class Before: @Ts.prim_func - def full(rxplaceholder: T.Buffer((), "int32"), T_full: T.Buffer((T.int64(2), T.int64(3)), "int32")): + def full(rxplaceholder: T.Tensor((), "int32"), T_full: T.Tensor((T.int64(2), T.int64(3)), "int32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -291,8 +291,8 @@ def full(rxplaceholder: T.Buffer((), "int32"), T_full: T.Buffer((T.int64(2), T.i class Expected: @Ts.prim_func def full( - rxplaceholder: T.Buffer((), "int32"), - T_full: T.Buffer((T.int64(2), T.int64(3)), "int32"), + rxplaceholder: T.Tensor((), "int32"), + T_full: T.Tensor((T.int64(2), T.int64(3)), "int32"), ): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -327,8 +327,8 @@ def test_scheduled(): class Scheduled: @Ts.prim_func def full( - rxplaceholder: T.Buffer((), "int32"), - T_full: T.Buffer((T.int64(2), T.int64(3)), "int32"), + rxplaceholder: T.Tensor((), "int32"), + T_full: T.Tensor((T.int64(2), T.int64(3)), "int32"), ): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -362,7 +362,7 @@ def test_multiple(): @tvm.script.ir_module class Before: @Ts.prim_func - def add(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_add: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def add(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_add: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_add"): @@ -372,7 +372,7 @@ def add(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32") T_add[ax0, ax1, ax2, ax3] = rxplaceholder[T.int64(0), ax2, ax3] + rxplaceholder_1[ax0, ax1, ax2, T.int64(0)] @Ts.prim_func - def full(rxplaceholder: T.Buffer((), "int32"), T_full: T.Buffer((T.int64(2), T.int64(3)), "int32")): + def full(rxplaceholder: T.Tensor((), "int32"), T_full: T.Tensor((T.int64(2), T.int64(3)), "int32")): T.func_attr({"tirx.noalias": True}) for i0, i1 in T.grid(T.int64(2), T.int64(3)): with Ts.sblock("T_full"): @@ -385,11 +385,11 @@ def full(rxplaceholder: T.Buffer((), "int32"), T_full: T.Buffer((T.int64(2), T.i class Expected: @Ts.prim_func def add( - rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), - rxplaceholder_1: T.Buffer( + rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), + rxplaceholder_1: T.Tensor( (T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32" ), - T_add: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32"), + T_add: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32"), ): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -432,8 +432,8 @@ def add( @Ts.prim_func def full( - rxplaceholder: T.Buffer((), "int32"), - T_full: T.Buffer((T.int64(2), T.int64(3)), "int32"), + rxplaceholder: T.Tensor((), "int32"), + T_full: T.Tensor((T.int64(2), T.int64(3)), "int32"), ): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): @@ -465,7 +465,7 @@ def test_add_on_metal(): @tvm.script.ir_module class Before: @Ts.prim_func - def add(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_add: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def add(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_add: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.noalias": True}) for i0, i1, i2, i3 in T.grid(T.int64(4), T.int64(3), T.int64(2), T.int64(3)): with Ts.sblock("T_add"): @@ -477,7 +477,7 @@ def add(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32") @tvm.script.ir_module class Expected: @Ts.prim_func - def add(rxplaceholder: T.Buffer((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_add: T.Buffer((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): + def add(rxplaceholder: T.Tensor((T.int64(1), T.int64(2), T.int64(3)), "float32"), rxplaceholder_1: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(1)), "float32"), T_add: T.Tensor((T.int64(4), T.int64(3), T.int64(2), T.int64(3)), "float32")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) for i0_i1_i2_i3_fused_0 in T.thread_binding(T.int64(1), thread="blockIdx.x"): for i0_i1_i2_i3_fused_1 in T.thread_binding(T.int64(72), thread="threadIdx.x"): @@ -503,7 +503,7 @@ def test_scalar_add(): @tvm.script.ir_module class Before: @Ts.prim_func - def add(rxplaceholder: T.Buffer((), "int64"), T_add: T.Buffer((), "int64")): + def add(rxplaceholder: T.Tensor((), "int64"), T_add: T.Tensor((), "int64")): T.func_attr({"tirx.noalias": True}) with Ts.sblock("T_add"): vi = Ts.axis.spatial(1, T.int64(0)) @@ -514,7 +514,7 @@ def add(rxplaceholder: T.Buffer((), "int64"), T_add: T.Buffer((), "int64")): @tvm.script.ir_module class Expected: @Ts.prim_func - def add(rxplaceholder: T.Buffer((), "int64"), T_add: T.Buffer((), "int64")): + def add(rxplaceholder: T.Tensor((), "int64"), T_add: T.Tensor((), "int64")): T.func_attr({"tirx.is_scheduled": True, "tirx.noalias": True}) # with Ts.sblock("root"): for u_fused_0 in T.thread_binding(1, thread="blockIdx.x"): @@ -539,7 +539,7 @@ def test_sum(): @tvm.script.ir_module class Before: @Ts.prim_func - def sum(A: T.Buffer((T.int64(2), T.int64(2)), "float64"), A_red: T.Buffer((), "float64")): + def sum(A: T.Tensor((T.int64(2), T.int64(2)), "float64"), A_red: T.Tensor((), "float64")): for k0, k1 in T.grid(T.int64(2), T.int64(2)): with Ts.sblock("A_red"): v_k0, v_k1 = Ts.axis.remap("RR", [k0, k1]) @@ -550,7 +550,7 @@ def sum(A: T.Buffer((T.int64(2), T.int64(2)), "float64"), A_red: T.Buffer((), "f @tvm.script.ir_module class Expected: @Ts.prim_func - def sum(A: T.Buffer((T.int64(2), T.int64(2)), "float64"), A_red: T.Buffer((), "float64")): + def sum(A: T.Tensor((T.int64(2), T.int64(2)), "float64"), A_red: T.Tensor((), "float64")): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): for u_fused_0 in T.thread_binding(1, thread="blockIdx.x"): @@ -580,14 +580,14 @@ def test_scalar_block_no_loops(): @tvm.script.ir_module class Before: @Ts.prim_func - def scalar_add(a: T.Buffer((), "float32"), b: T.Buffer((), "float32"), c: T.Buffer((), "float32")): + def scalar_add(a: T.Tensor((), "float32"), b: T.Tensor((), "float32"), c: T.Tensor((), "float32")): with Ts.sblock("scalar_add"): c[()] = a[()] + b[()] @tvm.script.ir_module class Expected: @Ts.prim_func - def scalar_add(a: T.Buffer((), "float32"), b: T.Buffer((), "float32"), c: T.Buffer((), "float32")): + def scalar_add(a: T.Tensor((), "float32"), b: T.Tensor((), "float32"), c: T.Tensor((), "float32")): T.func_attr({"tirx.is_scheduled": True}) # with Ts.sblock("root"): for u_fused_0 in T.thread_binding(1, thread="blockIdx.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 425d15d8db15..2da45f57034f 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 @@ -31,7 +31,7 @@ def _narrow(func): def test_block(): @Ts.prim_func(private=True) - def before(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def before(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): for i in T.serial(0, T.int64(16)): for j in T.serial(0, T.int64(8)): with Ts.sblock(): @@ -39,7 +39,7 @@ def before(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): B[vi] = A[vi] + T.float32(1) @Ts.prim_func(private=True) - def expected(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def expected(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): for i in T.serial(0, T.int32(16)): for j in T.serial(0, T.int32(8)): with Ts.sblock(): @@ -54,8 +54,8 @@ def test_block_iters_used_only_in_regions(): @Ts.prim_func(private=True) def before( - A: T.Buffer((T.int64(16), T.int64(16)), "float32"), - B: T.Buffer((T.int64(16), T.int64(16)), "float32"), + A: T.Tensor((T.int64(16), T.int64(16)), "float32"), + B: T.Tensor((T.int64(16), T.int64(16)), "float32"), ): for i_o, j_o in T.grid(T.int64(2), T.int64(2)): with Ts.sblock("tile_o"): @@ -94,7 +94,7 @@ def before( B_tile[vi_i, vj_i] = A_tile[vi_i, vj_i] + T.float32(1) @Ts.prim_func(private=True) - def expected(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")): + def expected(A: T.Tensor((16, 16), "float32"), B: T.Tensor((16, 16), "float32")): for i_o, j_o in T.grid(2, 2): with Ts.sblock("tile_o"): vi_o, vj_o = Ts.axis.remap("SS", [i_o, j_o]) @@ -116,7 +116,7 @@ def expected(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16), "float32")) def test_fail_on_buffer_param(): @Ts.prim_func(private=True) - def func(A: T.Buffer((128,), "int64"), B: T.Buffer((128,), "int64")): + def func(A: T.Tensor((128,), "int64"), B: T.Tensor((128,), "int64")): for i in T.serial(0, 16): for j in T.serial(0, 8): with Ts.sblock(): @@ -129,7 +129,7 @@ def func(A: T.Buffer((128,), "int64"), B: T.Buffer((128,), "int64")): def test_fail_on_block_alloc_buffer(): @Ts.prim_func(private=True) - def func(A: T.Buffer((128,), "int32"), B: T.Buffer((128,), "int32")): + def func(A: T.Tensor((128,), "int32"), B: T.Tensor((128,), "int32")): C = Ts.sblock_alloc_buffer((128,), "int64") for i in T.serial(0, 16): for j in T.serial(0, 8): @@ -154,9 +154,9 @@ def test_metal_simdgroup_matmul_builds(): @Ts.prim_func def main( - A: T.Buffer((T.int64(1), n, T.int64(256)), "float16"), - B: T.Buffer((T.int64(256), T.int64(256)), "float16"), - C: T.Buffer((T.int64(1), n, T.int64(256)), "float16"), + A: T.Tensor((T.int64(1), n, T.int64(256)), "float16"), + B: T.Tensor((T.int64(256), T.int64(256)), "float16"), + C: T.Tensor((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)): with Ts.sblock("NT_matmul"): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_hoist_expression.py b/tests/python/s_tir/transform/test_s_tir_transform_hoist_expression.py index fe141edf1a4f..33374b2dbc51 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_hoist_expression.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_hoist_expression.py @@ -42,13 +42,13 @@ def _run_transform(before, hoisted_conditionals, hoisted_let_bindings): def test_hoist_to_top_if_else_stmt(): @Ts.prim_func(private=True) - def before(A: T.Buffer((16,), "float32"), n: T.int32): + def before(A: T.Tensor((16,), "float32"), n: T.int32): for i in T.serial(16): if n != 0: A[i] = 0.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((16,), "float32"), n: T.int32): + def expected(A: T.Tensor((16,), "float32"), n: T.int32): if n != 0: for i in T.serial(16): A[i] = 0.0 @@ -59,13 +59,13 @@ def expected(A: T.Buffer((16,), "float32"), n: T.int32): def test_hoist_to_top_all(): @Ts.prim_func(private=True) - def before(A: T.Buffer((16,), "float32"), n: T.int32): + def before(A: T.Tensor((16,), "float32"), n: T.int32): for i in T.serial(16): if n != 0: A[i] = 0.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((16,), "float32"), n: T.int32): + def expected(A: T.Tensor((16,), "float32"), n: T.int32): if n != 0: for i in T.serial(16): A[i] = 0.0 @@ -76,7 +76,7 @@ def expected(A: T.Buffer((16,), "float32"), n: T.int32): def test_suppress_hoist_if_else_never(): @Ts.prim_func(private=True) - def before(A: T.Buffer((16,), "float32"), n: T.int32): + def before(A: T.Tensor((16,), "float32"), n: T.int32): for i in T.serial(16): if n != 0: A[i] = 0.0 @@ -89,7 +89,7 @@ def before(A: T.Buffer((16,), "float32"), n: T.int32): def test_suppress_hoist_if_else_expr_only(): @Ts.prim_func(private=True) - def before(A: T.Buffer((16,), "float32"), n: T.int32): + def before(A: T.Tensor((16,), "float32"), n: T.int32): for i in T.serial(16): if n != 0: A[i] = 0.0 @@ -102,7 +102,7 @@ def before(A: T.Buffer((16,), "float32"), n: T.int32): def test_hoist_block_var(): @Ts.prim_func(private=True) - def before(A: T.Buffer((128, 16), "float32"), n: T.int32): + def before(A: T.Tensor((128, 16), "float32"), n: T.int32): i = T.env_thread("threadIdx.x") T.launch_thread(i, 128) @@ -111,7 +111,7 @@ def before(A: T.Buffer((128, 16), "float32"), n: T.int32): A[i, j] = 0.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((128, 16), "float32"), n: T.int32): + def expected(A: T.Tensor((128, 16), "float32"), n: T.int32): i = T.env_thread("threadIdx.x") T.launch_thread(i, 128) @@ -125,7 +125,7 @@ def expected(A: T.Buffer((128, 16), "float32"), n: T.int32): def test_suppress_hoist_block_var(): @Ts.prim_func(private=True) - def before(A: T.Buffer((128, 16), "float32"), n: T.int32): + def before(A: T.Tensor((128, 16), "float32"), n: T.int32): thread_x = T.env_thread("threadIdx.x") T.launch_thread(thread_x, 128) @@ -146,7 +146,7 @@ def before(A: T.Buffer((128, 16), "float32"), n: T.int32): def test_hoist_across_block_var(): @Ts.prim_func(private=True) - def before(A: T.Buffer((128, 16), "float32"), n: T.int32): + def before(A: T.Tensor((128, 16), "float32"), n: T.int32): thread_x = T.env_thread("threadIdx.x") T.launch_thread(thread_x, 128) @@ -156,7 +156,7 @@ def before(A: T.Buffer((128, 16), "float32"), n: T.int32): A[i, j] = 0.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((128, 16), "float32"), n: T.int32): + def expected(A: T.Tensor((128, 16), "float32"), n: T.int32): thread_x = T.env_thread("threadIdx.x") if n == 0: @@ -171,7 +171,7 @@ def expected(A: T.Buffer((128, 16), "float32"), n: T.int32): def test_suppress_hoist_across_block_var(): @Ts.prim_func(private=True) - def before(A: T.Buffer((128, 16), "float32"), n: T.int32): + def before(A: T.Tensor((128, 16), "float32"), n: T.int32): thread_x = T.env_thread("threadIdx.x") T.launch_thread(thread_x, 128) @@ -181,7 +181,7 @@ def before(A: T.Buffer((128, 16), "float32"), n: T.int32): A[i, j] = 0.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((128, 16), "float32"), n: T.int32): + def expected(A: T.Tensor((128, 16), "float32"), n: T.int32): thread_x = T.env_thread("threadIdx.x") T.launch_thread(thread_x, 128) @@ -200,14 +200,14 @@ def expected(A: T.Buffer((128, 16), "float32"), n: T.int32): def test_hoist_to_middle(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): for j in T.serial(4): if i < 3: A[i, j] = 0.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): if i < 3: for j in T.serial(4): @@ -219,7 +219,7 @@ def expected(A: T.Buffer((4, 4), "float32")): def test_hoist_with_let(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): for j in T.serial(4): condition: T.let[T.bool] = i < 3 @@ -227,7 +227,7 @@ def before(A: T.Buffer((4, 4), "float32")): A[i, j] = 0.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): condition: T.let[T.bool] = i < 3 # noqa: F841 if i < 3: @@ -248,7 +248,7 @@ def test_hoist_disable_let(): """ @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): for j in T.serial(4): condition: T.let[T.bool] = i < 3 @@ -256,7 +256,7 @@ def before(A: T.Buffer((4, 4), "float32")): A[i, j] = 0.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32")): for i, j in T.grid(4, 4): condition: T.let[T.bool] = i < 3 # noqa: F841 if i < 3: @@ -268,7 +268,7 @@ def expected(A: T.Buffer((4, 4), "float32")): def test_hoist_if_else(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): for j in T.serial(4): if i < 3: @@ -277,7 +277,7 @@ def before(A: T.Buffer((4, 4), "float32")): A[i, j] = 1.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): if i < 3: for j in T.serial(4): @@ -292,7 +292,7 @@ def expected(A: T.Buffer((4, 4), "float32")): def test_hoist_sequential_assign(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32"), B: T.Tensor((4, 4), "float32")): for i in T.serial(4): for j in T.serial(4): if i < 3: @@ -303,7 +303,7 @@ def before(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4, 4), "float32")): B[i, j] = 1.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32"), B: T.Tensor((4, 4), "float32")): for i in T.serial(4): if i < 3: for j in T.serial(4): @@ -320,7 +320,7 @@ def expected(A: T.Buffer((4, 4), "float32"), B: T.Buffer((4, 4), "float32")): def test_hoist_multi_if(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): for j in T.serial(4): for k in T.serial(4): @@ -329,7 +329,7 @@ def before(A: T.Buffer((4, 4), "float32")): A[i, j] = 0.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): if i < 2: for j in T.serial(4): @@ -343,13 +343,13 @@ def expected(A: T.Buffer((4, 4), "float32")): def test_hoist_complex_conditional(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i, j, k in T.grid(4, 4, 4): if j < 3 and i < 2: A[i, j] = 0.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): if i < 2: for j in T.serial(4): @@ -363,13 +363,13 @@ def expected(A: T.Buffer((4, 4), "float32")): def test_suppress_splitting_conditional(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i, j, k in T.grid(4, 4, 4): if j < 3 and i < 2: A[i, j] = 0.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32")): for i, j in T.grid(4, 4): if j < 3 and i < 2: for k in T.serial(4): @@ -385,7 +385,7 @@ def expected(A: T.Buffer((4, 4), "float32")): def test_hoist_multi_if_else(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): for j in T.serial(4): for k in T.serial(4): @@ -401,7 +401,7 @@ def before(A: T.Buffer((4, 4), "float32")): A[i, j] = 3.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): if i < 2: for j in T.serial(4): @@ -426,7 +426,7 @@ def expected(A: T.Buffer((4, 4), "float32")): def test_hoist_multi_if_else_different_branches(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): for j in T.serial(4): for k in T.serial(4): @@ -442,7 +442,7 @@ def before(A: T.Buffer((4, 4), "float32")): A[i, j] = 3.0 @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): if i < 2: if i < 1: @@ -476,12 +476,12 @@ def expected(A: T.Buffer((4, 4), "float32")): def test_hoist_if_else_expr(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i, j in T.grid(4, 4): A[i, j] = T.if_then_else(i < 2, 1.0, 2.0) @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): if i < 2: for j in T.serial(4): @@ -496,7 +496,7 @@ def expected(A: T.Buffer((4, 4), "float32")): def test_suppress_hoist_if_else_expr(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i, j in T.grid(4, 4): A[i, j] = T.if_then_else(i < 2, 1.0, 2.0) @@ -512,13 +512,13 @@ def before(A: T.Buffer((4, 4), "float32")): def test_hoist_let_expr(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i, j in T.grid(4, 4): x = T.float32() A[i, j] = T.Let(5.0 * x + T.cast(j, "float32"), where={x: T.cast(i + 1, "float32")}) @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32")): for i in T.serial(4): x: T.let[T.float32] = T.cast(i + 1, "float32") # noqa: F841 for j in T.serial(4): @@ -530,7 +530,7 @@ def expected(A: T.Buffer((4, 4), "float32")): def test_suppress_hoist_let_expr(): @Ts.prim_func(private=True) - def before(A: T.Buffer((4, 4), "float32")): + def before(A: T.Tensor((4, 4), "float32")): for i, j in T.grid(4, 4): x = T.float32() A[i, j] = T.Let(5.0 * x + T.cast(j, "float32"), where={x: T.cast(i + 1, "float32")}) 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 71c374ef1988..a52f565c43fb 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 @@ -117,7 +117,7 @@ def func(l: T.int32, m: T.int32, n: T.int32): def test_hoist_no_match_for(): @Ts.prim_func(private=True) def func(data: T.handle("float32"), l: T.int32, m: T.int32, n: T.int32): - data_ptr = T.decl_buffer(1, "float32", data=data) + data_ptr = T.decl_tensor(1, "float32", data=data) for i in T.serial(l): for j in T.serial(m): data_ptr[i * 3 + j] = data_ptr[i * 3 + j] + T.float32(0.5) @@ -163,7 +163,7 @@ def test_attr_stmt(): @Ts.prim_func(private=True) def func(data: T.handle("float32"), l: T.int32, m: T.int32, n: T.int32): - data_ptr = T.decl_buffer(1, "float32", data=data) + data_ptr = T.decl_tensor(1, "float32", data=data) tx = T.launch_thread("threadIdx.x", dshape[0]) bx = T.launch_thread("blockIdx.x", dshape[1]) for i in T.serial(l): @@ -194,7 +194,7 @@ def func(data: T.handle("float32"), l: T.int32, m: T.int32, n: T.int32): def test_nested_for(): @Ts.prim_func(private=True) def func(data: T.handle("float32")): - data_ptr = T.decl_buffer(1, "float32", data=data) + data_ptr = T.decl_tensor(1, "float32", data=data) for i in range(5): for j in range(10): if i >= 3: @@ -228,7 +228,7 @@ def test_if_block(): @I.ir_module class Module: @Ts.prim_func(private=True) - def main(data: T.Buffer((1,), "float32"), n: T.int32): + def main(data: T.Tensor((1,), "float32"), n: T.int32): # First loop nest: i, j, k, l for i in T.serial(5): for j in T.serial(10): @@ -273,7 +273,7 @@ def main(data: T.Buffer((1,), "float32"), n: T.int32): def test_multi_if(): @Ts.prim_func(private=True) def func(data: T.handle("float32")): - data_ptr = T.decl_buffer(1, "float32", data=data) + data_ptr = T.decl_tensor(1, "float32", data=data) for i in range(10): for j in range(10): for k in range(10): @@ -299,7 +299,7 @@ def func(data: T.handle("float32")): def test_no_hoisting_1(): @Ts.prim_func(private=True) def func(data: T.handle("float32")): - data_ptr = T.decl_buffer(1, "float32", data=data) + data_ptr = T.decl_tensor(1, "float32", data=data) for i in range(10): for j in range(10): for k in range(10): @@ -323,7 +323,7 @@ def func(data: T.handle("float32")): def test_no_hoisting_2(): @Ts.prim_func(private=True) def func(data: T.handle("float32")): - data_ptr = T.decl_buffer(1, "float32", data=data) + data_ptr = T.decl_tensor(1, "float32", data=data) for i in range(10): for j in range(10): for k in range(10): @@ -358,7 +358,7 @@ def test_no_hoisting_4(): @I.ir_module class Module: @Ts.prim_func(private=True) - def main(data: T.Buffer((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): + def main(data: T.Tensor((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): bx = T.launch_thread("blockIdx.x", dshape[1]) for i in T.serial(l): for j in T.serial(m): @@ -390,7 +390,7 @@ def test_no_hoisting_6(): @I.ir_module class Module: @Ts.prim_func(private=True) - def main(data: T.Buffer((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): + def main(data: T.Tensor((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): tx = T.launch_thread("threadIdx.x", dshape[0]) bx = T.launch_thread("blockIdx.x", dshape[1]) for i in T.serial(l): @@ -418,7 +418,7 @@ def test_no_hoisting_7(): @I.ir_module class Module: @Ts.prim_func(private=True) - def main(data: T.Buffer((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): + def main(data: T.Tensor((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): tx = T.launch_thread("threadIdx.x", dshape[0]) bx = T.launch_thread("blockIdx.x", dshape[1]) for i in T.serial(l): @@ -453,7 +453,7 @@ def test_hoisting_block_scope_2(): @I.ir_module class Module: @Ts.prim_func(private=True) - def main(data: T.Buffer((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): + def main(data: T.Tensor((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): tx = T.launch_thread("threadIdx.x", dshape[0]) for i in T.serial(l): for j in T.serial(m): @@ -487,7 +487,7 @@ def test_hoisting_block_scope_5(): @I.ir_module class Module: @Ts.prim_func(private=True) - def main(data: T.Buffer((1,), "float32"), l: T.int32, m: T.int32, n: T.int32, g: T.int32): + def main(data: T.Tensor((1,), "float32"), l: T.int32, m: T.int32, n: T.int32, g: T.int32): for i in T.serial(l): for j in T.serial(m): for k in T.serial(n): @@ -516,7 +516,7 @@ def test_hoisting_block_scope_6(): @I.ir_module class Module: @Ts.prim_func(private=True) - def main(data: T.Buffer((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): + def main(data: T.Tensor((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): tx = T.launch_thread("threadIdx.x", dshape[0]) bx = T.launch_thread("blockIdx.x", dshape[1]) for i in T.serial(l): @@ -544,7 +544,7 @@ def test_hoisting_block_scope_7(): @I.ir_module class Module: @Ts.prim_func(private=True) - def main(data: T.Buffer((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): + def main(data: T.Tensor((1,), "float32"), l: T.int32, m: T.int32, n: T.int32): tx = T.launch_thread("threadIdx.x", dshape[0]) bx = T.launch_thread("blockIdx.x", dshape[1]) for i in T.serial(l): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_inject_double_buffer.py b/tests/python/s_tir/transform/test_s_tir_transform_inject_double_buffer.py index e00e38d8558f..a9d49efa0dc1 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_inject_double_buffer.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_inject_double_buffer.py @@ -42,11 +42,11 @@ def test_double_buffer(): class Module: @Ts.prim_func def db(A: T.handle("float32"), C: T.handle("float32")): - A_buf = T.decl_buffer((n * m,), "float32", data=A) - C_buf = T.decl_buffer((m,), "float32", data=C) + A_buf = T.decl_tensor((n * m,), "float32", data=A) + C_buf = T.decl_tensor((m,), "float32", data=C) tx = T.launch_thread("threadIdx.x", 1) for i in range(n): - B = T.alloc_buffer((m,), "float32", scope="shared") + B = T.alloc_tensor((m,), "float32", scope="shared") with T.attr(B.data, "double_buffer_scope", 1): for j in range(m): B[j] = A_buf[i * 4 + j] @@ -68,7 +68,7 @@ def db(A: T.handle("float32"), C: T.handle("float32")): def visitor(op): nonlocal allocate_node - if _is_buffer_binding(op, "tirx.alloc_buffer") and "B" in str(op.var.data): + if _is_buffer_binding(op, "tirx.alloc_tensor") and "B" in str(op.var.data): allocate_node = op tvm_ffi.structural_walk(stmt, visitor) @@ -97,9 +97,9 @@ def test_double_buffer_transform(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer([16, 32], "float32"), B: T.Buffer(16, "float32")): + def main(A: T.Tensor([16, 32], "float32"), B: T.Tensor(16, "float32")): for i in range(16): - cache = T.alloc_buffer((32,), "float32") + cache = T.alloc_tensor((32,), "float32") T.attr(cache.data, "double_buffer_scope", 1) @@ -116,7 +116,7 @@ def main(A: T.Buffer([16, 32], "float32"), B: T.Buffer(16, "float32")): def visitor(op): nonlocal allocate_node - if _is_buffer_binding(op, "tirx.alloc_buffer"): + if _is_buffer_binding(op, "tirx.alloc_tensor"): allocate_node = op tvm_ffi.structural_walk(After["main"].body, visitor) @@ -137,9 +137,9 @@ def test_double_buffer_with_decl_buffer(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((16, 32), "float32"), B: T.Buffer(16, "float32")): + def main(A: T.Tensor((16, 32), "float32"), B: T.Tensor(16, "float32")): for i in range(16): - cache = T.decl_buffer(32, "float32") + cache = T.decl_tensor(32, "float32") T.attr(cache.data, "double_buffer_scope", 1) for j in range(32): @@ -152,8 +152,8 @@ def main(A: T.Buffer((16, 32), "float32"), B: T.Buffer(16, "float32")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((16, 32), "float32"), B: T.Buffer(16, "float32")): - cache = T.decl_buffer(64, "float32") + def main(A: T.Tensor((16, 32), "float32"), B: T.Tensor(16, "float32")): + cache = T.decl_tensor(64, "float32") for j in range(32): cache[j] = A[0, j] diff --git a/tests/python/s_tir/transform/test_s_tir_transform_inject_permuted_layout.py b/tests/python/s_tir/transform/test_s_tir_transform_inject_permuted_layout.py index de0bfc02f853..3ee02c829ffc 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_inject_permuted_layout.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_inject_permuted_layout.py @@ -38,7 +38,7 @@ def _check_primfunc_transform(before: PrimFunc, expected: PrimFunc): def test_backward_compatibility_shared_a(): # fmt: off @Ts.prim_func - def before(X: T.Buffer((4096, 4096), "float16")): + def before(X: T.Tensor((4096, 4096), "float16")): # with Ts.sblock("root"): for blockIdx_y in T.thread_binding(256, thread="blockIdx.y"): for threadIdx_y in T.thread_binding(4, thread="threadIdx.y"): @@ -71,7 +71,7 @@ def before(X: T.Buffer((4096, 4096), "float16")): T.ptx_legacy.ldmatrix("float16", T.bool(False), 4, ".b16", X_reindex_shared_dyn_m16n8k8_matrixA.data, ax0_0 * 8, T.tvm_access_ptr(T.type_annotation("float16"), X_reindex_shared_dyn.data, threadIdx_y // 2 * 2048 + ax0_0 * 1024 + ax2_0_1 * 8, 1024, 1), threadIdx_x * 32) @Ts.prim_func - def expected(X: T.Buffer((4096, 4096), "float16")): + def expected(X: T.Tensor((4096, 4096), "float16")): for blockIdx_y in T.thread_binding(256, thread="blockIdx.y"): for threadIdx_y in T.thread_binding(4, thread="threadIdx.y"): for threadIdx_x in T.thread_binding(32, thread="threadIdx.x"): @@ -101,7 +101,7 @@ def expected(X: T.Buffer((4096, 4096), "float16")): def test_backward_compatibility_shared_a_and_b(): # fmt: off @Ts.prim_func - def before(X: T.Buffer((4096, 4096), "float16"), Y: T.Buffer((4096, 4096), "float16")): + def before(X: T.Tensor((4096, 4096), "float16"), Y: T.Tensor((4096, 4096), "float16")): for blockIdx_x in T.thread_binding(4, thread="blockIdx.x"): for blockIdx_y in T.thread_binding(256, thread="blockIdx.y"): for threadIdx_y in T.thread_binding(4, thread="threadIdx.y"): @@ -139,7 +139,7 @@ def before(X: T.Buffer((4096, 4096), "float16"), Y: T.Buffer((4096, 4096), "floa T.ptx_legacy.ldmatrix("float16", T.bool(True), 4, ".b16", Y_reindex_shared_dyn_m16n8k8_matrixB.data, ax1_0 * 8, T.tvm_access_ptr(T.type_annotation("float16"), Y_reindex_shared_dyn.data, ax2_0_1 * 1024 + threadIdx_y % 2 * 64 + ax1_0 * 32, 1024, 1), threadIdx_x % 8 * 128 + threadIdx_x // 8 * 8) @Ts.prim_func - def expected(X: T.Buffer((4096, 4096), "float16"), Y: T.Buffer((4096, 4096), "float16")): + def expected(X: T.Tensor((4096, 4096), "float16"), Y: T.Tensor((4096, 4096), "float16")): for blockIdx_x in T.thread_binding(4, thread="blockIdx.x"): for blockIdx_y in T.thread_binding(256, thread="blockIdx.y"): for threadIdx_y in T.thread_binding(4, thread="threadIdx.y"): @@ -186,7 +186,7 @@ def expected(X: T.Buffer((4096, 4096), "float16"), Y: T.Buffer((4096, 4096), "fl def test_buffer_a(): # fmt: off @Ts.prim_func - def before(A: T.Buffer((T.int64(128), T.int64(32)), 'float16')): + def before(A: T.Tensor((T.int64(128), T.int64(32)), 'float16')): A_shared_dyn = Ts.sblock_alloc_buffer((T.int64(128), T.int64(32)), "float16", scope="shared.dyn") A_warp = Ts.sblock_alloc_buffer((T.int64(4), T.int64(1), T.int64(32), T.int64(8)), "float16", scope="warp") @@ -222,7 +222,7 @@ def before(A: T.Buffer((T.int64(128), T.int64(32)), 'float16')): ) @Ts.prim_func - def expected(A: T.Buffer((T.int64(128), T.int64(32)), "float16")): + def expected(A: T.Tensor((T.int64(128), T.int64(32)), "float16")): A_shared_dyn = Ts.sblock_alloc_buffer((T.int64(128), T.int64(32)), "float16", scope="shared.dyn") A_warp = Ts.sblock_alloc_buffer((T.int64(4), T.int64(1), T.int64(32), T.int64(8)), "float16", scope="warp") for threadIdx_z in T.thread_binding(T.int64(2), thread="threadIdx.z"): @@ -250,7 +250,7 @@ def expected(A: T.Buffer((T.int64(128), T.int64(32)), "float16")): def test_buffer_b(): # fmt: off @Ts.prim_func - def before(B: T.Buffer((T.int64(128), T.int64(32)), "float16")): + def before(B: T.Tensor((T.int64(128), T.int64(32)), "float16")): B_shared_dyn = Ts.sblock_alloc_buffer((T.int64(128), T.int64(32)), "float16", scope="shared.dyn") for threadIdx_z in T.thread_binding(T.int64(2), thread="threadIdx.z"): for threadIdx_y in T.thread_binding(T.int64(2), thread="threadIdx.y"): @@ -272,7 +272,7 @@ def before(B: T.Buffer((T.int64(128), T.int64(32)), "float16")): T.ptx_legacy.ldmatrix("float16", T.bool(False), 4, ".b16", B_warp.data, v1 * T.int64(256) + threadIdx_x * T.int64(8), T.tvm_access_ptr(T.type_annotation("float16"), B_shared_dyn.data, threadIdx_y * T.int64(2048) + v1 * T.int64(512) + v0 * T.int64(16), T.int64(512), 1), threadIdx_x // T.int64(16) * T.int64(256) + threadIdx_x % T.int64(8) * T.int64(32) + threadIdx_x % T.int64(16) // T.int64(8) * T.int64(8)) @Ts.prim_func - def expected(B: T.Buffer((T.int64(128), T.int64(32)), "float16")): + def expected(B: T.Tensor((T.int64(128), T.int64(32)), "float16")): B_shared_dyn = Ts.sblock_alloc_buffer((T.int64(128), T.int64(32)), "float16", scope="shared.dyn") for threadIdx_z in T.thread_binding(T.int64(2), thread="threadIdx.z"): for threadIdx_y in T.thread_binding(T.int64(2), thread="threadIdx.y"): @@ -302,7 +302,7 @@ def expected(B: T.Buffer((T.int64(128), T.int64(32)), "float16")): def test_buffer_c_fp32(): # fmt: off @Ts.prim_func - def before(O: T.Buffer((T.int64(128), T.int64(128)), 'float16')): + def before(O: T.Tensor((T.int64(128), T.int64(128)), 'float16')): O_shared_dyn = Ts.sblock_alloc_buffer((T.int64(128), T.int64(128)), scope="shared.dyn") O_warp = Ts.sblock_alloc_buffer((T.int64(4), T.int64(4), T.int64(32), T.int64(8)), scope="warp") @@ -322,7 +322,7 @@ def before(O: T.Buffer((T.int64(128), T.int64(128)), 'float16')): O[v0 * T.int64(8) + threadIdx_z * T.int64(4) + threadIdx_y * T.int64(2) + threadIdx_x // T.int64(16), threadIdx_x % T.int64(16) * T.int64(8) + v1] = T.Cast("float16", O_shared_dyn[v0 * T.int64(8) + threadIdx_z * T.int64(4) + threadIdx_y * T.int64(2) + threadIdx_x // T.int64(16), threadIdx_x % T.int64(16) * T.int64(8) + v1]) @Ts.prim_func - def expected(O: T.Buffer((T.int64(128), T.int64(128)), "float16")): + def expected(O: T.Tensor((T.int64(128), T.int64(128)), "float16")): # with Ts.sblock("root"): O_shared_dyn = Ts.sblock_alloc_buffer((T.int64(128), T.int64(128)), scope="shared.dyn") O_warp = Ts.sblock_alloc_buffer((T.int64(4), T.int64(4), T.int64(32), T.int64(8)), scope="warp") 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 af1ebe08ac10..e71e0dff1b3e 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 @@ -35,7 +35,7 @@ def test_cp_async_raw_dtype_round_trips(): # (it prints dtype-first via tirx.ptx.cp_async_raw). Guards the regression # where the element dtype was dropped after the flat op was phased out. @T.prim_func - def f(A: T.Buffer((128,), "float16"), B: T.Buffer((128,), "float16")): + def f(A: T.Tensor((128,), "float16"), B: T.Tensor((128,), "float16")): T.func_attr({"global_symbol": "f"}) for i in T.serial(8): T.s_tir.cp_async_raw("float16", B.data, i * 16, A.data, i * 16, 16) @@ -63,7 +63,7 @@ def generate_global_to_shared_vectorized_copy(dtype, vector_size): @Ts.prim_func def ptx_global_to_shared_copy( - A: T.Buffer((32, 128), dtype), B: T.Buffer((32, 128), dtype) + A: T.Tensor((32, 128), dtype), B: T.Tensor((32, 128), dtype) ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) bx = T.env_thread("blockIdx.x") @@ -91,7 +91,7 @@ def ptx_global_to_shared_copy( @Ts.prim_func def ptx_global_to_shared_copy_fp32x1( - A: T.Buffer((32, 128), "float32"), B: T.Buffer((32, 128), "float32") + A: T.Tensor((32, 128), "float32"), B: T.Tensor((32, 128), "float32") ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) bx = T.env_thread("blockIdx.x") @@ -116,9 +116,9 @@ def ptx_global_to_shared_copy_fp32x1( @Ts.prim_func def ptx_global_to_shared_dyn_copy_fp16x8( - A: T.Buffer((32, 128), "float16"), - B: T.Buffer((32, 128), "float16"), - C: T.Buffer((32, 128), "float16"), + A: T.Tensor((32, 128), "float16"), + B: T.Tensor((32, 128), "float16"), + C: T.Tensor((32, 128), "float16"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) bx = T.env_thread("blockIdx.x") @@ -378,9 +378,9 @@ def tvm_callback_cuda_postproc(code, _): def test_cp_async_in_if_then_else(postproc_if_missing_async_support): @Ts.prim_func def simple_compute( - A: T.Buffer((16, 14), "float32"), - B: T.Buffer((16, 14), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 14), "float32"), + B: T.Tensor((16, 14), "float32"), + C: T.Tensor((16, 16), "float32"), ): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for tx in T.thread_binding(0, 16, thread="threadIdx.x"): @@ -461,9 +461,9 @@ def test_vectorize_cp_async_in_if_then_else(postproc_if_missing_async_support): @Ts.prim_func def complex_compute( - A: T.Buffer((2, 16, 16, 1280), "float16"), - W: T.Buffer((1280, 3, 3, 1280), "float16"), - Conv: T.Buffer((512, 1280), "float16"), + A: T.Tensor((2, 16, 16, 1280), "float16"), + W: T.Tensor((1280, 3, 3, 1280), "float16"), + Conv: T.Tensor((512, 1280), "float16"), ): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # with Ts.sblock("root"): @@ -903,10 +903,10 @@ def test_multiplication_nodes_are_inlined(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((32, 128), "float16")): + def main(A: T.Tensor((32, 128), "float16")): tx = T.launch_thread("threadIdx.x", T.int64(32)) - A_flattened = T.decl_buffer((4096,), "float16", data=A.data) - A_shared = T.decl_buffer([4096], "float16", scope="shared") + A_flattened = T.decl_tensor((4096,), "float16", data=A.data) + A_shared = T.decl_tensor([4096], "float16", scope="shared") T.attr("default", "async_scope", 1) for i in range(16): @@ -920,10 +920,10 @@ def main(A: T.Buffer((32, 128), "float16")): @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer((32, 128), "float16")): + def main(A: T.Tensor((32, 128), "float16")): tx = T.launch_thread("threadIdx.x", T.int64(32)) - A_flattened = T.decl_buffer((4096,), "float16", data=A.data) - A_shared = T.decl_buffer((4096,), "float16", scope="shared") + A_flattened = T.decl_tensor((4096,), "float16", data=A.data) + A_shared = T.decl_tensor((4096,), "float16", scope="shared") for i in range(16): cse_v1: T.int64 = T.Cast("int64", i) T.s_tir.cp_async_raw( 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 6e093cc430f4..438dcaa923a4 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 @@ -56,7 +56,7 @@ def _check_error(func): @Ts.prim_func -def trivial_pipeline(A: T.Buffer((16, 1), "float32"), C: T.Buffer((16, 1), "float32")): +def trivial_pipeline(A: T.Tensor((16, 1), "float32"), C: T.Tensor((16, 1), "float32")): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( 0, 1, annotations={"software_pipeline_stage": [0, 1], "software_pipeline_order": [0, 1]} @@ -77,7 +77,7 @@ def trivial_pipeline(A: T.Buffer((16, 1), "float32"), C: T.Buffer((16, 1), "floa @Ts.prim_func def transformed_trivial_pipeline( - A: T.Buffer((16, 1), "float32"), C: T.Buffer((16, 1), "float32") + A: T.Tensor((16, 1), "float32"), C: T.Tensor((16, 1), "float32") ) -> None: for tx in T.thread_binding(16, thread="threadIdx.x"): with Ts.sblock(): @@ -100,7 +100,7 @@ def transformed_trivial_pipeline( def gen_simple_compute(num_stages): @Ts.prim_func - def simple_compute(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): + def simple_compute(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( 0, @@ -128,7 +128,7 @@ def simple_compute(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "floa @Ts.prim_func def transformed_simple_compute( - A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32") ) -> None: for tx in T.thread_binding(0, 16, thread="threadIdx.x"): with Ts.sblock(): @@ -161,7 +161,7 @@ def transformed_simple_compute( @Ts.prim_func -def dynamic_compute(A: T.Buffer((16, k), "float32"), C: T.Buffer((16, k), "float32")): +def dynamic_compute(A: T.Tensor((16, k), "float32"), C: T.Tensor((16, k), "float32")): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( 0, @@ -189,7 +189,7 @@ def dynamic_compute(A: T.Buffer((16, k), "float32"), C: T.Buffer((16, k), "float @Ts.prim_func -def transformed_dynamic_compute(A: T.Buffer((16, k), "float32"), C: T.Buffer((16, k), "float32")): +def transformed_dynamic_compute(A: T.Tensor((16, k), "float32"), C: T.Tensor((16, k), "float32")): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): with Ts.sblock(): Ts.reads(A[tx, 0 : T.max(1, k)]) @@ -227,7 +227,7 @@ def transformed_dynamic_compute(A: T.Buffer((16, k), "float32"), C: T.Buffer((16 @Ts.prim_func def simple_compute_with_other_annotation( - A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32") ): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( @@ -255,7 +255,7 @@ def simple_compute_with_other_annotation( @Ts.prim_func def transformed_simple_compute_with_other_annotation( - A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32") ) -> None: for tx in T.thread_binding(0, 16, thread="threadIdx.x"): with Ts.sblock(): @@ -289,7 +289,7 @@ def transformed_simple_compute_with_other_annotation( @Ts.prim_func -def three_stage_compute(A: T.Buffer((16, 16), "float32"), D: T.Buffer((16, 16), "float32")): +def three_stage_compute(A: T.Tensor((16, 16), "float32"), D: T.Tensor((16, 16), "float32")): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( 0, @@ -320,7 +320,7 @@ def three_stage_compute(A: T.Buffer((16, 16), "float32"), D: T.Buffer((16, 16), @Ts.prim_func def transformed_three_stage_compute( - A: T.Buffer((16, 16), "float32"), D: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), D: T.Tensor((16, 16), "float32") ) -> None: for tx in T.thread_binding(16, thread="threadIdx.x"): with Ts.sblock(): @@ -374,9 +374,9 @@ def transformed_three_stage_compute( @Ts.prim_func def dag_interleaving( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( @@ -418,9 +418,9 @@ def dag_interleaving( @Ts.prim_func def transformed_dag_interleaving( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: for tx in T.thread_binding(16, thread="threadIdx.x"): with Ts.sblock(): @@ -483,7 +483,7 @@ def transformed_dag_interleaving( @Ts.prim_func def nested_pipeline_simple( - A: T.Buffer((16, 16, 16), "float32"), C: T.Buffer((16, 16, 16), "float32") + A: T.Tensor((16, 16, 16), "float32"), C: T.Tensor((16, 16, 16), "float32") ): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( @@ -527,7 +527,7 @@ def nested_pipeline_simple( @Ts.prim_func def transformed_nested_pipeline_simple( - A: T.Buffer((16, 16, 16), "float32"), C: T.Buffer((16, 16, 16), "float32") + A: T.Tensor((16, 16, 16), "float32"), C: T.Tensor((16, 16, 16), "float32") ) -> None: for tx in T.thread_binding(0, 16, thread="threadIdx.x"): with Ts.sblock(): @@ -604,7 +604,7 @@ def transformed_nested_pipeline_simple( @Ts.prim_func def nested_pipeline_prefetch_inner( - A: T.Buffer((16, 16, 16), "float32"), C: T.Buffer((16, 16, 16), "float32") + A: T.Tensor((16, 16, 16), "float32"), C: T.Tensor((16, 16, 16), "float32") ): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( @@ -648,7 +648,7 @@ def nested_pipeline_prefetch_inner( @Ts.prim_func def transformed_nested_pipeline_prefetch_inner( - A: T.Buffer((16, 16, 16), "float32"), C: T.Buffer((16, 16, 16), "float32") + A: T.Tensor((16, 16, 16), "float32"), C: T.Tensor((16, 16, 16), "float32") ) -> None: for tx in T.thread_binding(0, 16, thread="threadIdx.x"): with Ts.sblock(): @@ -728,7 +728,7 @@ def transformed_nested_pipeline_prefetch_inner( @Ts.prim_func def nested_pipeline_interleaving( - A: T.Buffer((16, 16, 16), "float32"), C: T.Buffer((16, 16, 16), "float32") + A: T.Tensor((16, 16, 16), "float32"), C: T.Tensor((16, 16, 16), "float32") ): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( @@ -778,7 +778,7 @@ def nested_pipeline_interleaving( @Ts.prim_func def transformed_nested_pipeline_interleaving( - A: T.Buffer((16, 16, 16), "float32"), C: T.Buffer((16, 16, 16), "float32") + A: T.Tensor((16, 16, 16), "float32"), C: T.Tensor((16, 16, 16), "float32") ) -> None: for tx in T.thread_binding(0, 16, thread="threadIdx.x"): with Ts.sblock(): @@ -887,7 +887,7 @@ def transformed_nested_pipeline_interleaving( @Ts.prim_func def nested_pipeline_double_buffer( - A: T.Buffer((16, 16, 16), "float32"), C: T.Buffer((16, 16, 16), "float32") + A: T.Tensor((16, 16, 16), "float32"), C: T.Tensor((16, 16, 16), "float32") ): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( @@ -938,7 +938,7 @@ def nested_pipeline_double_buffer( @Ts.prim_func def transformed_nested_pipeline_double_buffer( - A: T.Buffer((16, 16, 16), "float32"), C: T.Buffer((16, 16, 16), "float32") + A: T.Tensor((16, 16, 16), "float32"), C: T.Tensor((16, 16, 16), "float32") ) -> None: for tx in T.thread_binding(0, 16, thread="threadIdx.x"): with Ts.sblock(): @@ -1051,7 +1051,7 @@ def transformed_nested_pipeline_double_buffer( @Ts.prim_func def simple_compute_incorrect_reorder( - A: T.Buffer((16, 16), "float32"), D: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), D: T.Tensor((16, 16), "float32") ): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( @@ -1083,7 +1083,7 @@ def simple_compute_incorrect_reorder( @Ts.prim_func def simple_compute_conflicting_order( - A: T.Buffer((16, 16), "float32"), D: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), D: T.Tensor((16, 16), "float32") ): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial( @@ -1115,7 +1115,7 @@ def simple_compute_conflicting_order( @Ts.prim_func def simple_compute_missing_annotation( - A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32") ): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in T.serial(0, 16, annotations={"software_pipeline_stage": [0, 1]}): @@ -1194,7 +1194,7 @@ def test_simple_compute_async(): mod = tvm.s_tir.transform.InjectSoftwarePipeline()(sch.mod) @Ts.prim_func - def ref(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): + def ref(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")): for tx in T.thread_binding(16, thread="threadIdx.x"): with Ts.sblock(): Ts.reads(A[tx, 0:16]) @@ -1241,7 +1241,7 @@ def ref(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): mod = tvm.s_tir.transform.InjectSoftwarePipeline()(sch.mod) @Ts.prim_func - def ref(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: + def ref(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for tx in T.thread_binding(16, thread="threadIdx.x"): with Ts.sblock(): Ts.reads(A[tx, 0:16]) @@ -1294,9 +1294,9 @@ def ref(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> N def test_async_producer_interleaving(): @Ts.prim_func def simple_compute( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ): for tx in T.thread_binding(0, 16, thread="threadIdx.x"): for i in range(16): @@ -1329,9 +1329,9 @@ def simple_compute( @Ts.prim_func def ref( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: for tx in T.thread_binding(16, thread="threadIdx.x"): with Ts.sblock(): @@ -1408,7 +1408,7 @@ def test_three_stage_compute_two_stage_async(): mod = tvm.s_tir.transform.InjectSoftwarePipeline()(sch.mod) @Ts.prim_func - def ref(A: T.Buffer((16, 16), "float32"), D: T.Buffer((16, 16), "float32")) -> None: + def ref(A: T.Tensor((16, 16), "float32"), D: T.Tensor((16, 16), "float32")) -> None: for tx in T.thread_binding(16, thread="threadIdx.x"): with Ts.sblock(): Ts.reads(A[tx, 0:16]) @@ -1645,7 +1645,7 @@ def test_async_nested_pipeline_mma_gemm_ideal_annotation(): def test_less_loop_than_num_stage(): @Ts.prim_func - def before(A: T.Buffer((2,), "float32"), E: T.Buffer((2,), "float32")): + def before(A: T.Tensor((2,), "float32"), E: T.Tensor((2,), "float32")): for i in T.serial( 0, 2, @@ -1668,7 +1668,7 @@ def before(A: T.Buffer((2,), "float32"), E: T.Buffer((2,), "float32")): E[i] = D[0] + T.float32(5) @Ts.prim_func - def after(A: T.Buffer((2,), "float32"), E: T.Buffer((2,), "float32")): + def after(A: T.Tensor((2,), "float32"), E: T.Tensor((2,), "float32")): with Ts.sblock("root"): Ts.reads() Ts.writes() @@ -1722,7 +1722,7 @@ def test_less_loop_than_num_stage_dynamic(): K = T.dynamic("K", "int32") @Ts.prim_func - def before(A: T.Buffer([K], "float32"), E: T.Buffer([K], "float32")): + def before(A: T.Tensor([K], "float32"), E: T.Tensor([K], "float32")): for i in T.serial( 0, K, @@ -1747,7 +1747,7 @@ def before(A: T.Buffer([K], "float32"), E: T.Buffer([K], "float32")): K = T.dynamic("K", "int32") @Ts.prim_func - def after(A: T.Buffer([K], "float32"), E: T.Buffer([K], "float32")): + def after(A: T.Tensor([K], "float32"), E: T.Tensor([K], "float32")): with Ts.sblock("root"): Ts.reads() Ts.writes() diff --git a/tests/python/s_tir/transform/test_s_tir_transform_inject_virtual_thread.py b/tests/python/s_tir/transform/test_s_tir_transform_inject_virtual_thread.py index 953fad702493..36ea17d68cf8 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_inject_virtual_thread.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_inject_virtual_thread.py @@ -43,12 +43,12 @@ def test_vthread(): class Module: @Ts.prim_func def main(A: T.handle("float32"), C: T.handle("float32")): - A_buf = T.decl_buffer((n * nthread,), "float32", data=A) - C_buf = T.decl_buffer((n * nthread,), "float32", data=C) + A_buf = T.decl_tensor((n * nthread,), "float32", data=A) + C_buf = T.decl_tensor((n * nthread,), "float32", data=C) for i in range(n): vt_x = T.launch_thread("vthread", nthread) vt_y = T.launch_thread("vthread", nthread) - B = T.alloc_buffer((m,), scope="shared") + B = T.alloc_tensor((m,), scope="shared") B[i] = A_buf[i * nthread + vt_x] T.evaluate( T.call_extern( @@ -69,7 +69,7 @@ def main(A: T.handle("float32"), C: T.handle("float32")): allocates = [] def find_allocates(node): - if _is_buffer_binding(node, "tirx.alloc_buffer"): + if _is_buffer_binding(node, "tirx.alloc_tensor"): allocates.append(node) tvm_ffi.structural_walk(stmt.body, find_allocates) @@ -91,9 +91,9 @@ def main(): for i in range(n): vt_x = T.launch_thread("vthread", nthread) vt_y = T.launch_thread("vthread", nthread) - A = T.alloc_buffer((m,), scope="shared") - B = T.alloc_buffer((m,), scope="shared") - C = T.alloc_buffer((m,), scope="shared") + A = T.alloc_tensor((m,), scope="shared") + B = T.alloc_tensor((m,), scope="shared") + C = T.alloc_tensor((m,), scope="shared") A[vt_x] = T.Cast("float32", vt_x) + T.float32(1) B[vt_y] = T.Cast("float32", vt_y) + T.float32(1) T.evaluate( @@ -118,7 +118,7 @@ def main(): allocates = [] def find_allocates(node): - if _is_buffer_binding(node, "tirx.alloc_buffer"): + if _is_buffer_binding(node, "tirx.alloc_tensor"): allocates.append(node) tvm_ffi.structural_walk(stmt.body, find_allocates) @@ -137,10 +137,10 @@ class Module: @Ts.prim_func def main(A: T.handle("float32")): T.func_attr({"global_symbol": "main"}) - A_buf = T.decl_buffer((100 * nthread,), "float32", data=A) + A_buf = T.decl_tensor((100 * nthread,), "float32", data=A) for i in range(100): vt = T.launch_thread("vthread", nthread) - B = T.alloc_buffer((128,), scope="shared") + B = T.alloc_tensor((128,), scope="shared") if i == 0: B[i] = A_buf[i * nthread + vt] else: @@ -176,12 +176,12 @@ def test_vthread_simplified(): def before_func(): vthread = T.env_thread("vthread") T.launch_thread(vthread, 4) - B = T.alloc_buffer((4,), "int32", scope="shared") + B = T.alloc_tensor((4,), "int32", scope="shared") B[T.ramp(0, 1, 4)] = T.broadcast(vthread, 4) @Ts.prim_func def expected_func(): - B = T.alloc_buffer((16,), "int32", scope="shared") + B = T.alloc_tensor((16,), "int32", scope="shared") # The indices for B should each be a single Ramp node, and # should not be the sum of a Ramp and Broadcast node. B[T.ramp(T.Mul(0, 4), 1, 4)] = T.broadcast(0, 4) @@ -203,7 +203,7 @@ def test_vthread_vectorized(): def before_func(): vthread = T.env_thread("vthread") T.launch_thread(vthread, 4) - B = T.alloc_buffer((4,), "int32", scope="shared") + B = T.alloc_tensor((4,), "int32", scope="shared") B[T.ramp(0, 1, 4)] = T.broadcast(vthread, 4) before_mod = tvm.IRModule.from_expr(before_func.with_attr("global_symbol", "main")) @@ -216,7 +216,7 @@ def before_func(): def visitor(op): nonlocal allocate_node - if _is_buffer_binding(op, "tirx.alloc_buffer") and "shared" in str(op.var.data.ty): + if _is_buffer_binding(op, "tirx.alloc_tensor") and "shared" in str(op.var.data.ty): allocate_node = op tvm_ffi.structural_walk(after_func.body, visitor) @@ -230,7 +230,7 @@ def test_vthread_rewrites_masked_accesses(): def before_func(): vthread = T.env_thread("vthread") T.launch_thread(vthread, 2) - B = T.alloc_buffer((4,), "float32", scope="shared") + B = T.alloc_tensor((4,), "float32", scope="shared") mask = T.meta_var(T.Broadcast(T.bool(True), 4)) loaded = T.meta_var(T.masked_load("float32x4", B, T.Ramp(0, 1, 4), mask)) value = T.meta_var(loaded + T.Broadcast(T.Cast("float32", vthread), 4)) 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 b0ee057861dd..c15e9c1ce960 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 @@ -26,7 +26,7 @@ def test_lift_tx_beyond_local(): n = T.dynamic("n", "int32") @Ts.prim_func - def before(A: T.Buffer((32, 1, 128)), B: T.Buffer((32, n, 128)), C: T.Buffer((32, 1, n))): + def before(A: T.Tensor((32, 1, 128)), B: T.Tensor((32, n, 128)), C: T.Tensor((32, 1, n))): for ax0_ax1_fused in T.thread_binding(n * 32, thread="blockIdx.x"): with Ts.sblock(""): @@ -80,7 +80,7 @@ def before(A: T.Buffer((32, 1, 128)), B: T.Buffer((32, n, 128)), C: T.Buffer((32 n = T.dynamic("n", "int32") @Ts.prim_func - def expected(A: T.Buffer((32, 1, 128), "float32"), B: T.Buffer((32, n, 128)), C: T.Buffer((32, 1, n))): + def expected(A: T.Tensor((32, 1, 128), "float32"), B: T.Tensor((32, n, 128)), C: T.Tensor((32, 1, n))): # with Ts.sblock("root"): for blockIdx_x in T.thread_binding(n * 32, thread="blockIdx.x"): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_loop_partition.py b/tests/python/s_tir/transform/test_s_tir_transform_loop_partition.py index 97f739d6c5bc..25d48f99a8b3 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_loop_partition.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_loop_partition.py @@ -119,8 +119,8 @@ def func(m: T.int64, n: T.int64): def test_oneD_pool(): @Ts.prim_func def func(m: T.int64, data: T.handle("float32"), out: T.handle("float32")): - data_ptr = T.decl_buffer((16,), "float32", data=data) - out_ptr = T.decl_buffer((16,), "float32", data=out) + data_ptr = T.decl_tensor((16,), "float32", data=data) + out_ptr = T.decl_tensor((16,), "float32", data=out) for ow in range(16): for kw in range(3): if T.likely(ow > 0): @@ -151,7 +151,7 @@ def test_cce_loop_1(): m = 514 @Ts.prim_func - def func(A: T.Buffer((n * m,), "float16"), B: T.Buffer((n * m,), "float16")): + def func(A: T.Tensor((n * m,), "float16"), B: T.Tensor((n * m,), "float16")): for i in range(11): for j in range(160): if T.likely(i * 160 + j < 1600): @@ -211,7 +211,7 @@ def func(): @Ts.prim_func def partitioned_concat( - A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32"), C: T.Buffer((32,), "float32") + A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32"), C: T.Tensor((32,), "float32") ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i in T.serial(0, 16): @@ -234,15 +234,15 @@ def partition_from_scheduled_tir(prim_func, pass_cfg, do_flatten=True): @Ts.prim_func def partitioned_concat_3( - placeholder: T.Buffer((1, 64, 28, 28), "int8"), - placeholder_1: T.Buffer((1, 32, 28, 28), "int8"), - placeholder_2: T.Buffer((1, 32, 28, 28), "int8"), - T_concat: T.Buffer((1, 128, 28, 28), "int8"), + placeholder: T.Tensor((1, 64, 28, 28), "int8"), + placeholder_1: T.Tensor((1, 32, 28, 28), "int8"), + placeholder_2: T.Tensor((1, 32, 28, 28), "int8"), + T_concat: T.Tensor((1, 128, 28, 28), "int8"), ) -> None: - placeholder_flat = T.decl_buffer([50176], "int8", data=placeholder.data) - placeholder_1_flat = T.decl_buffer([25088], "int8", data=placeholder_1.data) - placeholder_2_flat = T.decl_buffer([25088], "int8", data=placeholder_2.data) - T_concat_flat = T.decl_buffer([100352], "int8", data=T_concat.data) + placeholder_flat = T.decl_tensor([50176], "int8", data=placeholder.data) + placeholder_1_flat = T.decl_tensor([25088], "int8", data=placeholder_1.data) + placeholder_2_flat = T.decl_tensor([25088], "int8", data=placeholder_2.data) + T_concat_flat = T.decl_tensor([100352], "int8", data=T_concat.data) for i1, i2, i3 in T.grid(64, 28, 28): T_concat_flat[i1 * 784 + i2 * 28 + i3] = placeholder_flat[i1 * 784 + i2 * 28 + i3] for i1, i2, i3 in T.grid(32, 28, 28): @@ -253,15 +253,15 @@ def partitioned_concat_3( @Ts.prim_func def concat_func_3( - placeholder: T.Buffer((1, 64, 28, 28), "int8"), - placeholder_1: T.Buffer((1, 32, 28, 28), "int8"), - placeholder_2: T.Buffer((1, 32, 28, 28), "int8"), - T_concat: T.Buffer((1, 128, 28, 28), "int8"), + placeholder: T.Tensor((1, 64, 28, 28), "int8"), + placeholder_1: T.Tensor((1, 32, 28, 28), "int8"), + placeholder_2: T.Tensor((1, 32, 28, 28), "int8"), + T_concat: T.Tensor((1, 128, 28, 28), "int8"), ) -> None: - placeholder_flat = T.decl_buffer([50176], "int8", data=placeholder.data) - placeholder_1_flat = T.decl_buffer([25088], "int8", data=placeholder_1.data) - placeholder_2_flat = T.decl_buffer([25088], "int8", data=placeholder_2.data) - T_concat_flat = T.decl_buffer([100352], "int8", data=T_concat.data) + placeholder_flat = T.decl_tensor([50176], "int8", data=placeholder.data) + placeholder_1_flat = T.decl_tensor([25088], "int8", data=placeholder_1.data) + placeholder_2_flat = T.decl_tensor([25088], "int8", data=placeholder_2.data) + T_concat_flat = T.decl_tensor([100352], "int8", data=T_concat.data) for i1 in T.serial(128, annotations={"pragma_loop_partition_hint": 1}): for i2, i3 in T.grid(28, 28): if 96 <= i1: @@ -288,10 +288,10 @@ def test_condition_mutually_exclusive(): def test_loop_partition_unroll_hint(): @Ts.prim_func def main( - A_arg: T.Buffer((1, 3, 224, 224), "int8"), B_arg: T.Buffer((1, 224, 7, 16), "int8") + A_arg: T.Tensor((1, 3, 224, 224), "int8"), B_arg: T.Tensor((1, 224, 7, 16), "int8") ) -> None: - A = T.decl_buffer(150528, "int8", data=A_arg.data) - B = T.decl_buffer(25088, "int8", data=B_arg.data) + A = T.decl_tensor(150528, "int8", data=A_arg.data) + B = T.decl_tensor(25088, "int8", data=B_arg.data) for ax0 in T.serial( 112, annotations={"pragma_loop_partition_hint": True}, @@ -302,10 +302,10 @@ def main( @Ts.prim_func def partitioned_main( - A_arg: T.Buffer((1, 3, 224, 224), "int8"), B_arg: T.Buffer((1, 224, 7, 16), "int8") + A_arg: T.Tensor((1, 3, 224, 224), "int8"), B_arg: T.Tensor((1, 224, 7, 16), "int8") ) -> None: - A = T.decl_buffer(150528, dtype="int8", data=A_arg.data) - B = T.decl_buffer(25088, dtype="int8", data=B_arg.data) + A = T.decl_tensor(150528, dtype="int8", data=A_arg.data) + B = T.decl_tensor(25088, dtype="int8", data=B_arg.data) # body for ax1, ax2, ax3 in T.grid(224, 7, 16): if 3 <= ax2 and ax3 < 3: @@ -338,10 +338,10 @@ def partitioned_main( def test_loop_partition_recursive_unroll_hint(): @Ts.prim_func def main(): - placeholder_0_dm = T.decl_buffer([1, 32, 32, 16], dtype="int8") + placeholder_0_dm = T.decl_tensor([1, 32, 32, 16], dtype="int8") for i3_0 in T.serial(5, annotations={"pragma_loop_partition_hint": 1}): for i2_0 in T.serial(2, annotations={"pragma_loop_partition_hint": 1}): - pad_temp = T.decl_buffer([1, 16, 16, 16], dtype="int8") + pad_temp = T.decl_tensor([1, 16, 16, 16], dtype="int8") for ax0, ax1, ax2 in T.grid(16, 16, 16): if ( 6 <= i2_0 * 4 + ax0 @@ -363,17 +363,17 @@ def main(): @Ts.prim_func def partitioned_main(): - placeholder_0_dm = T.decl_buffer((16384,), "int8") + placeholder_0_dm = T.decl_tensor((16384,), "int8") for i3_0 in T.unroll(2): for i2_0 in T.unroll(2): - pad_temp = T.decl_buffer((4096,), "int8") + pad_temp = T.decl_tensor((4096,), "int8") for ax0, ax1, ax2 in T.grid(16, 16, 16): if 6 <= i2_0 * 4 + ax0 and 6 <= i3_0 * 4 + ax1: pad_temp[ax0 * 256 + ax1 * 16 + ax2] = placeholder_0_dm[ i2_0 * 2048 + ax0 * 512 + i3_0 * 64 + ax1 * 16 + ax2 ] for i2_0 in T.unroll(2): - pad_temp = T.decl_buffer((4096,), "int8") + pad_temp = T.decl_tensor((4096,), "int8") for ax0, ax1, ax2 in T.grid(16, 16, 16): if 6 <= i2_0 * 4 + ax0: pad_temp[ax0 * 256 + ax1 * 16 + ax2] = placeholder_0_dm[ @@ -381,7 +381,7 @@ def partitioned_main(): ] for i3_0 in T.unroll(2): for i2_0 in T.unroll(2): - pad_temp = T.decl_buffer((4096,), "int8") + pad_temp = T.decl_tensor((4096,), "int8") for ax0, ax1, ax2 in T.grid(16, 16, 16): if 6 <= i2_0 * 4 + ax0 and i3_0 * 4 + ax1 < 14: pad_temp[ax0 * 256 + ax1 * 16 + ax2] = placeholder_0_dm[ @@ -402,7 +402,7 @@ def partitioned_main(): def test_loop_partition_keep_loop_annotations(): @Ts.prim_func - def before(A: T.Buffer(160, "int32"), B: T.Buffer(160, "int32")) -> None: + def before(A: T.Tensor(160, "int32"), B: T.Tensor(160, "int32")) -> None: for i in T.serial( 160, annotations={"pragma_loop_partition_hint": True, "key": "value"}, @@ -415,9 +415,9 @@ def before(A: T.Buffer(160, "int32"), B: T.Buffer(160, "int32")) -> None: B[i] = A[i] + 3 @Ts.prim_func - def after(A: T.Buffer(160, "int32"), B: T.Buffer(160, "int32")) -> None: - A_1 = T.decl_buffer((160,), "int32", data=A.data) - B_1 = T.decl_buffer((160,), "int32", data=B.data) + def after(A: T.Tensor(160, "int32"), B: T.Tensor(160, "int32")) -> None: + A_1 = T.decl_tensor((160,), "int32", data=A.data) + B_1 = T.decl_tensor((160,), "int32", data=B.data) for i in T.serial(10, annotations={"key": "value"}): B_1[i] = A_1[i] + 1 for i in T.serial(140, annotations={"key": "value"}): @@ -439,10 +439,10 @@ def after(A: T.Buffer(160, "int32"), B: T.Buffer(160, "int32")) -> None: def test_loop_partition_with_unit_loop_in_condition(): @Ts.prim_func def before( - placeholder: T.Buffer((50176,), "int8"), - placeholder_1: T.Buffer((25088,), "int8"), - placeholder_2: T.Buffer((25088,), "int8"), - T_concat: T.Buffer((100352,), "int8"), + placeholder: T.Tensor((50176,), "int8"), + placeholder_1: T.Tensor((25088,), "int8"), + placeholder_2: T.Tensor((25088,), "int8"), + T_concat: T.Tensor((100352,), "int8"), ) -> None: for k in range(1, annotations={"preserve_unit_loop": True}): for i1 in T.serial(128, annotations={"pragma_loop_partition_hint": 1}): @@ -460,15 +460,15 @@ def before( @Ts.prim_func def after( - placeholder: T.Buffer(50176, "int8"), - placeholder_1: T.Buffer(25088, "int8"), - placeholder_2: T.Buffer(25088, "int8"), - T_concat: T.Buffer(100352, "int8"), + placeholder: T.Tensor(50176, "int8"), + placeholder_1: T.Tensor(25088, "int8"), + placeholder_2: T.Tensor(25088, "int8"), + T_concat: T.Tensor(100352, "int8"), ) -> None: - placeholder_3 = T.decl_buffer((50176,), "int8", data=placeholder.data) - placeholder_1_1 = T.decl_buffer((25088,), "int8", data=placeholder_1.data) - placeholder_2_1 = T.decl_buffer((25088,), "int8", data=placeholder_2.data) - T_concat_1 = T.decl_buffer((100352,), "int8", data=T_concat.data) + placeholder_3 = T.decl_tensor((50176,), "int8", data=placeholder.data) + placeholder_1_1 = T.decl_tensor((25088,), "int8", data=placeholder_1.data) + placeholder_2_1 = T.decl_tensor((25088,), "int8", data=placeholder_2.data) + T_concat_1 = T.decl_tensor((100352,), "int8", data=T_concat.data) for k in T.serial(1, annotations={"preserve_unit_loop": True}): for i1, i2, i3 in T.grid(64, 28, 28): T_concat_1[i1 * 784 + i2 * 28 + i3] = placeholder_3[i1 * 784 + i2 * 28 + i3] @@ -492,10 +492,10 @@ def after( @Ts.prim_func def concat_func_single_point( - placeholder: T.Buffer((28, 64), "int8"), - placeholder_1: T.Buffer((28, 1), "int8"), - placeholder_2: T.Buffer((28, 63), "int8"), - T_concat: T.Buffer((28, 128), "int8"), + placeholder: T.Tensor((28, 64), "int8"), + placeholder_1: T.Tensor((28, 1), "int8"), + placeholder_2: T.Tensor((28, 63), "int8"), + T_concat: T.Tensor((28, 128), "int8"), ) -> None: for i0 in range(28): for i1 in T.serial(128, annotations={"pragma_loop_partition_hint": 1}): @@ -509,15 +509,15 @@ def concat_func_single_point( @Ts.prim_func def expected_partitioned_concat_single_point( - placeholder: T.Buffer((28, 64), "int8"), - placeholder_1: T.Buffer((28, 1), "int8"), - placeholder_2: T.Buffer((28, 63), "int8"), - T_concat: T.Buffer((28, 128), "int8"), + placeholder: T.Tensor((28, 64), "int8"), + placeholder_1: T.Tensor((28, 1), "int8"), + placeholder_2: T.Tensor((28, 63), "int8"), + T_concat: T.Tensor((28, 128), "int8"), ): - placeholder_3 = T.decl_buffer((1792,), "int8", data=placeholder.data) - placeholder_1_1 = T.decl_buffer((28,), "int8", data=placeholder_1.data) - placeholder_2_1 = T.decl_buffer((1764,), "int8", data=placeholder_2.data) - T_concat_1 = T.decl_buffer((3584,), "int8", data=T_concat.data) + placeholder_3 = T.decl_tensor((1792,), "int8", data=placeholder.data) + placeholder_1_1 = T.decl_tensor((28,), "int8", data=placeholder_1.data) + placeholder_2_1 = T.decl_tensor((1764,), "int8", data=placeholder_2.data) + T_concat_1 = T.decl_tensor((3584,), "int8", data=T_concat.data) for i0 in range(28): for i1 in range(63): T_concat_1[i0 * 128 + i1] = placeholder_2_1[i0 * 63 + i1] @@ -528,10 +528,10 @@ def expected_partitioned_concat_single_point( @Ts.prim_func def concat_func_start_point_equality( - placeholder: T.Buffer((28, 64), "int8"), - placeholder_1: T.Buffer((28, 1), "int8"), - placeholder_2: T.Buffer((28, 63), "int8"), - T_concat: T.Buffer((28, 128), "int8"), + placeholder: T.Tensor((28, 64), "int8"), + placeholder_1: T.Tensor((28, 1), "int8"), + placeholder_2: T.Tensor((28, 63), "int8"), + T_concat: T.Tensor((28, 128), "int8"), ) -> None: for i0 in range(28): for i1 in range(128, annotations={"pragma_loop_partition_hint": 1}): @@ -548,15 +548,15 @@ def concat_func_start_point_equality( @Ts.prim_func def concat_func_start_point_equality_expected( - placeholder: T.Buffer((28, 64), "int8"), - placeholder_1: T.Buffer((28, 1), "int8"), - placeholder_2: T.Buffer((28, 63), "int8"), - T_concat: T.Buffer((28, 128), "int8"), + placeholder: T.Tensor((28, 64), "int8"), + placeholder_1: T.Tensor((28, 1), "int8"), + placeholder_2: T.Tensor((28, 63), "int8"), + T_concat: T.Tensor((28, 128), "int8"), ): - placeholder_3 = T.decl_buffer((1792,), "int8", data=placeholder.data) - placeholder_1_1 = T.decl_buffer((28,), "int8", data=placeholder_1.data) - placeholder_2_1 = T.decl_buffer((1764,), "int8", data=placeholder_2.data) - T_concat_1 = T.decl_buffer((3584,), "int8", data=T_concat.data) + placeholder_3 = T.decl_tensor((1792,), "int8", data=placeholder.data) + placeholder_1_1 = T.decl_tensor((28,), "int8", data=placeholder_1.data) + placeholder_2_1 = T.decl_tensor((1764,), "int8", data=placeholder_2.data) + T_concat_1 = T.decl_tensor((3584,), "int8", data=T_concat.data) for i0 in range(28): T_concat_1[i0 * 128] = placeholder_1_1[i0] for i1 in range(63): @@ -567,10 +567,10 @@ def concat_func_start_point_equality_expected( @Ts.prim_func def concat_func_end_point_equality( - placeholder: T.Buffer((28, 64), "int8"), - placeholder_1: T.Buffer((28, 1), "int8"), - placeholder_2: T.Buffer((28, 63), "int8"), - T_concat: T.Buffer((28, 128), "int8"), + placeholder: T.Tensor((28, 64), "int8"), + placeholder_1: T.Tensor((28, 1), "int8"), + placeholder_2: T.Tensor((28, 63), "int8"), + T_concat: T.Tensor((28, 128), "int8"), ) -> None: for i0 in range(28): for i1 in range(128, annotations={"pragma_loop_partition_hint": 1}): @@ -587,15 +587,15 @@ def concat_func_end_point_equality( @Ts.prim_func def concat_func_end_point_equality_expected( - placeholder: T.Buffer((28, 64), "int8"), - placeholder_1: T.Buffer((28, 1), "int8"), - placeholder_2: T.Buffer((28, 63), "int8"), - T_concat: T.Buffer((28, 128), "int8"), + placeholder: T.Tensor((28, 64), "int8"), + placeholder_1: T.Tensor((28, 1), "int8"), + placeholder_2: T.Tensor((28, 63), "int8"), + T_concat: T.Tensor((28, 128), "int8"), ): - placeholder_3 = T.decl_buffer((1792,), "int8", data=placeholder.data) - placeholder_1_1 = T.decl_buffer((28,), "int8", data=placeholder_1.data) - placeholder_2_1 = T.decl_buffer((1764,), "int8", data=placeholder_2.data) - T_concat_1 = T.decl_buffer((3584,), "int8", data=T_concat.data) + placeholder_3 = T.decl_tensor((1792,), "int8", data=placeholder.data) + placeholder_1_1 = T.decl_tensor((28,), "int8", data=placeholder_1.data) + placeholder_2_1 = T.decl_tensor((1764,), "int8", data=placeholder_2.data) + T_concat_1 = T.decl_tensor((3584,), "int8", data=T_concat.data) for i0 in range(28): for i1 in range(64): T_concat_1[i0 * 128 + i1] = placeholder_2_1[i0 * 63 + i1] @@ -606,10 +606,10 @@ def concat_func_end_point_equality_expected( @Ts.prim_func def concat_func_edge_equalities( - placeholder: T.Buffer((28, 64), "int8"), - placeholder_1: T.Buffer((28, 1), "int8"), - placeholder_2: T.Buffer((28, 1), "int8"), - T_concat: T.Buffer((28, 66), "int8"), + placeholder: T.Tensor((28, 64), "int8"), + placeholder_1: T.Tensor((28, 1), "int8"), + placeholder_2: T.Tensor((28, 1), "int8"), + T_concat: T.Tensor((28, 66), "int8"), ) -> None: for i0 in range(28): for i1 in range( @@ -628,15 +628,15 @@ def concat_func_edge_equalities( @Ts.prim_func def concat_func_edge_equalities_expected( - placeholder: T.Buffer((28, 64), "int8"), - placeholder_1: T.Buffer((28, 1), "int8"), - placeholder_2: T.Buffer((28, 1), "int8"), - T_concat: T.Buffer((28, 66), "int8"), + placeholder: T.Tensor((28, 64), "int8"), + placeholder_1: T.Tensor((28, 1), "int8"), + placeholder_2: T.Tensor((28, 1), "int8"), + T_concat: T.Tensor((28, 66), "int8"), ): - placeholder_3 = T.decl_buffer((1792,), "int8", data=placeholder.data) - placeholder_1_1 = T.decl_buffer((28,), "int8", data=placeholder_1.data) - placeholder_2_1 = T.decl_buffer((28,), "int8", data=placeholder_2.data) - T_concat_1 = T.decl_buffer((1848,), "int8", data=T_concat.data) + placeholder_3 = T.decl_tensor((1792,), "int8", data=placeholder.data) + placeholder_1_1 = T.decl_tensor((28,), "int8", data=placeholder_1.data) + placeholder_2_1 = T.decl_tensor((28,), "int8", data=placeholder_2.data) + T_concat_1 = T.decl_tensor((1848,), "int8", data=T_concat.data) for i0 in range(28): T_concat_1[i0 * 66] = placeholder_2_1[i0] for i1 in range(64): @@ -646,12 +646,12 @@ def concat_func_edge_equalities_expected( @Ts.prim_func def concat_five_buffers_with_equalities( - buffer_a: T.Buffer((28, 1), "int8"), # Used for i1 == 0 - buffer_b: T.Buffer((28, 63), "int8"), # Fills i1 from 1 to 63 - buffer_c: T.Buffer((28, 1), "int8"), # Used for i1 == 64 - buffer_d: T.Buffer((28, 63), "int8"), # Fills i1 from 65 to 128 - buffer_e: T.Buffer((28, 1), "int8"), # Used for i1 == 129 - T_concat: T.Buffer((28, 129), "int8"), + buffer_a: T.Tensor((28, 1), "int8"), # Used for i1 == 0 + buffer_b: T.Tensor((28, 63), "int8"), # Fills i1 from 1 to 63 + buffer_c: T.Tensor((28, 1), "int8"), # Used for i1 == 64 + buffer_d: T.Tensor((28, 63), "int8"), # Fills i1 from 65 to 128 + buffer_e: T.Tensor((28, 1), "int8"), # Used for i1 == 129 + T_concat: T.Tensor((28, 129), "int8"), ) -> None: for i0 in range(28): for i1 in range(130, annotations={"pragma_loop_partition_hint": 1}): @@ -669,19 +669,19 @@ def concat_five_buffers_with_equalities( @Ts.prim_func def concat_five_buffers_with_equalities_expected( - buffer_a: T.Buffer((28, 1), "int8"), # Used for i1 == 0 - buffer_b: T.Buffer((28, 63), "int8"), # Fills i1 from 1 to 63 - buffer_c: T.Buffer((28, 1), "int8"), # Used for i1 == 64 - buffer_d: T.Buffer((28, 63), "int8"), # Fills i1 from 65 to 128 - buffer_e: T.Buffer((28, 1), "int8"), # Used for i1 == 129 - T_concat: T.Buffer((28, 129), "int8"), + buffer_a: T.Tensor((28, 1), "int8"), # Used for i1 == 0 + buffer_b: T.Tensor((28, 63), "int8"), # Fills i1 from 1 to 63 + buffer_c: T.Tensor((28, 1), "int8"), # Used for i1 == 64 + buffer_d: T.Tensor((28, 63), "int8"), # Fills i1 from 65 to 128 + buffer_e: T.Tensor((28, 1), "int8"), # Used for i1 == 129 + T_concat: T.Tensor((28, 129), "int8"), ): - buffer_a_1 = T.decl_buffer((28,), "int8", data=buffer_a.data) - buffer_b_1 = T.decl_buffer((1764,), "int8", data=buffer_b.data) - buffer_c_1 = T.decl_buffer((28,), "int8", data=buffer_c.data) - buffer_d_1 = T.decl_buffer((1764,), "int8", data=buffer_d.data) - buffer_e_1 = T.decl_buffer((28,), "int8", data=buffer_e.data) - T_concat_1 = T.decl_buffer((3612,), "int8", data=T_concat.data) + buffer_a_1 = T.decl_tensor((28,), "int8", data=buffer_a.data) + buffer_b_1 = T.decl_tensor((1764,), "int8", data=buffer_b.data) + buffer_c_1 = T.decl_tensor((28,), "int8", data=buffer_c.data) + buffer_d_1 = T.decl_tensor((1764,), "int8", data=buffer_d.data) + buffer_e_1 = T.decl_tensor((28,), "int8", data=buffer_e.data) + T_concat_1 = T.decl_tensor((3612,), "int8", data=T_concat.data) for i0 in range(28): T_concat_1[i0 * 129] = buffer_a_1[i0] for i1 in range(63): @@ -693,7 +693,7 @@ def concat_five_buffers_with_equalities_expected( @Ts.prim_func -def nested_partition_with_single_points(A: T.Buffer((25,), "int32")): +def nested_partition_with_single_points(A: T.Tensor((25,), "int32")): for i in T.serial(5, annotations={"pragma_loop_partition_hint": 1}): if i == 1: for j in T.serial(5, annotations={"pragma_loop_partition_hint": 1}): @@ -706,8 +706,8 @@ def nested_partition_with_single_points(A: T.Buffer((25,), "int32")): @Ts.prim_func -def nested_partition_with_single_points_expected(A: T.Buffer((25,), "int32")): - A_1 = T.decl_buffer((25,), "int32", data=A.data) +def nested_partition_with_single_points_expected(A: T.Tensor((25,), "int32")): + A_1 = T.decl_tensor((25,), "int32", data=A.data) for j in range(2): A_1[j + 3] = j + 3 for j in range(2): @@ -744,7 +744,7 @@ def test_single_point_partition(origin, expected): def test_equation_on_floordiv(): @Ts.prim_func - def before(A: T.Buffer((2, 2, 20), "int32")): + def before(A: T.Tensor((2, 2, 20), "int32")): for i in T.serial(5, annotations={"pragma_loop_partition_hint": 1}): if i == 1: for vv in T.vectorized(640, annotations={"pragma_loop_partition_hint": 1}): @@ -752,7 +752,7 @@ def before(A: T.Buffer((2, 2, 20), "int32")): A[i - 1, i * 2 + vv // 320 - 3, vv % 320 // 16] = 1 @Ts.prim_func - def expected(A: T.Buffer((2, 2, 20), "int32")): + def expected(A: T.Tensor((2, 2, 20), "int32")): for vv in T.vectorized(320): A[0, 0, vv // 16] = 1 @@ -767,9 +767,9 @@ def test_ignore_loop_partition_hint(): """Skip unroll body and prologue for pipeline case""" @Ts.prim_func - def before(A: T.Buffer((10), "float32"), D: T.Buffer((10), "float32")): - B = T.decl_buffer([2], "float32") - C = T.decl_buffer([2], "float32") + def before(A: T.Tensor((10), "float32"), D: T.Tensor((10), "float32")): + B = T.decl_tensor([2], "float32") + C = T.decl_tensor([2], "float32") for i in T.serial(12, annotations={"pragma_loop_partition_hint": 1}): if T.ignore_loop_partition(i < 10): B[i % 2] = A[i] + 1.0 @@ -779,9 +779,9 @@ def before(A: T.Buffer((10), "float32"), D: T.Buffer((10), "float32")): D[i - 2] = C[i % 2] + 3.0 @Ts.prim_func - def expected(A: T.Buffer((10), "float32"), D: T.Buffer((10), "float32")): - B = T.decl_buffer([2], "float32") - C = T.decl_buffer([2], "float32") + def expected(A: T.Tensor((10), "float32"), D: T.Tensor((10), "float32")): + B = T.decl_tensor([2], "float32") + C = T.decl_tensor([2], "float32") for i in range(2): B[i] = A[i] + 1.0 if i == 1: 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 c375eb7c1ddc..5543f3c7d9ac 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 @@ -45,7 +45,7 @@ def _check_fail(original): @Ts.prim_func def loop_split( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: for i, ko in T.grid(128, 4): for ki in T.thread_binding(0, 32, thread="threadIdx.x"): @@ -61,7 +61,7 @@ def loop_split( @Ts.prim_func def lowered_loop_split( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: reduce_temp0 = Ts.sblock_alloc_buffer([1], dtype="float32", strides=[1], scope="local") normal_reduce_temp0 = Ts.sblock_alloc_buffer([1], dtype="float32", strides=[1], scope="local") @@ -105,7 +105,7 @@ def lowered_loop_split( @Ts.prim_func def no_normal_reduction( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: for i in T.serial(0, 128): for k in T.thread_binding(0, 128, thread="threadIdx.x"): @@ -121,7 +121,7 @@ def no_normal_reduction( # complains that k is defined outside of a block @Ts.prim_func(check_well_formed=False) def lowered_no_normal_reduction( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: reduce_temp0 = Ts.sblock_alloc_buffer([1], dtype="float32", strides=[1], scope="local") for i in T.serial(0, 128): @@ -146,7 +146,7 @@ def lowered_no_normal_reduction( @Ts.prim_func def two_bound_loops( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: for i in T.serial(0, 128): for ko in T.thread_binding(0, 4, thread="threadIdx.x"): @@ -164,7 +164,7 @@ def two_bound_loops( # complains that ko is defined outside of a block @Ts.prim_func(check_well_formed=False) def lowered_two_bound_loops( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: reduce_temp0 = Ts.sblock_alloc_buffer([1], dtype="float32", strides=[1], scope="local") for i in T.serial(0, 128): @@ -195,7 +195,7 @@ def lowered_two_bound_loops( @Ts.prim_func def multiple_blocks_under_reduction_loop( - A: T.Buffer([16, 16, 16], dtype="float32"), B: T.Buffer([16], dtype="float32") + A: T.Tensor([16, 16, 16], dtype="float32"), B: T.Tensor([16], dtype="float32") ) -> None: B_rf_local = Ts.sblock_alloc_buffer([16, 16], dtype="float32", scope="local") for i in T.thread_binding(0, 16, thread="blockIdx.x"): @@ -222,7 +222,7 @@ def multiple_blocks_under_reduction_loop( @Ts.prim_func def lowered_multiple_blocks_under_reduction_loop( - A: T.Buffer([16, 16, 16], dtype="float32"), B: T.Buffer([16], dtype="float32") + A: T.Tensor([16, 16, 16], dtype="float32"), B: T.Tensor([16], dtype="float32") ) -> None: B_rf_local = Ts.sblock_alloc_buffer([16, 16], dtype="float32", scope="local") reduce_temp0 = Ts.sblock_alloc_buffer([1], dtype="float32", strides=[1], scope="local") @@ -276,7 +276,7 @@ def lowered_multiple_blocks_under_reduction_loop( @Ts.prim_func def with_block_predicate( - A: T.Buffer([128, 120], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 120], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: for i, ko in T.grid(128, 4): for ki in T.thread_binding(0, 32, thread="threadIdx.x"): @@ -293,7 +293,7 @@ def with_block_predicate( @Ts.prim_func def lowered_with_block_predicate( - A: T.Buffer([128, 120], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 120], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: reduce_temp0 = Ts.sblock_alloc_buffer([1], dtype="float32", strides=[1], scope="local") normal_reduce_temp0 = Ts.sblock_alloc_buffer([1], dtype="float32", strides=[1], scope="local") @@ -338,7 +338,7 @@ def lowered_with_block_predicate( @Ts.prim_func def single_reduction_loop_with_block_predicate( - A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32") + A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32") ) -> None: T_softmax_maxelem_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") T_softmax_expsum_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") @@ -388,7 +388,7 @@ def single_reduction_loop_with_block_predicate( @Ts.prim_func def lowered_single_reduction_loop_with_block_predicate( - A: T.Buffer((256, 256), "float32"), T_softmax_norm: T.Buffer((256, 256), "float32") + A: T.Tensor((256, 256), "float32"), T_softmax_norm: T.Tensor((256, 256), "float32") ) -> None: T_softmax_maxelem_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") T_softmax_expsum_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") @@ -494,9 +494,9 @@ def lowered_single_reduction_loop_with_block_predicate( @Ts.prim_func def spatial_reduction_with_shared_prefetch( - A: T.Buffer((128, 150528), "float32"), - B: T.Buffer((128, 150528), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 150528), "float32"), + B: T.Tensor((128, 150528), "float32"), + C: T.Tensor((128, 128), "float32"), ): C_local = Ts.sblock_alloc_buffer((128, 128), scope="local") A_shared = Ts.sblock_alloc_buffer((128, 150528), scope="shared") @@ -589,9 +589,9 @@ def spatial_reduction_with_shared_prefetch( @Ts.prim_func def lowered_spatial_reduction_with_shared_prefetch( - A: T.Buffer((128, 150528), "float32"), - B: T.Buffer((128, 150528), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 150528), "float32"), + B: T.Tensor((128, 150528), "float32"), + C: T.Tensor((128, 128), "float32"), ): C_local = Ts.sblock_alloc_buffer((128, 128), scope="local") A_shared = Ts.sblock_alloc_buffer((128, 150528), scope="shared") @@ -712,7 +712,7 @@ def lowered_spatial_reduction_with_shared_prefetch( @Ts.prim_func -def spatial_reduction_loop_predicate(A: T.Buffer((2, 32), "float32"), B: T.Buffer((2,), "float32")): +def spatial_reduction_loop_predicate(A: T.Tensor((2, 32), "float32"), B: T.Tensor((2,), "float32")): for i_0 in range(1): for i_1 in T.thread_binding(16, thread="threadIdx.y"): for k_0 in range(1): @@ -733,7 +733,7 @@ def spatial_reduction_loop_predicate(A: T.Buffer((2, 32), "float32"), B: T.Buffe @Ts.prim_func def lowered_reduction_spatial_loop_predicate( - A: T.Buffer((2, 32), "float32"), B: T.Buffer((2,), "float32") + A: T.Tensor((2, 32), "float32"), B: T.Tensor((2,), "float32") ): cross_thread_B = Ts.sblock_alloc_buffer((1,), strides=(1,), scope="local") in_thread_B = Ts.sblock_alloc_buffer((1,), strides=(1,), scope="local") @@ -773,9 +773,9 @@ def lowered_reduction_spatial_loop_predicate( @Ts.prim_func def single_reduction_loop_with_tensorize( - input_A: T.Buffer((1, 64, 7, 7, 32), "uint8"), - input_B: T.Buffer((16, 64, 1, 1, 8, 32, 4), "int8"), - output: T.Buffer((1, 16, 7, 7, 32), "int32"), + input_A: T.Tensor((1, 64, 7, 7, 32), "uint8"), + input_B: T.Tensor((16, 64, 1, 1, 8, 32, 4), "int8"), + output: T.Tensor((1, 16, 7, 7, 32), "int32"), ) -> None: # body # with Ts.sblock("root") @@ -838,9 +838,9 @@ def single_reduction_loop_with_tensorize( @Ts.prim_func def nested_reduction_loop_with_inner_match_buffers( - in0: T.Buffer((4, 16), "int8"), - in1: T.Buffer((4, 16), "int8"), - out: T.Buffer((4, 4), "int32"), + in0: T.Tensor((4, 16), "int8"), + in1: T.Tensor((4, 16), "int8"), + out: T.Tensor((4, 4), "int32"), ) -> None: # body # with Ts.sblock("root") @@ -888,7 +888,7 @@ def nested_reduction_loop_with_inner_match_buffers( @Ts.prim_func def reducer_max( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: for i in T.serial(0, 128): for k in T.thread_binding(0, 128, thread="threadIdx.x"): @@ -904,7 +904,7 @@ def reducer_max( # complains that k is defined outside of a block @Ts.prim_func(check_well_formed=False) def lowered_reducer_max( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: reduce_temp0 = Ts.sblock_alloc_buffer([1], dtype="float32", strides=[1], scope="local") for i in T.serial(0, 128): @@ -928,7 +928,7 @@ def lowered_reducer_max( @Ts.prim_func -def zero_rank_buffer(A: T.Buffer([128], dtype="float32"), B: T.Buffer([], dtype="float32")) -> None: +def zero_rank_buffer(A: T.Tensor([128], dtype="float32"), B: T.Tensor([], dtype="float32")) -> None: for k in T.thread_binding(0, 128, thread="threadIdx.x"): with Ts.sblock("B"): vk = Ts.axis.reduce(128, k) @@ -942,7 +942,7 @@ def zero_rank_buffer(A: T.Buffer([128], dtype="float32"), B: T.Buffer([], dtype= # complains that k is defined outside of a block @Ts.prim_func(check_well_formed=False) def lowered_zero_rank_buffer( - A: T.Buffer([128], dtype="float32"), B: T.Buffer([], dtype="float32") + A: T.Tensor([128], dtype="float32"), B: T.Tensor([], dtype="float32") ) -> None: reduce_temp0 = Ts.sblock_alloc_buffer([1], dtype="float32", strides=[1], scope="local") for k in T.thread_binding(0, 128, thread="threadIdx.x"): @@ -965,7 +965,7 @@ def lowered_zero_rank_buffer( @Ts.prim_func def multiple_bufferstore( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: C = Ts.sblock_alloc_buffer([], dtype="float32") for i in T.serial(0, 128): @@ -982,7 +982,7 @@ def multiple_bufferstore( @Ts.prim_func def reduction_loop_not_deepest( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: for k in T.thread_binding(0, 128, thread="threadIdx.x"): for i in T.serial(0, 128): @@ -997,7 +997,7 @@ def reduction_loop_not_deepest( @Ts.prim_func def reduction_loop_bound_to_blockidx( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: for i in T.serial(0, 128): for k in T.thread_binding(0, 128, thread="blockIdx.x"): @@ -1012,7 +1012,7 @@ def reduction_loop_bound_to_blockidx( @Ts.prim_func def different_access_indices( - A: T.Buffer([128, 128, 128], dtype="float32"), B: T.Buffer([128, 128], dtype="float32") + A: T.Tensor([128, 128, 128], dtype="float32"), B: T.Tensor([128, 128], dtype="float32") ) -> None: for i, j in T.grid(128, 128): for k in T.thread_binding(0, 128, thread="threadIdx.x"): @@ -1034,7 +1034,7 @@ def different_access_indices( @Ts.prim_func def invalid_reducer( - A: T.Buffer([128, 128], dtype="float32"), B: T.Buffer([128], dtype="float32") + A: T.Tensor([128, 128], dtype="float32"), B: T.Tensor([128], dtype="float32") ) -> None: for i in T.serial(0, 128): for k in T.thread_binding(0, 128, thread="threadIdx.x"): @@ -1049,7 +1049,7 @@ def invalid_reducer( @Ts.prim_func def softmax( - A: T.Buffer([256, 256], dtype="float32"), T_softmax_norm: T.Buffer([256, 256], dtype="float32") + A: T.Tensor([256, 256], dtype="float32"), T_softmax_norm: T.Tensor([256, 256], dtype="float32") ) -> None: T_softmax_maxelem_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") T_softmax_expsum_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") @@ -1107,7 +1107,7 @@ def softmax( @Ts.prim_func def lowered_softmax( - A: T.Buffer([256, 256], dtype="float32"), T_softmax_norm: T.Buffer([256, 256], dtype="float32") + A: T.Tensor([256, 256], dtype="float32"), T_softmax_norm: T.Tensor([256, 256], dtype="float32") ) -> None: T_softmax_maxelem_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") T_softmax_expsum_shared = Ts.sblock_alloc_buffer([256], dtype="float32", scope="shared") @@ -1217,10 +1217,10 @@ def lowered_softmax( @Ts.prim_func def argmax_split( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0 in T.grid(128, 4): for i1_1 in T.thread_binding(32, thread="threadIdx.x"): @@ -1244,10 +1244,10 @@ def argmax_split( @Ts.prim_func def lowered_argmax_split( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmax_v0: T.Buffer((128,), "int32"), - argmax_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmax_v0: T.Tensor((128,), "int32"), + argmax_v1: T.Tensor((128,), "float32"), ) -> None: cross_thread_argmax_v0 = Ts.sblock_alloc_buffer([1], dtype="int32", strides=[1], scope="local") cross_thread_argmax_v1 = Ts.sblock_alloc_buffer( @@ -1312,10 +1312,10 @@ def lowered_argmax_split( @Ts.prim_func def argmin_split_init_update_reordered( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmin_v0: T.Buffer((128,), "int32"), - argmin_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmin_v0: T.Tensor((128,), "int32"), + argmin_v1: T.Tensor((128,), "float32"), ) -> None: for i0, i1_0 in T.grid(128, 4): for i1_1 in T.thread_binding(32, thread="threadIdx.x"): @@ -1339,10 +1339,10 @@ def argmin_split_init_update_reordered( @Ts.prim_func def lowered_argmin_split_init_update_reordered( - idx: T.Buffer((128, 128), "int32"), - val: T.Buffer((128, 128), "float32"), - argmin_v0: T.Buffer((128,), "int32"), - argmin_v1: T.Buffer((128,), "float32"), + idx: T.Tensor((128, 128), "int32"), + val: T.Tensor((128, 128), "float32"), + argmin_v0: T.Tensor((128,), "int32"), + argmin_v1: T.Tensor((128,), "float32"), ) -> None: cross_thread_argmin_v0 = Ts.sblock_alloc_buffer([1], dtype="int32", strides=[1], scope="local") cross_thread_argmin_v1 = Ts.sblock_alloc_buffer( @@ -1407,10 +1407,10 @@ def lowered_argmin_split_init_update_reordered( @Ts.prim_func def layer_norm_tuple_sum( - data: T.Buffer((128, 768), "float32"), - gamma: T.Buffer(768, "float32"), - bias: T.Buffer(768, "float32"), - T_layer_norm: T.Buffer((128, 768), "float32"), + data: T.Tensor((128, 768), "float32"), + gamma: T.Tensor(768, "float32"), + bias: T.Tensor(768, "float32"), + T_layer_norm: T.Tensor((128, 768), "float32"), ) -> None: data_red_temp_v0 = Ts.sblock_alloc_buffer([128], dtype="float32") data_red_temp_v1 = Ts.sblock_alloc_buffer([128], dtype="float32") @@ -1457,10 +1457,10 @@ def layer_norm_tuple_sum( @Ts.prim_func def lowered_layer_norm_tuple_sum( - data: T.Buffer((128, 768), "float32"), - gamma: T.Buffer(768, "float32"), - bias: T.Buffer(768, "float32"), - T_layer_norm: T.Buffer((128, 768), "float32"), + data: T.Tensor((128, 768), "float32"), + gamma: T.Tensor(768, "float32"), + bias: T.Tensor(768, "float32"), + T_layer_norm: T.Tensor((128, 768), "float32"), ) -> None: # with Ts.sblock("root") data_red_temp_v0 = Ts.sblock_alloc_buffer([128], dtype="float32") @@ -1551,7 +1551,7 @@ def lowered_layer_norm_tuple_sum( @Ts.prim_func -def thread_broadcast_1(A: T.Buffer((256, 256), "float32"), B: T.Buffer((256,), "float32")): +def thread_broadcast_1(A: T.Tensor((256, 256), "float32"), B: T.Tensor((256,), "float32")): temp_local = Ts.sblock_alloc_buffer((256,), scope="local") for i in T.thread_binding(256, thread="blockIdx.x"): for k in T.thread_binding(256, thread="threadIdx.x"): @@ -1571,7 +1571,7 @@ def thread_broadcast_1(A: T.Buffer((256, 256), "float32"), B: T.Buffer((256,), " # complains that k is defined outside of a block @Ts.prim_func(check_well_formed=False) -def lowered_thread_broadcast_1(A: T.Buffer((256, 256), "float32"), B: T.Buffer((256,), "float32")): +def lowered_thread_broadcast_1(A: T.Tensor((256, 256), "float32"), B: T.Tensor((256,), "float32")): temp_local = Ts.sblock_alloc_buffer((256,), scope="local") cross_thread_temp_local = Ts.sblock_alloc_buffer((1,), strides=(1,), scope="local") for i in T.thread_binding(256, thread="blockIdx.x"): @@ -1606,7 +1606,7 @@ def lowered_thread_broadcast_1(A: T.Buffer((256, 256), "float32"), B: T.Buffer(( 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"), lv1606: T.Buffer((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), lv1582: T.Buffer((T.int64(1), T.int64(1), T.int64(1), n), 'float16'), var_compute_intermediate: T.Buffer((T.int64(1), T.int64(32), T.int64(1), n))): +def thread_broadcast_2(lv1605: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16"), lv1606: T.Tensor((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), lv1582: T.Tensor((T.int64(1), T.int64(1), T.int64(1), n), 'float16'), var_compute_intermediate: T.Tensor((T.int64(1), T.int64(32), T.int64(1), n))): var_NT_matmul_intermediate_local = Ts.sblock_alloc_buffer((T.int64(1), T.int64(32), T.int64(1), n), "float16", scope="local") var_NT_matmul_intermediate_rf_local = Ts.sblock_alloc_buffer((T.int64(256), T.int64(1), T.int64(32), T.int64(1), n), "float16", scope="local") @@ -1652,7 +1652,7 @@ def thread_broadcast_2(lv1605: T.Buffer((T.int64(1), T.int64(32), T.int64(1), T. 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"), lv1606: T.Buffer((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), lv1582: T.Buffer((T.int64(1), T.int64(1), T.int64(1), n), 'float16'), var_compute_intermediate: T.Buffer((T.int64(1), T.int64(32), T.int64(1), n))): +def lowered_thread_broadcast_2(lv1605: T.Tensor((T.int64(1), T.int64(32), T.int64(1), T.int64(128)), "float16"), lv1606: T.Tensor((T.int64(1), T.int64(32), n, T.int64(128)), 'float16'), lv1582: T.Tensor((T.int64(1), T.int64(1), T.int64(1), n), 'float16'), var_compute_intermediate: T.Tensor((T.int64(1), T.int64(32), T.int64(1), n))): var_NT_matmul_intermediate_local = Ts.sblock_alloc_buffer((T.int64(1), T.int64(32), T.int64(1), n), "float16", scope="local") var_NT_matmul_intermediate_rf_local = Ts.sblock_alloc_buffer((T.int64(256), T.int64(1), T.int64(32), T.int64(1), n), "float16", scope="local") @@ -1715,7 +1715,7 @@ def lowered_thread_broadcast_2(lv1605: T.Buffer((T.int64(1), T.int64(32), T.int6 @Ts.prim_func -def no_thread_broadcast(A: T.Buffer((256, 256), "float32"), B: T.Buffer((256, 256), "float32")): +def no_thread_broadcast(A: T.Tensor((256, 256), "float32"), B: T.Tensor((256, 256), "float32")): temp_1_local = Ts.sblock_alloc_buffer((256,), scope="local") temp_2_local = Ts.sblock_alloc_buffer((1,), scope="local") for i in T.thread_binding(256, thread="blockIdx.x"): @@ -1743,7 +1743,7 @@ def no_thread_broadcast(A: T.Buffer((256, 256), "float32"), B: T.Buffer((256, 25 # complains that k is defined outside of a block @Ts.prim_func(check_well_formed=False) def lowered_no_thread_broadcast( - A: T.Buffer((256, 256), "float32"), B: T.Buffer((256, 256), "float32") + A: T.Tensor((256, 256), "float32"), B: T.Tensor((256, 256), "float32") ): temp_1_local = Ts.sblock_alloc_buffer((256,), scope="local") temp_2_local = Ts.sblock_alloc_buffer((1,), scope="local") diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lower_init_block.py b/tests/python/s_tir/transform/test_s_tir_transform_lower_init_block.py index f34c1e93e64d..c2dfa45d8fca 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lower_init_block.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lower_init_block.py @@ -26,7 +26,7 @@ @tvm.script.ir_module class WithInit: @Ts.prim_func - def main(A: T.Buffer([64, 64, 64]), B: T.Buffer([64])) -> None: + def main(A: T.Tensor([64, 64, 64]), B: T.Tensor([64])) -> None: for i0, j0 in T.grid(64, 64): for k0 in T.serial(32, 64): with Ts.sblock(): @@ -39,7 +39,7 @@ def main(A: T.Buffer([64, 64, 64]), B: T.Buffer([64])) -> None: @tvm.script.ir_module class WithBranch: @Ts.prim_func - def main(A: T.Buffer([64, 64, 64]), B: T.Buffer([64])) -> None: + def main(A: T.Tensor([64, 64, 64]), B: T.Tensor([64])) -> None: for i0, j0 in T.grid(64, 64): for k0 in T.serial(32, 64): with Ts.sblock(): @@ -54,7 +54,7 @@ def main(A: T.Buffer([64, 64, 64]), B: T.Buffer([64])) -> None: @tvm.script.ir_module class InitWithMatchBuffer: @Ts.prim_func - def main(A: T.Buffer([64, 64, 64]), B: T.Buffer([64])) -> None: + def main(A: T.Tensor([64, 64, 64]), B: T.Tensor([64])) -> None: for i0, j0 in T.grid(64, 64): for k0 in T.serial(32, 64): with Ts.sblock(): @@ -69,7 +69,7 @@ def main(A: T.Buffer([64, 64, 64]), B: T.Buffer([64])) -> None: @tvm.script.ir_module class BranchWithMatchBuffer: @Ts.prim_func - def main(A: T.Buffer([64, 64, 64]), B: T.Buffer([64])) -> None: + def main(A: T.Tensor([64, 64, 64]), B: T.Tensor([64])) -> None: for i0, j0 in T.grid(64, 64): for k0 in T.serial(32, 64): with Ts.sblock(): 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 e1aab5cd42a0..2f0f121be5d6 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 @@ -40,7 +40,7 @@ def _check_fail(original): @Ts.prim_func -def buffer_load_store(A: T.Buffer((16, 16, 16)), C: T.Buffer((16, 16))) -> None: +def buffer_load_store(A: T.Tensor((16, 16, 16)), C: T.Tensor((16, 16))) -> None: for i, j, k in T.grid(4, 16, 8): with Ts.sblock(): Ts.reads(C[i * 4 : i * 4 + 4, k * 2 : k * 2 + 2]) @@ -56,7 +56,7 @@ def buffer_load_store(A: T.Buffer((16, 16, 16)), C: T.Buffer((16, 16))) -> None: @Ts.prim_func -def transformed_buffer_load_store(A: T.Buffer((16, 16, 16)), C: T.Buffer((16, 16))) -> None: +def transformed_buffer_load_store(A: T.Tensor((16, 16, 16)), C: T.Tensor((16, 16))) -> None: for i, j, k in T.grid(4, 16, 8): with Ts.sblock(): Ts.reads(C[i * 4 : i * 4 + 4, k * 2 : k * 2 + 2]) @@ -78,7 +78,7 @@ def intrin_test(data, elem_offset, stride_0, stride_1, shape_0, shape_1): @Ts.prim_func -def opaque_access(A: T.Buffer((32, 64, 128)), B: T.Buffer((64, 64, 64))) -> None: +def opaque_access(A: T.Tensor((32, 64, 128)), B: T.Tensor((64, 64, 64))) -> None: for i, j, k in T.grid(2, 64, 8): with Ts.sblock(): Ts.reads([]) @@ -122,7 +122,7 @@ def opaque_access(A: T.Buffer((32, 64, 128)), B: T.Buffer((64, 64, 64))) -> None @Ts.prim_func -def transformed_opaque_access(A: T.Buffer((32, 64, 128)), B: T.Buffer((64, 64, 64))) -> None: +def transformed_opaque_access(A: T.Tensor((32, 64, 128)), B: T.Tensor((64, 64, 64))) -> None: for i, j, k in T.grid(2, 64, 8): with Ts.sblock(): Ts.reads([]) @@ -154,7 +154,7 @@ def transformed_opaque_access(A: T.Buffer((32, 64, 128)), B: T.Buffer((64, 64, 6 @Ts.prim_func -def opaque_buffer_data_projection(A: T.Buffer((16,))) -> None: +def opaque_buffer_data_projection(A: T.Tensor((16,))) -> None: with Ts.sblock(): Ts.reads([]) Ts.writes(A[4:8]) @@ -163,7 +163,7 @@ def opaque_buffer_data_projection(A: T.Buffer((16,))) -> None: @Ts.prim_func -def transformed_opaque_buffer_data_projection(A: T.Buffer((16,))) -> None: +def transformed_opaque_buffer_data_projection(A: T.Tensor((16,))) -> None: with Ts.sblock(): Ts.reads([]) Ts.writes(A[4:8]) @@ -175,7 +175,7 @@ def transformed_opaque_buffer_data_projection(A: T.Buffer((16,))) -> None: @Ts.prim_func -def high_dim_opaque_access(A: T.Buffer((16, 32, 64))) -> None: +def high_dim_opaque_access(A: T.Tensor((16, 32, 64))) -> None: for i, j, k in T.grid(16, 2, 4): with Ts.sblock(): Ts.reads([]) @@ -199,7 +199,7 @@ def high_dim_opaque_access(A: T.Buffer((16, 32, 64))) -> None: @Ts.prim_func -def transformed_high_dim_opaque_access(A: T.Buffer((16, 32, 64))) -> None: +def transformed_high_dim_opaque_access(A: T.Tensor((16, 32, 64))) -> None: for i, j, k in T.grid(16, 2, 4): with Ts.sblock(): Ts.reads([]) @@ -222,7 +222,7 @@ def transformed_high_dim_opaque_access(A: T.Buffer((16, 32, 64))) -> None: @Ts.prim_func def high_dim_opaque_access_with_source_strides( - A: T.Buffer((16, 32, 64), strides=[2576, 80, 1]), + A: T.Tensor((16, 32, 64), strides=[2576, 80, 1]), ) -> None: for i, j, k in T.grid(16, 2, 4): with Ts.sblock(): @@ -248,7 +248,7 @@ def high_dim_opaque_access_with_source_strides( @Ts.prim_func def transformed_high_dim_opaque_access_with_source_strides( - A: T.Buffer((16, 32, 64), strides=[2576, 80, 1]), + A: T.Tensor((16, 32, 64), strides=[2576, 80, 1]), ) -> None: for i, j, k in T.grid(16, 2, 4): with Ts.sblock(): @@ -273,7 +273,7 @@ def transformed_high_dim_opaque_access_with_source_strides( @Ts.prim_func -def recursive_match(A: T.Buffer((64, 64, 64)), B: T.Buffer((64, 64, 64))) -> None: +def recursive_match(A: T.Tensor((64, 64, 64)), B: T.Tensor((64, 64, 64))) -> None: for i, j, k in T.grid(64, 4, 4): with Ts.sblock(): Ts.reads([]) @@ -329,7 +329,7 @@ def recursive_match(A: T.Buffer((64, 64, 64)), B: T.Buffer((64, 64, 64))) -> Non @Ts.prim_func -def transformed_recursive_match(A: T.Buffer((64, 64, 64)), B: T.Buffer((64, 64, 64))) -> None: +def transformed_recursive_match(A: T.Tensor((64, 64, 64)), B: T.Tensor((64, 64, 64))) -> None: for i, j, k in T.grid(64, 4, 4): with Ts.sblock(): Ts.reads([]) @@ -376,8 +376,8 @@ def transformed_recursive_match(A: T.Buffer((64, 64, 64)), B: T.Buffer((64, 64, @Ts.prim_func def symbolic_match( - A: T.Buffer((n * m, m)), # noqa: F821 - B: T.Buffer((n * 2, m * 4)), # noqa: F821 + A: T.Tensor((n * m, m)), # noqa: F821 + B: T.Tensor((n * 2, m * 4)), # noqa: F821 n: T.int32, m: T.int32, ) -> None: @@ -406,8 +406,8 @@ def symbolic_match( @Ts.prim_func def transformed_symbolic_match( - A: T.Buffer((n * m, m)), # noqa: F821 - B: T.Buffer((n * 2, m * 4)), # noqa: F821 + A: T.Tensor((n * m, m)), # noqa: F821 + B: T.Tensor((n * 2, m * 4)), # noqa: F821 n: T.int32, m: T.int32, ) -> None: @@ -431,7 +431,7 @@ def transformed_symbolic_match( @Ts.prim_func -def rank0_buffer(A: T.Buffer((8, 8)), B: T.Buffer((8, 8))) -> None: +def rank0_buffer(A: T.Tensor((8, 8)), B: T.Tensor((8, 8))) -> None: for i, j in T.grid(8, 8): with Ts.sblock(): Ts.reads([]) @@ -452,7 +452,7 @@ def rank0_buffer(A: T.Buffer((8, 8)), B: T.Buffer((8, 8))) -> None: @Ts.prim_func -def transformed_rank0_buffer(A: T.Buffer((8, 8)), B: T.Buffer((8, 8))) -> None: +def transformed_rank0_buffer(A: T.Tensor((8, 8)), B: T.Tensor((8, 8))) -> None: for i, j in T.grid(8, 8): with Ts.sblock(): Ts.reads([]) @@ -471,7 +471,7 @@ def transformed_rank0_buffer(A: T.Buffer((8, 8)), B: T.Buffer((8, 8))) -> None: @Ts.prim_func -def fail_match_load(A: T.Buffer((8, 8))) -> None: +def fail_match_load(A: T.Tensor((8, 8))) -> None: for i, j in T.grid(8, 8): with Ts.sblock(): Ts.reads(A[i, j]) @@ -481,7 +481,7 @@ def fail_match_load(A: T.Buffer((8, 8))) -> None: @Ts.prim_func -def fail_match_store(A: T.Buffer((8, 8))) -> None: +def fail_match_store(A: T.Tensor((8, 8))) -> None: for i, j in T.grid(8, 8): with Ts.sblock(): Ts.reads([]) @@ -495,7 +495,7 @@ def fail_match_store(A: T.Buffer((8, 8))) -> None: @Ts.prim_func(check_well_formed=False) -def fail_buffer_bind(A: T.Buffer((8, 8))) -> None: +def fail_buffer_bind(A: T.Tensor((8, 8))) -> None: for i, j in T.grid(8, 2): with Ts.sblock(): sub_A = Ts.match_buffer( @@ -507,7 +507,7 @@ def fail_buffer_bind(A: T.Buffer((8, 8))) -> None: # well-formed checker complains about redefinition of a stride variable @Ts.prim_func(check_well_formed=False) -def fail_match_func_param(A: T.Buffer((8, 8)), m: T.int32, n: T.int32) -> None: +def fail_match_func_param(A: T.Tensor((8, 8)), m: T.int32, n: T.int32) -> None: for i, j in T.grid(8, 2): with Ts.sblock(): sub_A = Ts.match_buffer( @@ -560,7 +560,7 @@ def test_fail_match_func_param(): @Ts.prim_func -def scalar_match_buffer_type_coercion(A: T.Buffer((8, 8))) -> None: +def scalar_match_buffer_type_coercion(A: T.Tensor((8, 8))) -> None: for i, j in T.grid(8, 8): with Ts.sblock(""): vi = Ts.axis.spatial(8, i) @@ -573,7 +573,7 @@ def scalar_match_buffer_type_coercion(A: T.Buffer((8, 8))) -> None: @Ts.prim_func -def transformed_scalar_match_buffer_type_coercion(A: T.Buffer((8, 8))) -> None: +def transformed_scalar_match_buffer_type_coercion(A: T.Tensor((8, 8))) -> None: for i, j in T.grid(8, 8): with Ts.sblock(""): vi = Ts.axis.spatial(8, i) @@ -589,7 +589,7 @@ def test_scalar_match_buffer_type_coercion(): @Ts.prim_func -def masked_match_buffer(A: T.Buffer((8,), "float32")) -> None: +def masked_match_buffer(A: T.Tensor((8,), "float32")) -> None: with Ts.sblock(): Ts.reads(A[2:6]) sub_A = Ts.match_buffer(A[2:6], (4,), 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 3b04dcf0d994..3dffe078171b 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 @@ -35,7 +35,7 @@ def _check(original, transformed): @Ts.prim_func def compacted_elementwise_func( - A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32") ) -> None: for i in range(0, 16): with Ts.sblock(): @@ -56,10 +56,10 @@ def compacted_elementwise_func( @Ts.prim_func def transformed_elementwise_func( - A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32") ) -> None: for i in T.serial(0, 16): - B_new = T.alloc_buffer( + B_new = T.alloc_tensor( [1, 16], "float32", annotations={"buffer_allocated_addr": [], "buffer_data_alignment": 64}, @@ -71,7 +71,7 @@ def transformed_elementwise_func( @Ts.prim_func -def compacted_gpu_func(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")) -> None: +def compacted_gpu_func(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")) -> None: for i0 in T.thread_binding(0, 4, thread="blockIdx.x"): for i1 in T.thread_binding(0, 2, thread="threadIdx.x"): for i2 in T.thread_binding(0, 2, thread="vthread"): @@ -93,7 +93,7 @@ def compacted_gpu_func(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), " @Ts.prim_func def transformed_gpu_func( - A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32") ) -> None: i0 = T.env_thread("blockIdx.x") i1 = T.env_thread("threadIdx.x") @@ -102,7 +102,7 @@ def transformed_gpu_func( T.launch_thread(i0, 4) T.launch_thread(i1, 2) T.launch_thread(i2, 2) - B = T.alloc_buffer( + B = T.alloc_tensor( [1, 16], "float32", scope="local", @@ -116,8 +116,8 @@ def transformed_gpu_func( @Ts.prim_func def compacted_symbolic_func( - A: T.Buffer((n, m), "float32"), # noqa: F821 - C: T.Buffer((n, m), "float32"), # noqa: F821 + A: T.Tensor((n, m), "float32"), # noqa: F821 + C: T.Tensor((n, m), "float32"), # noqa: F821 n: T.int32, m: T.int32, ) -> None: @@ -140,13 +140,13 @@ def compacted_symbolic_func( @Ts.prim_func def transformed_symbolic_func( - A: T.Buffer((n, m), "float32"), # noqa: F821 - C: T.Buffer((n, m), "float32"), # noqa: F821 + A: T.Tensor((n, m), "float32"), # noqa: F821 + C: T.Tensor((n, m), "float32"), # noqa: F821 n: T.int32, m: T.int32, ) -> None: for i in range(0, n): - B = T.alloc_buffer( + B = T.alloc_tensor( [m], "float32", annotations={"buffer_allocated_addr": [], "buffer_data_alignment": 64}, @@ -158,7 +158,7 @@ def transformed_symbolic_func( @Ts.prim_func -def compacted_predicate_func(A: T.Buffer(32, "float32"), C: T.Buffer(32, "float32")) -> None: +def compacted_predicate_func(A: T.Tensor(32, "float32"), C: T.Tensor(32, "float32")) -> None: for i, j in T.grid(5, 7): with Ts.sblock(): Ts.reads(A[i * 7 + j]) @@ -168,14 +168,14 @@ def compacted_predicate_func(A: T.Buffer(32, "float32"), C: T.Buffer(32, "float3 @Ts.prim_func -def transformed_predicate_func(A: T.Buffer(32, "float32"), C: T.Buffer(32, "float32")) -> None: +def transformed_predicate_func(A: T.Tensor(32, "float32"), C: T.Tensor(32, "float32")) -> None: for i, j in T.grid(5, 7): if i * 7 + j < 32: C[i * 7 + j] = A[i * 7 + j] + 1.0 @Ts.prim_func -def compacted_unit_loop_func(A: T.Buffer(32, "float32"), C: T.Buffer(32, "float32")) -> None: +def compacted_unit_loop_func(A: T.Tensor(32, "float32"), C: T.Tensor(32, "float32")) -> None: for x, y, z in T.grid(4, 1, 8): with Ts.sblock(): Ts.reads(A[x * 8 + y * 8 + z]) @@ -184,13 +184,13 @@ def compacted_unit_loop_func(A: T.Buffer(32, "float32"), C: T.Buffer(32, "float3 @Ts.prim_func -def transformed_unit_loop_func(A: T.Buffer(32, "float32"), C: T.Buffer(32, "float32")) -> None: +def transformed_unit_loop_func(A: T.Tensor(32, "float32"), C: T.Tensor(32, "float32")) -> None: for x, z in T.grid(4, 8): C[x * 8 + z] = A[x * 8 + z] + 1.0 @Ts.prim_func -def compacted_multi_alloc_func(A: T.Buffer(32, "float32"), D: T.Buffer(32, "float32")) -> None: +def compacted_multi_alloc_func(A: T.Tensor(32, "float32"), D: T.Tensor(32, "float32")) -> None: for i in range(0, 32): with Ts.sblock(): Ts.reads(A[i]) @@ -203,14 +203,14 @@ def compacted_multi_alloc_func(A: T.Buffer(32, "float32"), D: T.Buffer(32, "floa @Ts.prim_func -def transformed_multi_alloc_func(A: T.Buffer(32, "float32"), D: T.Buffer(32, "float32")) -> None: +def transformed_multi_alloc_func(A: T.Tensor(32, "float32"), D: T.Tensor(32, "float32")) -> None: for i in range(0, 32): - B = T.alloc_buffer( + B = T.alloc_tensor( (32,), "float32", annotations={"buffer_allocated_addr": [], "buffer_data_alignment": 64}, ) - C = T.alloc_buffer( + C = T.alloc_tensor( (32,), "float32", annotations={"buffer_allocated_addr": [], "buffer_data_alignment": 64}, @@ -222,7 +222,7 @@ def transformed_multi_alloc_func(A: T.Buffer(32, "float32"), D: T.Buffer(32, "fl @Ts.prim_func def compacted_strided_buffer_func( - A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32") ) -> None: for i0 in range(0, 4): with Ts.sblock(): @@ -245,11 +245,11 @@ def compacted_strided_buffer_func( @Ts.prim_func def transformed_strided_buffer_func( - A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32") + A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32") ) -> None: # body for i0 in T.serial(4): - B = T.alloc_buffer( + B = T.alloc_tensor( [4, 16], "float32", strides=[17, 1], @@ -265,7 +265,7 @@ def transformed_strided_buffer_func( @Ts.prim_func -def compacted_symbolic_strided_buffer_func(A: T.Buffer((1, n, 10240))) -> None: +def compacted_symbolic_strided_buffer_func(A: T.Tensor((1, n, 10240))) -> None: padded_size = T.meta_var(T.min((n + 63) // 64 * 64, 96)) # with Ts.sblock("root"): for i, j, k in T.grid(((n + 63) // 64 * 4 + 7) // 8, 2, 160): @@ -287,10 +287,10 @@ def compacted_symbolic_strided_buffer_func(A: T.Buffer((1, n, 10240))) -> None: @Ts.prim_func -def transformed_symbolic_strided_buffer_func(A: T.Buffer((1, n, 10240))): +def transformed_symbolic_strided_buffer_func(A: T.Tensor((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): - A_pad_shared_dyn = T.alloc_buffer( + A_pad_shared_dyn = T.alloc_tensor( (1, padded_size, 64), strides=(72 * padded_size, 72, 1), scope="shared.dyn", @@ -306,13 +306,13 @@ def transformed_symbolic_strided_buffer_func(A: T.Buffer((1, n, 10240))): @Ts.prim_func -def annotated_loops(A: T.Buffer((16,), "float32")) -> None: +def annotated_loops(A: T.Tensor((16,), "float32")) -> None: for i in range(0, 16, annotations={"pragma_1": "str_value", "pragma_2": 1, "pragma_3": 0.0}): A[i] = 0.0 @Ts.prim_func -def boolean_handling_before(a: T.Buffer(10, "bool"), b: T.Buffer(10, "bool")) -> None: +def boolean_handling_before(a: T.Tensor(10, "bool"), b: T.Tensor(10, "bool")) -> None: for i0 in T.serial(10): with Ts.sblock("b"): Ts.reads(a[i0]) @@ -321,7 +321,7 @@ def boolean_handling_before(a: T.Buffer(10, "bool"), b: T.Buffer(10, "bool")) -> @Ts.prim_func -def boolean_handling_after(a: T.Buffer(10, "bool"), b: T.Buffer(10, "bool")) -> None: +def boolean_handling_after(a: T.Tensor(10, "bool"), b: T.Tensor(10, "bool")) -> None: # body for i0 in T.serial(10): b[i0] = a[i0] @@ -393,14 +393,14 @@ def annotated_block() -> None: def test_preserved_annotations(): @Ts.prim_func - def before(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): + def before(A: T.Tensor(8, "float32"), B: T.Tensor(8, "float32")): for i in T.serial(8, annotations={"k_0": 1, "k_1": [2, 3], "k_2": 3.14}): with Ts.sblock("block"): Ts.sblock_attr({"k_3": "oops"}) B[i] = A[i] + 1.0 @Ts.prim_func - def after(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): + def after(A: T.Tensor(8, "float32"), B: T.Tensor(8, "float32")): for i in T.serial(8, annotations={"k_0": 1, "k_1": [2, 3], "k_2": 3.14}): B[i] = A[i] + 1.0 @@ -411,13 +411,13 @@ def after(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): def test_none_pragma_annotation(): @Ts.prim_func - def before(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): + def before(A: T.Tensor(8, "float32"), B: T.Tensor(8, "float32")): for i in T.serial(8, annotations={"pragma_unroll_explicit": None}): with Ts.sblock("block"): B[i] = A[i] + 1.0 @Ts.prim_func - def after(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): + def after(A: T.Tensor(8, "float32"), B: T.Tensor(8, "float32")): for i in T.serial(8): B[i] = A[i] + 1.0 diff --git a/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py b/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py index d10e32fb8297..0ab0f63c2647 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_lower_thread_all_reduce.py @@ -40,7 +40,7 @@ def _has_volatile_alloc_buffer(mod): def visit(node): nonlocal has_volatile_alloc - if _is_buffer_binding(node, "tirx.alloc_buffer") and "tirx.volatile" in node.value.attrs: + if _is_buffer_binding(node, "tirx.alloc_tensor") and "tirx.volatile" in node.value.attrs: has_volatile_alloc = has_volatile_alloc or node.value.attrs["tirx.volatile"] is True tvm_ffi.structural_walk(mod["main"].body, visit) @@ -53,15 +53,15 @@ def test_basic(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((128, 32), "float32"), B: T.Buffer(128, "float32")): + def main(A: T.Tensor((128, 32), "float32"), B: T.Tensor(128, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) - A_flat = T.decl_buffer(4096, data=A.data) + A_flat = T.decl_tensor(4096, data=A.data) for i in range(128): threadIdx_x = T.launch_thread("threadIdx.x", 32) - reduce = T.alloc_buffer((1,), scope="local") - reduce_1 = T.decl_buffer(1, data=reduce.data, scope="local") + reduce = T.alloc_tensor((1,), scope="local") + reduce_1 = T.decl_tensor(1, data=reduce.data, scope="local") with T.attr( T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), @@ -93,14 +93,14 @@ def test_basic_with_decl_buffer(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((128, 32), "float32"), B: T.Buffer(128, "float32")): + def main(A: T.Tensor((128, 32), "float32"), B: T.Tensor(128, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) - A_flat = T.decl_buffer(4096, data=A.data) + A_flat = T.decl_tensor(4096, data=A.data) for i in range(128): threadIdx_x = T.launch_thread("threadIdx.x", 32) - reduce = T.decl_buffer(1, dtype="float32", scope="local") + reduce = T.decl_tensor(1, dtype="float32", scope="local") with T.attr( T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), @@ -130,18 +130,18 @@ def test_reduce_summation(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((128, 128), "float32"), B: T.Buffer(128, "float32")): + def main(A: T.Tensor((128, 128), "float32"), B: T.Tensor(128, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) - A_flat = T.decl_buffer(16384, data=A.data) + A_flat = T.decl_tensor(16384, data=A.data) for i in range(128): threadIdx_x = T.launch_thread("threadIdx.x", 32) - normal_reduce = T.alloc_buffer((1,), scope="local") - normal_reduce_1 = T.decl_buffer(1, data=normal_reduce.data, scope="local") + normal_reduce = T.alloc_tensor((1,), scope="local") + normal_reduce_1 = T.decl_tensor(1, data=normal_reduce.data, scope="local") - reduce = T.alloc_buffer((1,), scope="local") - reduce_1 = T.decl_buffer(1, data=reduce.data, scope="local") + reduce = T.alloc_tensor((1,), scope="local") + reduce_1 = T.decl_tensor(1, data=reduce.data, scope="local") normal_reduce_1[0] = T.float32(0) @@ -177,18 +177,18 @@ def test_multi_group_reduction(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((32, 32), "float32"), B: T.Buffer((32,), "float32")): + def main(A: T.Tensor((32, 32), "float32"), B: T.Tensor((32,), "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) threadIdx_y = T.launch_thread("threadIdx.y", 32) - cross_thread_B = T.alloc_buffer((1,), scope="local") + cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 32) - cross_thread_B_1 = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0), ): - A_1 = T.decl_buffer((1024,), data=A.data) + A_1 = T.decl_tensor((1024,), data=A.data) T.tvm_thread_allreduce( T.uint32(1), A_1[threadIdx_y * 32 + threadIdx_x], @@ -197,7 +197,7 @@ def main(A: T.Buffer((32, 32), "float32"), B: T.Buffer((32,), "float32")): threadIdx_x, ) if threadIdx_x == 0: - B_1 = T.decl_buffer((32,), data=B.data) + B_1 = T.decl_tensor((32,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] After = transform(Before) @@ -212,18 +212,18 @@ def test_multi_group_reduction_consumed_through_alias(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((4, 128), "float32"), B: T.Buffer((4,), "float32")): + def main(A: T.Tensor((4, 128), "float32"), B: T.Tensor((4,), "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) threadIdx_y = T.launch_thread("threadIdx.y", 4) - cross_thread_B = T.alloc_buffer((1,), scope="local") + cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 128) - cross_thread_B_alias = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + cross_thread_B_alias = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0), ): - A_flat = T.decl_buffer((512,), data=A.data) + A_flat = T.decl_tensor((512,), data=A.data) T.tvm_thread_allreduce( T.uint32(1), A_flat[threadIdx_y * 128 + threadIdx_x], @@ -233,7 +233,7 @@ def main(A: T.Buffer((4, 128), "float32"), B: T.Buffer((4,), "float32")): ) cross_thread_B_alias[0] = cross_thread_B[0] if threadIdx_x == 0: - B_flat = T.decl_buffer((4,), data=B.data) + B_flat = T.decl_tensor((4,), data=B.data) B_flat[threadIdx_y] = cross_thread_B_alias[0] After = transform(Before) @@ -252,17 +252,17 @@ def test_multi_group_reduction_with_alias_declared_after_allreduce(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((4, 128), "float32"), B: T.Buffer((4,), "float32")): + def main(A: T.Tensor((4, 128), "float32"), B: T.Tensor((4,), "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) threadIdx_y = T.launch_thread("threadIdx.y", 4) - cross_thread_B = T.alloc_buffer((1,), scope="local") + cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 128) with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0), ): - A_flat = T.decl_buffer((512,), data=A.data) + A_flat = T.decl_tensor((512,), data=A.data) T.tvm_thread_allreduce( T.uint32(1), A_flat[threadIdx_y * 128 + threadIdx_x], @@ -270,10 +270,10 @@ def main(A: T.Buffer((4, 128), "float32"), B: T.Buffer((4,), "float32")): cross_thread_B[0], threadIdx_x, ) - cross_thread_B_alias = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + cross_thread_B_alias = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") cross_thread_B[0] = cross_thread_B_alias[0] if threadIdx_x == 0: - B_flat = T.decl_buffer((4,), data=B.data) + B_flat = T.decl_tensor((4,), data=B.data) B_flat[threadIdx_y] = cross_thread_B_alias[0] After = transform(Before) @@ -292,18 +292,18 @@ def test_multi_group_mask1(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((32, 8), "float32"), B: T.Buffer((32,), "float32")): + def main(A: T.Tensor((32, 8), "float32"), B: T.Tensor((32,), "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) threadIdx_y = T.launch_thread("threadIdx.y", 32) - cross_thread_B = T.alloc_buffer((1,), scope="local") + cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 8) - cross_thread_B_1 = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0), ): - A_1 = T.decl_buffer((256,), data=A.data) + A_1 = T.decl_tensor((256,), data=A.data) T.tvm_thread_allreduce( T.uint32(1), A_1[threadIdx_y * 8 + threadIdx_x], @@ -312,7 +312,7 @@ def main(A: T.Buffer((32, 8), "float32"), B: T.Buffer((32,), "float32")): threadIdx_x, ) if threadIdx_x == 0: - B_1 = T.decl_buffer((32,), data=B.data) + B_1 = T.decl_tensor((32,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] After = transform(Before) @@ -327,18 +327,18 @@ def test_multi_warp_reduce1(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128,), "float32")): + def main(A: T.Tensor((128, 128), "float32"), B: T.Tensor((128,), "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) for i in range(128): threadIdx_x = T.launch_thread("threadIdx.x", 128) - cross_thread_B = T.alloc_buffer((1,), scope="local") - cross_thread_B_1 = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + cross_thread_B = T.alloc_tensor((1,), scope="local") + cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0), ): - A_1 = T.decl_buffer((16384,), data=A.data) + A_1 = T.decl_tensor((16384,), data=A.data) T.tvm_thread_allreduce( T.uint32(1), A_1[i * 128 + threadIdx_x], @@ -347,7 +347,7 @@ def main(A: T.Buffer((128, 128), "float32"), B: T.Buffer((128,), "float32")): threadIdx_x, ) if threadIdx_x == 0: - B_1 = T.decl_buffer((128,), data=B.data) + B_1 = T.decl_tensor((128,), data=B.data) B_1[i] = cross_thread_B_1[0] After = transform(Before) @@ -363,22 +363,22 @@ def test_multi_warp_reduce2(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((1, 1024), "float32"), B: T.Buffer((1,), "float32")): + def main(A: T.Tensor((1, 1024), "float32"), B: T.Tensor((1,), "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) threadIdx_x = T.launch_thread("threadIdx.x", 1024) - cross_thread_B = T.alloc_buffer((1,), scope="local") - cross_thread_B_1 = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + cross_thread_B = T.alloc_tensor((1,), scope="local") + cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0), ): - A_1 = T.decl_buffer((1024,), data=A.data) + A_1 = T.decl_tensor((1024,), data=A.data) T.tvm_thread_allreduce( T.uint32(1), A_1[threadIdx_x], T.bool(True), cross_thread_B_1[0], threadIdx_x ) if threadIdx_x == 0: - B_1 = T.decl_buffer((1,), data=B.data) + B_1 = T.decl_tensor((1,), data=B.data) B_1[0] = cross_thread_B_1[0] After = transform(Before) @@ -394,18 +394,18 @@ def test_multi_group_multi_warp_reduction(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((4, 128), "float32"), B: T.Buffer((4,), "float32")): + def main(A: T.Tensor((4, 128), "float32"), B: T.Tensor((4,), "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) threadIdx_y = T.launch_thread("threadIdx.y", 4) - cross_thread_B = T.alloc_buffer((1,), scope="local") + cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 128) - cross_thread_B_1 = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0), ): - A_1 = T.decl_buffer((512,), data=A.data) + A_1 = T.decl_tensor((512,), data=A.data) T.tvm_thread_allreduce( T.uint32(1), A_1[threadIdx_y * 128 + threadIdx_x], @@ -414,7 +414,7 @@ def main(A: T.Buffer((4, 128), "float32"), B: T.Buffer((4,), "float32")): threadIdx_x, ) if threadIdx_x == 0: - B_1 = T.decl_buffer((4,), data=B.data) + B_1 = T.decl_tensor((4,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] After = transform(Before) @@ -430,18 +430,18 @@ def test_multi_group_multi_warp_predicated_reduction(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((2, 70), "float32"), B: T.Buffer((2,), "float32")): + def main(A: T.Tensor((2, 70), "float32"), B: T.Tensor((2,), "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) threadIdx_y = T.launch_thread("threadIdx.y", 2) - in_thread_B = T.alloc_buffer((1,), scope="local") - cross_thread_B = T.alloc_buffer((1,), scope="local") + in_thread_B = T.alloc_tensor((1,), scope="local") + cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 512) - in_thread_B_1 = T.decl_buffer((1,), data=in_thread_B.data, scope="local") + in_thread_B_1 = T.decl_tensor((1,), data=in_thread_B.data, scope="local") in_thread_B_1[0] = T.float32(0) if threadIdx_x < 70: - A_1 = T.decl_buffer((140,), data=A.data) + A_1 = T.decl_tensor((140,), data=A.data) in_thread_B_1[0] = in_thread_B_1[0] + A_1[threadIdx_y * 70 + threadIdx_x] - cross_thread_B_1 = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", @@ -451,7 +451,7 @@ def main(A: T.Buffer((2, 70), "float32"), B: T.Buffer((2,), "float32")): T.uint32(1), in_thread_B_1[0], T.bool(True), cross_thread_B_1[0], threadIdx_x ) if threadIdx_x == 0: - B_1 = T.decl_buffer((2,), data=B.data) + B_1 = T.decl_tensor((2,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] After = transform(Before) @@ -467,7 +467,7 @@ def test_metal_no_mask(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((1, 1, 2, 128), "float32"), B: T.Buffer((1, 1, 2), "float32")): + def main(A: T.Tensor((1, 1, 2, 128), "float32"), B: T.Tensor((1, 1, 2), "float32")): T.func_attr( { "target": T.target( @@ -481,17 +481,17 @@ def main(A: T.Buffer((1, 1, 2, 128), "float32"), B: T.Buffer((1, 1, 2), "float32 } ) blockIdx_x = T.launch_thread("blockIdx.x", 1) - cross_thread_B = T.alloc_buffer((1,), scope="local") + cross_thread_B = T.alloc_tensor((1,), scope="local") threadIdx_z = T.launch_thread("threadIdx.z", 1) threadIdx_y = T.launch_thread("threadIdx.y", 2) threadIdx_x = T.launch_thread("threadIdx.x", 128) - cross_thread_B_1 = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0), ): - A_1 = T.decl_buffer((256,), data=A.data) + A_1 = T.decl_tensor((256,), data=A.data) T.tvm_thread_allreduce( T.uint32(1), A_1[threadIdx_y * 128 + threadIdx_x], @@ -500,7 +500,7 @@ def main(A: T.Buffer((1, 1, 2, 128), "float32"), B: T.Buffer((1, 1, 2), "float32 threadIdx_x, ) if threadIdx_x == 0: - B_1 = T.decl_buffer((2,), data=B.data) + B_1 = T.decl_tensor((2,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] After = transform(Before) @@ -517,7 +517,7 @@ def test_webgpu_warp_reduce(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((128, 32), "float32"), B: T.Buffer(128, "float32")): + def main(A: T.Tensor((128, 32), "float32"), B: T.Tensor(128, "float32")): T.func_attr( { "target": T.target( @@ -529,14 +529,14 @@ def main(A: T.Buffer((128, 32), "float32"), B: T.Buffer(128, "float32")): ), } ) - A_flat = T.decl_buffer(4096, data=A.data) + A_flat = T.decl_tensor(4096, data=A.data) for i in range(128): threadIdx_x = T.launch_thread("threadIdx.x", 32) - reduce_data = T.alloc_buffer((1,), "float32", scope="local") - reduce = T.decl_buffer(1, data=reduce_data.data, scope="local") - reduce_alias = T.decl_buffer(1, data=reduce.data, scope="local") + reduce_data = T.alloc_tensor((1,), "float32", scope="local") + reduce = T.decl_tensor(1, data=reduce_data.data, scope="local") + reduce_alias = T.decl_tensor(1, data=reduce.data, scope="local") with T.attr( T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), @@ -569,7 +569,7 @@ def test_webgpu_multi_warp_reduce(): @I.ir_module class Before: @Ts.prim_func(private=True) - def main(A: T.Buffer((1, 1, 2, 128), "float32"), B: T.Buffer((1, 1, 2), "float32")): + def main(A: T.Tensor((1, 1, 2, 128), "float32"), B: T.Tensor((1, 1, 2), "float32")): T.func_attr( { "target": T.target( @@ -583,17 +583,17 @@ def main(A: T.Buffer((1, 1, 2, 128), "float32"), B: T.Buffer((1, 1, 2), "float32 } ) blockIdx_x = T.launch_thread("blockIdx.x", 1) - cross_thread_B = T.alloc_buffer((1,), "float32", scope="local") + cross_thread_B = T.alloc_tensor((1,), "float32", scope="local") threadIdx_z = T.launch_thread("threadIdx.z", 1) threadIdx_y = T.launch_thread("threadIdx.y", 2) threadIdx_x = T.launch_thread("threadIdx.x", 128) - cross_thread_B_1 = T.decl_buffer((1,), data=cross_thread_B.data, scope="local") + cross_thread_B_1 = T.decl_tensor((1,), data=cross_thread_B.data, scope="local") with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", T.int32(0), ): - A_1 = T.decl_buffer((256,), data=A.data) + A_1 = T.decl_tensor((256,), data=A.data) T.tvm_thread_allreduce( T.uint32(1), A_1[threadIdx_y * 128 + threadIdx_x], @@ -602,7 +602,7 @@ def main(A: T.Buffer((1, 1, 2, 128), "float32"), B: T.Buffer((1, 1, 2), "float32 threadIdx_x, ) if threadIdx_x == 0: - B_1 = T.decl_buffer((2,), data=B.data) + B_1 = T.decl_tensor((2,), data=B.data) B_1[threadIdx_y] = cross_thread_B_1[0] After = transform(Before) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_manifest_shared_memory_local_stage.py b/tests/python/s_tir/transform/test_s_tir_transform_manifest_shared_memory_local_stage.py index a770963b5d58..1545f9180695 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_manifest_shared_memory_local_stage.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_manifest_shared_memory_local_stage.py @@ -28,7 +28,7 @@ @tvm.script.ir_module class MatmulBefore: @Ts.prim_func - def main(A: T.Buffer((1024, 1024), "float32"), B: T.Buffer((1024, 1024), "float32"), C: T.Buffer((1024, 1024), "float32")) -> None: + def main(A: T.Tensor((1024, 1024), "float32"), B: T.Tensor((1024, 1024), "float32"), C: T.Tensor((1024, 1024), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) # body @@ -69,7 +69,7 @@ def main(A: T.Buffer((1024, 1024), "float32"), B: T.Buffer((1024, 1024), "float3 @tvm.script.ir_module class MatmulAfter: @Ts.prim_func - def main(A: T.Buffer((1024, 1024), "float32"), B: T.Buffer((1024, 1024), "float32"), C: T.Buffer((1024, 1024), "float32")) -> None: + def main(A: T.Tensor((1024, 1024), "float32"), B: T.Tensor((1024, 1024), "float32"), C: T.Tensor((1024, 1024), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) # body 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 52f66e62b3be..249e45b04a21 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 @@ -39,7 +39,7 @@ def _is_buffer_binding(node, *op_names): @tvm.script.ir_module class Transpose: @Ts.prim_func - def main(A: T.Buffer([1024, 1024]), B: T.Buffer([1024, 1024])) -> None: + def main(A: T.Tensor([1024, 1024]), B: T.Tensor([1024, 1024])) -> None: with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for ty in T.thread_binding(8, thread="threadIdx.y"): @@ -60,7 +60,7 @@ def main(A: T.Buffer([1024, 1024]), B: T.Buffer([1024, 1024])) -> None: @tvm.script.ir_module class GlobalToShared: @Ts.prim_func - def main(A: T.Buffer([1024, 1024]), B: T.Buffer([1024, 1024])) -> None: + def main(A: T.Tensor([1024, 1024]), B: T.Tensor([1024, 1024])) -> None: with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): @@ -82,7 +82,7 @@ def main(A: T.Buffer([1024, 1024]), B: T.Buffer([1024, 1024])) -> None: @tvm.script.ir_module class SharedToGlobal: @Ts.prim_func - def main(A: T.Buffer([1024, 1024]), B: T.Buffer([1024, 1024])) -> None: + def main(A: T.Tensor([1024, 1024]), B: T.Tensor([1024, 1024])) -> None: with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): @@ -104,7 +104,7 @@ def main(A: T.Buffer([1024, 1024]), B: T.Buffer([1024, 1024])) -> None: @tvm.script.ir_module class GlobalToSharedWithLocalStage: @Ts.prim_func - def main(A: T.Buffer([1024, 1024]), B: T.Buffer([1024, 1024])) -> None: + def main(A: T.Tensor([1024, 1024]), B: T.Tensor([1024, 1024])) -> None: with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): @@ -172,7 +172,7 @@ def main() -> None: @tvm.script.ir_module class WmmaToGlobal: @Ts.prim_func - def main(C: T.Buffer([1024, 1024])) -> None: + def main(C: T.Tensor([1024, 1024])) -> None: with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): @@ -191,7 +191,7 @@ def main(C: T.Buffer([1024, 1024])) -> None: @tvm.script.ir_module class WmmaToGlobalWithFusion: @Ts.prim_func - def main(A: T.Buffer([1024]), C: T.Buffer([1024, 1024])) -> None: + def main(A: T.Tensor([1024]), C: T.Tensor([1024, 1024])) -> None: with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): @@ -212,7 +212,7 @@ def main(A: T.Buffer([1024]), C: T.Buffer([1024, 1024])) -> None: @tvm.script.ir_module class MmaToGlobal: @Ts.prim_func - def main(C: T.Buffer([1024, 1024])) -> None: + def main(C: T.Tensor([1024, 1024])) -> None: with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): @@ -231,7 +231,7 @@ def main(C: T.Buffer([1024, 1024])) -> None: @tvm.script.ir_module class TransformedGlobalToShared: @Ts.prim_func - def main(A: T.Buffer([1024, 1024]), B: T.Buffer([1024, 1024])) -> None: + def main(A: T.Tensor([1024, 1024]), B: T.Tensor([1024, 1024])) -> None: with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): @@ -270,7 +270,7 @@ def main(A: T.Buffer([1024, 1024]), B: T.Buffer([1024, 1024])) -> None: @tvm.script.ir_module class TransformedSharedToGlobal: @Ts.prim_func - def main(A: T.Buffer([1024, 1024]), B: T.Buffer([1024, 1024])) -> None: + def main(A: T.Tensor([1024, 1024]), B: T.Tensor([1024, 1024])) -> None: with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): @@ -311,7 +311,7 @@ def main(A: T.Buffer([1024, 1024]), B: T.Buffer([1024, 1024])) -> None: @tvm.script.ir_module class TransformedGlobalToSharedWithLocalStage: @Ts.prim_func - def main(A: T.Buffer((1024, 1024)), B: T.Buffer((1024, 1024))): + def main(A: T.Tensor((1024, 1024)), B: T.Tensor((1024, 1024))): with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): @@ -581,7 +581,7 @@ def main() -> None: @tvm.script.ir_module class TransformedWmmaToGlobal: @Ts.prim_func - def main(C: T.Buffer((1024, 1024), "float32")): + def main(C: T.Tensor((1024, 1024), "float32")): with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): @@ -782,7 +782,7 @@ def main(C: T.Buffer((1024, 1024), "float32")): @tvm.script.ir_module class TransformedWmmaToGlobalWithFusion: @Ts.prim_func - def main(A: T.Buffer((1024,), "float32"), C: T.Buffer((1024, 1024), "float32")) -> None: + def main(A: T.Tensor((1024,), "float32"), C: T.Tensor((1024, 1024), "float32")) -> None: # body with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) @@ -1007,7 +1007,7 @@ def main(A: T.Buffer((1024,), "float32"), C: T.Buffer((1024, 1024), "float32")) @tvm.script.ir_module class TransformedMmaToGlobal: @Ts.prim_func - def main(C: T.Buffer((1024, 1024), "float32")): + def main(C: T.Tensor((1024, 1024), "float32")): with Ts.sblock("root"): Ts.sblock_attr({"warp_execution": True}) for bx in T.thread_binding(8, thread="blockIdx.x"): @@ -1138,7 +1138,7 @@ def verify_single_allocation(stmt, alloc_size=None): alloc_extents = [] def verify(n): - if _is_buffer_binding(n, "tirx.alloc_buffer") and n.var.scope() == "shared.dyn": + if _is_buffer_binding(n, "tirx.alloc_tensor") and n.var.scope() == "shared.dyn": num_alloc[0] += 1 alloc_extents.append(n.var.shape) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_merge_dynamic_shared_memory_allocations.py b/tests/python/s_tir/transform/test_s_tir_transform_merge_dynamic_shared_memory_allocations.py index 1859f4d55b66..ae942e457d20 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_merge_dynamic_shared_memory_allocations.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_merge_dynamic_shared_memory_allocations.py @@ -27,32 +27,32 @@ def test_matmul_t_buffer(): - """Shared allocations should be merged, preserving DeclBuffer if present + """Shared allocations should be merged, preserving DeclTensor if present This test uses a matmul PrimFunc adapted from - test_matmul_dyn_shared, using `T.Buffer` (Allocate without - DeclBuffer) for the replaced allocations. + test_matmul_dyn_shared, using `T.Tensor` (Allocate without + DeclTensor) for the replaced allocations. """ transform = tvm.s_tir.transform.MergeSharedMemoryAllocations() - buffer_func = T.Buffer + buffer_func = T.Tensor @I.ir_module class Before: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float16"), - B: T.Buffer((1024, 1024), "float16"), - matmul: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float16"), + B: T.Tensor((1024, 1024), "float16"), + matmul: T.Tensor((1024, 1024), "float32"), ): - A_flat = T.decl_buffer(1048576, "float16", data=A.data) - B_flat = T.decl_buffer(1048576, "float16", data=B.data) - matmul_flat = T.decl_buffer(1048576, data=matmul.data) + A_flat = T.decl_tensor(1048576, "float16", data=A.data) + B_flat = T.decl_tensor(1048576, "float16", data=B.data) + matmul_flat = T.decl_tensor(1048576, data=matmul.data) threadIdx_x = T.launch_thread("threadIdx.x", 16) - C_local = T.alloc_buffer((1,), "float32", scope="local") - A_sh = T.alloc_buffer((256,), "float16", scope="shared.dyn") - B_sh = T.alloc_buffer((256,), "float16", scope="shared.dyn") - C_sh = T.alloc_buffer((256,), "float32", scope="shared.dyn") + C_local = T.alloc_tensor((1,), "float32", scope="local") + A_sh = T.alloc_tensor((256,), "float16", scope="shared.dyn") + B_sh = T.alloc_tensor((256,), "float16", scope="shared.dyn") + C_sh = T.alloc_tensor((256,), "float32", scope="shared.dyn") threadIdx_y = T.launch_thread("threadIdx.y", 16) blockIdx_x = T.launch_thread("blockIdx.x", 64) blockIdx_y = T.launch_thread("blockIdx.y", 64) @@ -85,22 +85,22 @@ def main( class Expected: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float16"), - B: T.Buffer((1024, 1024), "float16"), - matmul: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float16"), + B: T.Tensor((1024, 1024), "float16"), + matmul: T.Tensor((1024, 1024), "float32"), ): - A_flat = T.decl_buffer(1048576, "float16", data=A.data) - B_flat = T.decl_buffer(1048576, "float16", data=B.data) - matmul_flat = T.decl_buffer(1048576, data=matmul.data) + A_flat = T.decl_tensor(1048576, "float16", data=A.data) + B_flat = T.decl_tensor(1048576, "float16", data=B.data) + matmul_flat = T.decl_tensor(1048576, data=matmul.data) threadIdx_x = T.launch_thread("threadIdx.x", 16) - buf_dyn_shmem = T.alloc_buffer((1024,), "uint8", scope="shared.dyn") + buf_dyn_shmem = T.alloc_tensor((1024,), "uint8", scope="shared.dyn") - C_local = T.alloc_buffer((1,), "float32", scope="local") - A_sh = T.decl_buffer(256, "float16", data=buf_dyn_shmem.data, scope="shared.dyn") - B_sh = T.decl_buffer(256, "float16", data=buf_dyn_shmem.data, scope="shared.dyn") - C_sh = T.decl_buffer(256, "float32", data=buf_dyn_shmem.data, scope="shared.dyn") + C_local = T.alloc_tensor((1,), "float32", scope="local") + A_sh = T.decl_tensor(256, "float16", data=buf_dyn_shmem.data, scope="shared.dyn") + B_sh = T.decl_tensor(256, "float16", data=buf_dyn_shmem.data, scope="shared.dyn") + C_sh = T.decl_tensor(256, "float32", data=buf_dyn_shmem.data, scope="shared.dyn") threadIdx_y = T.launch_thread("threadIdx.y", 16) blockIdx_x = T.launch_thread("blockIdx.x", 64) @@ -128,7 +128,7 @@ def main( After = transform(Before) script = After["main"].script() # Verify merged allocation: one shared.dyn buffer of 1024 bytes (256*2 float16 + 256 float32) - assert "alloc_buffer((1024,)" in script + assert "alloc_tensor((1024,)" in script assert '"uint8"' in script assert '"shared.dyn"' in script # Verify storage sync calls preserved @@ -138,32 +138,32 @@ def main( def test_matmul_decl_buffer(): - """Shared allocations should be merged, preserving DeclBuffer if present + """Shared allocations should be merged, preserving DeclTensor if present This test uses a matmul PrimFunc adapted from - test_matmul_dyn_shared, using `T.decl_buffer` (Allocate followed by DeclBuffer) + test_matmul_dyn_shared, using `T.decl_tensor` (Allocate followed by DeclTensor) for the replaced allocations. """ transform = tvm.s_tir.transform.MergeSharedMemoryAllocations() - buffer_func = T.decl_buffer + buffer_func = T.decl_tensor @I.ir_module class Before: @Ts.prim_func def main( - A: T.Buffer((1024, 1024), "float16"), - B: T.Buffer((1024, 1024), "float16"), - matmul: T.Buffer((1024, 1024), "float32"), + A: T.Tensor((1024, 1024), "float16"), + B: T.Tensor((1024, 1024), "float16"), + matmul: T.Tensor((1024, 1024), "float32"), ): - A_flat = T.decl_buffer(1048576, "float16", data=A.data) - B_flat = T.decl_buffer(1048576, "float16", data=B.data) - matmul_flat = T.decl_buffer(1048576, data=matmul.data) + A_flat = T.decl_tensor(1048576, "float16", data=A.data) + B_flat = T.decl_tensor(1048576, "float16", data=B.data) + matmul_flat = T.decl_tensor(1048576, data=matmul.data) threadIdx_x = T.launch_thread("threadIdx.x", 16) - C_local = T.alloc_buffer((1,), "float32", scope="local") - A_sh = T.alloc_buffer((256,), "float16", scope="shared.dyn") - B_sh = T.alloc_buffer((256,), "float16", scope="shared.dyn") - C_sh = T.alloc_buffer((256,), "float32", scope="shared.dyn") + C_local = T.alloc_tensor((1,), "float32", scope="local") + A_sh = T.alloc_tensor((256,), "float16", scope="shared.dyn") + B_sh = T.alloc_tensor((256,), "float16", scope="shared.dyn") + C_sh = T.alloc_tensor((256,), "float32", scope="shared.dyn") threadIdx_y = T.launch_thread("threadIdx.y", 16) blockIdx_x = T.launch_thread("blockIdx.x", 64) blockIdx_y = T.launch_thread("blockIdx.y", 64) @@ -195,7 +195,7 @@ def main( After = transform(Before) script = After["main"].script() # Verify merged allocation: one shared.dyn buffer of 1024 bytes - assert "alloc_buffer((1024,)" in script + assert "alloc_tensor((1024,)" in script assert '"uint8"' in script assert '"shared.dyn"' in script assert "tvm_storage_sync" in script @@ -211,14 +211,14 @@ class Before: @Ts.prim_func def main(): threadIdx_x = T.launch_thread("threadIdx.x", 128) - A_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") - B_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") + A_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") + B_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") B_sh[threadIdx_x] = A_sh[threadIdx_x] After = transform(Before) script = After["main"].script() # Verify merged allocation: 1024 bytes (128*4 + 128*4) - assert "alloc_buffer((1024,)" in script + assert "alloc_tensor((1024,)" in script assert '"uint8"' in script assert '"shared.dyn"' in script # Verify offset indexing @@ -234,15 +234,15 @@ class Before: @Ts.prim_func def main(): threadIdx_x = T.launch_thread("threadIdx.x", 128) - A_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") - B_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") + A_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") + B_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") A_sh[threadIdx_x] = 0 B_sh[threadIdx_x] = 0 After = transform(Before) script = After["main"].script() # Verify merged allocation: 512 bytes (128*4, reusable) - assert "alloc_buffer((512,)" in script + assert "alloc_tensor((512,)" in script assert '"uint8"' in script assert '"shared.dyn"' in script @@ -254,10 +254,10 @@ def test_async_copy(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def main(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): threadIdx_x = T.launch_thread("threadIdx.x", 128) - A_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") - B_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") + A_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") + B_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") T.s_tir.cp_async_raw("float32", A_sh.data, threadIdx_x, A.data, threadIdx_x, 512) T.s_tir.cp_async_raw("float32", B_sh.data, threadIdx_x, B.data, threadIdx_x, 512) @@ -273,11 +273,11 @@ def main(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): # offsets remain in float32 elements and are scaled by the intrinsic # lowering, rather than being pre-scaled as byte offsets here. assert ( - 'A_sh = T.decl_buffer((128,), "float32", data=buf_dyn_shmem.data, ' + 'A_sh = T.decl_tensor((128,), "float32", data=buf_dyn_shmem.data, ' 'scope="shared.dyn")' in script ) assert ( - 'B_sh = T.decl_buffer((128,), "float32", data=buf_dyn_shmem.data, ' + 'B_sh = T.decl_tensor((128,), "float32", data=buf_dyn_shmem.data, ' 'scope="shared.dyn")' in script ) assert ( @@ -297,19 +297,19 @@ def test_decl_buffer_alias_extends_allocation_lifetime(): @I.ir_module class Before: @Ts.prim_func - def main(C: T.Buffer((128,), "float32")): + def main(C: T.Tensor((128,), "float32")): threadIdx_x = T.launch_thread("threadIdx.x", 128) - A_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") - B_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") - A_view = T.decl_buffer((128,), "float32", data=A_sh.data, scope="shared.dyn") - B_view = T.decl_buffer((128,), "float32", data=B_sh.data, scope="shared.dyn") + A_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") + B_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") + A_view = T.decl_tensor((128,), "float32", data=A_sh.data, scope="shared.dyn") + B_view = T.decl_tensor((128,), "float32", data=B_sh.data, scope="shared.dyn") A_view[threadIdx_x] = T.float32(1) B_view[threadIdx_x] = T.float32(2) C[threadIdx_x] = A_view[threadIdx_x] + B_view[threadIdx_x] After = transform(Before) script = After["main"].script() - assert 'alloc_buffer((1024,), "uint8", scope="shared.dyn")' in script + assert 'alloc_tensor((1024,), "uint8", scope="shared.dyn")' in script assert "B_view[threadIdx_x + 128]" in script assert "A_view[threadIdx_x]" in script @@ -328,17 +328,17 @@ def test_multi_thread_extent_blocks(): class Before: @Ts.prim_func(check_well_formed=False) def main( - X: T.Buffer((128,), "float32"), - Y: T.Buffer((128,), "float32"), + X: T.Tensor((128,), "float32"), + Y: T.Tensor((128,), "float32"), ): - X_flat = T.decl_buffer(128, data=X.data) - Y_flat = T.decl_buffer(128, data=Y.data) + X_flat = T.decl_tensor(128, data=X.data) + Y_flat = T.decl_tensor(128, data=Y.data) # First kernel launch tx0 = T.env_thread("threadIdx.x") with T.attr(tx0, "thread_extent", 128): - A_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") - B_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") + A_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") + B_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") A_sh[tx0] = X_flat[tx0] B_sh[tx0] = A_sh[tx0] X_flat[tx0] = B_sh[tx0] @@ -346,8 +346,8 @@ def main( # Second kernel launch — must NOT see kernel #0's merged buffer. tx1 = T.env_thread("threadIdx.x") with T.attr(tx1, "thread_extent", 128): - C_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") - D_sh = T.alloc_buffer((128,), "float32", scope="shared.dyn") + C_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") + D_sh = T.alloc_tensor((128,), "float32", scope="shared.dyn") C_sh[tx1] = Y_flat[tx1] D_sh[tx1] = C_sh[tx1] Y_flat[tx1] = D_sh[tx1] @@ -361,8 +361,8 @@ def main( assert script.count("shared.dyn") >= 2, ( "Expected at least two shared.dyn allocations (one per kernel)" ) - assert script.count("alloc_buffer") >= 2, ( - "Expected at least two alloc_buffer nodes (one merged buf per kernel)" + assert script.count("alloc_tensor") >= 2, ( + "Expected at least two alloc_tensor nodes (one merged buf per kernel)" ) # Both thread_extent blocks must contain their own merged buffer — diff --git a/tests/python/s_tir/transform/test_s_tir_transform_mma_buffer_layout.py b/tests/python/s_tir/transform/test_s_tir_transform_mma_buffer_layout.py index d37eafa7d196..946f11ce29c4 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_mma_buffer_layout.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_mma_buffer_layout.py @@ -27,7 +27,7 @@ ) @pytest.mark.parametrize("access_kind", ["load", "store"]) def test_explicit_matrix_ab_access_is_rejected(scope, shape, access_kind): - buffer = tirx.decl_buffer(shape, "float32", scope=scope) + buffer = tirx.decl_tensor(shape, "float32", scope=scope) if access_kind == "load": body = tirx.Evaluate(tirx.BufferLoad(buffer, [0, 0])) else: diff --git a/tests/python/s_tir/transform/test_s_tir_transform_plan_update_buffer_allocation_location.py b/tests/python/s_tir/transform/test_s_tir_transform_plan_update_buffer_allocation_location.py index 66650165ba81..c9cbf600ffa9 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_plan_update_buffer_allocation_location.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_plan_update_buffer_allocation_location.py @@ -33,7 +33,7 @@ def _check(original, transformed): @Ts.prim_func -def element_func(A: T.Buffer((16, 16)), C: T.Buffer((16, 16))) -> None: +def element_func(A: T.Tensor((16, 16)), C: T.Tensor((16, 16))) -> None: B = Ts.sblock_alloc_buffer((16, 16)) for i0 in range(0, 16): for j0 in range(0, 16): @@ -47,7 +47,7 @@ def element_func(A: T.Buffer((16, 16)), C: T.Buffer((16, 16))) -> None: @Ts.prim_func -def transformed_element_func(A: T.Buffer([16, 16]), C: T.Buffer([16, 16])) -> None: +def transformed_element_func(A: T.Tensor([16, 16]), C: T.Tensor([16, 16])) -> None: for i_0 in range(0, 16): with Ts.sblock(): Ts.reads([A[i_0, 0:16]]) @@ -158,7 +158,7 @@ def transformed_match_buffer_func() -> None: @Ts.prim_func -def opaque_access(A: T.Buffer([1024]), B: T.Buffer([1024])) -> None: +def opaque_access(A: T.Tensor([1024]), B: T.Tensor([1024])) -> None: A_cache = Ts.sblock_alloc_buffer([1024]) for i in T.serial(0, 8): with Ts.sblock(): @@ -188,7 +188,7 @@ def opaque_access(A: T.Buffer([1024]), B: T.Buffer([1024])) -> None: @Ts.prim_func -def transformed_opaque_access(A: T.Buffer([1024]), B: T.Buffer([1024])) -> None: +def transformed_opaque_access(A: T.Tensor([1024]), B: T.Tensor([1024])) -> None: for i in T.serial(0, 8): with Ts.sblock(): vi = Ts.axis.S(8, i) @@ -234,7 +234,7 @@ def test_loop_carried_dependency(): and the allocate buffer should keep the order.""" @Ts.prim_func - def before(A: T.Buffer((8, 8, 8), "int32"), B: T.Buffer((8, 8, 8), "int32")): + def before(A: T.Tensor((8, 8, 8), "int32"), B: T.Tensor((8, 8, 8), "int32")): C = Ts.sblock_alloc_buffer([8, 8, 8], dtype="int32") D = Ts.sblock_alloc_buffer([8, 8, 8], dtype="int32") for i in T.serial(8): @@ -258,7 +258,7 @@ def before(A: T.Buffer((8, 8, 8), "int32"), B: T.Buffer((8, 8, 8), "int32")): ) @Ts.prim_func - def after(A: T.Buffer((8, 8, 8), "int32"), B: T.Buffer((8, 8, 8), "int32")) -> None: + def after(A: T.Tensor((8, 8, 8), "int32"), B: T.Tensor((8, 8, 8), "int32")) -> None: for i in T.serial(8): with Ts.sblock(): Ts.reads(A[i, 0:8, 0:8]) @@ -292,7 +292,7 @@ def test_1D_cascade_op_rolling_buffer(): which is marked as opaque in consumer block's iter mappings.""" @Ts.prim_func - def before(A: T.Buffer((4, 16), "int32"), C: T.Buffer((4, 8), "int32")): + def before(A: T.Tensor((4, 16), "int32"), C: T.Tensor((4, 8), "int32")): B = Ts.sblock_alloc_buffer((4, 6), "int32") for c in T.serial(4): for i in T.serial(0, 2): @@ -318,7 +318,7 @@ def before(A: T.Buffer((4, 16), "int32"), C: T.Buffer((4, 8), "int32")): ) @Ts.prim_func - def after(A: T.Buffer((4, 16), "int32"), C: T.Buffer((4, 8), "int32")): + def after(A: T.Tensor((4, 16), "int32"), C: T.Tensor((4, 8), "int32")): for c in T.serial(4): with Ts.sblock(): Ts.reads(A[c, 0:12], C[c, 0:8]) @@ -350,14 +350,14 @@ def test_buffer_conditional_lowering(): Confirm that the `tirx.PlanAndUpdateBufferAllocationLocation` pass leaves (Buffer nodes corresponding to pointer-typed PrimFunc arguments) - unchanged, rather than lowering them to `reads`, `writes`, and `alloc_buffer` nodes. + unchanged, rather than lowering them to `reads`, `writes`, and `alloc_tensor` nodes. """ @Ts.prim_func def before(A: T.handle("float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i in range(1): - A_1 = T.decl_buffer((1,), data=A) + A_1 = T.decl_tensor((1,), data=A) A_1[i] = 0 after = before @@ -370,7 +370,7 @@ def test_dltensor_buffer_is_unlowered(): Confirm that the `tirx.PlanAndUpdateBufferAllocationLocation` pass leaves (Buffer nodes corresponding to PrimFunc DLTensor arguments) unchanged, rather than lowering them to `reads`, `writes`, and - `alloc_buffer` nodes. + `alloc_tensor` nodes. """ @Ts.prim_func @@ -383,14 +383,14 @@ def before(dlpack_handle: T.handle, axis: T.int64) -> T.int64: shape_ptr: T.let[T.handle("int64")] = T.tvm_struct_get( dlpack_handle, 0, 3, dtype=T.handle("int64").ty ) - shape = T.decl_buffer(ndim, "int64", data=shape_ptr) - product = T.decl_buffer([], "int64") + shape = T.decl_tensor(ndim, "int64", data=shape_ptr) + product = T.decl_tensor([], "int64") product[()] = 1 for dim in range(axis + 1, ndim): product[()] = product[()] * shape[dim] return product[()] else: - strides = T.decl_buffer(ndim, "int64", data=stride_ptr) + strides = T.decl_tensor(ndim, "int64", data=stride_ptr) stride: T.int64 = strides[axis] return stride @@ -402,7 +402,7 @@ def test_reduce_buffer_dominate_reduce_loops(): """Reduction write buffer allocation should dominate all reduce loops""" @Ts.prim_func - def before(x: T.Buffer((256, 256, 256), "float32"), x_red: T.Buffer((256, 256), "float32")): + def before(x: T.Tensor((256, 256, 256), "float32"), x_red: T.Tensor((256, 256), "float32")): x_red_ = Ts.sblock_alloc_buffer((256, 256)) for ax0_0, k1_0, ax1_0 in T.grid(4, 4, 4): for ax0_1, k1_1, ax1_1 in T.grid(64, 64, 64): @@ -420,7 +420,7 @@ def before(x: T.Buffer((256, 256, 256), "float32"), x_red: T.Buffer((256, 256), x_red[v0, v1] = x_red_[v0, v1] @Ts.prim_func - def after(x: T.Buffer((256, 256, 256), "float32"), x_red: T.Buffer((256, 256), "float32")): + def after(x: T.Tensor((256, 256, 256), "float32"), x_red: T.Tensor((256, 256), "float32")): for ax0_0 in range(4): with Ts.sblock(""): Ts.reads(x[ax0_0 * 64 : ax0_0 * 64 + 64, 0:256, 0:256]) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_remove_undef.py b/tests/python/s_tir/transform/test_s_tir_transform_remove_undef.py index 20682648a6f2..8162e0835f9b 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_remove_undef.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_remove_undef.py @@ -30,13 +30,13 @@ def test_remove_store_undef(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(1, "int32")): + def main(A: T.Tensor(1, "int32")): A[0] = T.undef() @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer(1, "int32")): + def main(A: T.Tensor(1, "int32")): T.evaluate(0) After = tvm.s_tir.transform.RemoveStoreUndef()(Before) @@ -49,13 +49,13 @@ def test_remove_store_undef_expression(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(1, "int32")): + def main(A: T.Tensor(1, "int32")): A[0] = 1 + T.undef() @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer(1, "int32")): + def main(A: T.Tensor(1, "int32")): T.evaluate(0) After = tvm.s_tir.transform.RemoveStoreUndef()(Before) @@ -68,7 +68,7 @@ def test_keep_other_call_nodes(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(1, "int32"), n: T.int32): + def main(A: T.Tensor(1, "int32"), n: T.int32): A[0] = T.shift_left(n, 1) Expected = Before @@ -83,14 +83,14 @@ def test_remove_let_undef(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(1, "int32")): + def main(A: T.Tensor(1, "int32")): val: T.let[T.int32] = T.undef() A[0] = val @I.ir_module class Expected: @Ts.prim_func - def main(A: T.Buffer(1, "int32")): + def main(A: T.Tensor(1, "int32")): T.evaluate(0) After = tvm.s_tir.transform.RemoveStoreUndef()(Before) @@ -103,7 +103,7 @@ def test_raise_error_for_undef_as_store_indices(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(1, "int32")): + def main(A: T.Tensor(1, "int32")): val: T.let[T.int32] = T.undef() A[val] = 5 @@ -121,7 +121,7 @@ def test_raise_error_for_undef_as_load_indices(): @I.ir_module class Before: @Ts.prim_func - def main(A: T.Buffer(1, "int32"), B: T.Buffer(1, "int32")): + def main(A: T.Tensor(1, "int32"), B: T.Tensor(1, "int32")): B[0] = A[T.undef()] with pytest.raises(RuntimeError): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_remove_weight_layout_rewrite_block.py b/tests/python/s_tir/transform/test_s_tir_transform_remove_weight_layout_rewrite_block.py index 4f8cc3f12e45..34883e4d67a7 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_remove_weight_layout_rewrite_block.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_remove_weight_layout_rewrite_block.py @@ -37,9 +37,9 @@ def _check(before, expect): def test_matmul(): @Ts.prim_func def before( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 16), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 16), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: T.func_attr({"layout_free_buffers": [1]}) B_ = Ts.sblock_alloc_buffer([16, 4, 4], dtype="float32") @@ -63,9 +63,9 @@ def before( @Ts.prim_func def after( - A: T.Buffer((16, 16), "float32"), - B: T.Buffer((16, 4, 4), "float32"), - C: T.Buffer((16, 16), "float32"), + A: T.Tensor((16, 16), "float32"), + B: T.Tensor((16, 4, 4), "float32"), + C: T.Tensor((16, 16), "float32"), ) -> None: T.func_attr({"layout_free_buffers": [1]}) for i0_o, i1_o in T.grid(16, 16): diff --git a/tests/python/s_tir/transform/test_s_tir_transform_renormalize_split_pattern.py b/tests/python/s_tir/transform/test_s_tir_transform_renormalize_split_pattern.py index cd98b4076bf8..40dba38499ec 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_renormalize_split_pattern.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_renormalize_split_pattern.py @@ -28,20 +28,20 @@ @tvm.script.ir_module class Before: @Ts.prim_func - def main(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Buffer((1, 8, 8, 256), "float32")) -> None: + def main(inputs: T.Tensor((1, 4, 4, 512), "float32"), weight: T.Tensor((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Tensor((1, 8, 8, 256), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - inputs_flat = T.decl_buffer([8192], dtype="float32", data=inputs.data) - weight_flat = T.decl_buffer([2097152], dtype="float32", data=weight.data) - conv2d_transpose_nhwc_flat = T.decl_buffer([16384], dtype="float32", data=conv2d_transpose_nhwc.data) + inputs_flat = T.decl_tensor([8192], dtype="float32", data=inputs.data) + weight_flat = T.decl_tensor([2097152], dtype="float32", data=weight.data) + conv2d_transpose_nhwc_flat = T.decl_tensor([16384], dtype="float32", data=conv2d_transpose_nhwc.data) # var definition threadIdx_x = T.env_thread("threadIdx.x") blockIdx_x = T.env_thread("blockIdx.x") # body T.launch_thread(blockIdx_x, 64) - conv2d_transpose_nhwc_local = T.decl_buffer([8], "float32", scope="local") - PadInput_shared = T.decl_buffer([768], "float32", scope="shared") - weight_shared = T.decl_buffer([4096], "float32", scope="shared") + conv2d_transpose_nhwc_local = T.decl_tensor([8], "float32", scope="local") + PadInput_shared = T.decl_tensor([768], "float32", scope="shared") + weight_shared = T.decl_tensor([4096], "float32", scope="shared") T.launch_thread(threadIdx_x, 32) for i2_3_init, i1_4_init, i2_4_init in T.grid(2, 2, 2): conv2d_transpose_nhwc_local[i1_4_init * 4 + i2_3_init * 2 + i2_4_init] = T.float32(0) @@ -59,20 +59,20 @@ def main(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 51 @tvm.script.ir_module class After: @Ts.prim_func - def main(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Buffer((1, 8, 8, 256), "float32")) -> None: + def main(inputs: T.Tensor((1, 4, 4, 512), "float32"), weight: T.Tensor((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Tensor((1, 8, 8, 256), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - inputs_flat = T.decl_buffer([8192], dtype="float32", data=inputs.data) - weight_flat = T.decl_buffer([2097152], dtype="float32", data=weight.data) - conv2d_transpose_nhwc_flat = T.decl_buffer([16384], dtype="float32", data=conv2d_transpose_nhwc.data) + inputs_flat = T.decl_tensor([8192], dtype="float32", data=inputs.data) + weight_flat = T.decl_tensor([2097152], dtype="float32", data=weight.data) + conv2d_transpose_nhwc_flat = T.decl_tensor([16384], dtype="float32", data=conv2d_transpose_nhwc.data) # var definition threadIdx_x = T.env_thread("threadIdx.x") blockIdx_x = T.env_thread("blockIdx.x") # body T.launch_thread(blockIdx_x, 64) - conv2d_transpose_nhwc_local = T.decl_buffer([8], "float32", scope="local") - PadInput_shared = T.decl_buffer([768], "float32", scope="shared") - weight_shared = T.decl_buffer([4096], "float32", scope="shared") + conv2d_transpose_nhwc_local = T.decl_tensor([8], "float32", scope="local") + PadInput_shared = T.decl_tensor([768], "float32", scope="shared") + weight_shared = T.decl_tensor([4096], "float32", scope="shared") T.launch_thread(threadIdx_x, 32) for i2_3_init, i1_4_init, i2_4_init in T.grid(2, 2, 2): conv2d_transpose_nhwc_local[i1_4_init * 4 + i2_3_init * 2 + i2_4_init] = T.float32(0) @@ -90,20 +90,20 @@ def main(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 51 @tvm.script.ir_module class After_simplified: @Ts.prim_func - def main(inputs: T.Buffer((1, 4, 4, 512), "float32"), weight: T.Buffer((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Buffer((1, 8, 8, 256), "float32")) -> None: + def main(inputs: T.Tensor((1, 4, 4, 512), "float32"), weight: T.Tensor((4, 4, 512, 256), "float32"), conv2d_transpose_nhwc: T.Tensor((1, 8, 8, 256), "float32")) -> None: # function attr dict T.func_attr({"global_symbol": "main", "tirx.noalias": True}) # var definition threadIdx_x = T.env_thread("threadIdx.x") blockIdx_x = T.env_thread("blockIdx.x") - inputs_flat = T.decl_buffer([8192], dtype="float32", data=inputs.data) - weight_flat = T.decl_buffer([2097152], dtype="float32", data=weight.data) - conv2d_transpose_nhwc_flat = T.decl_buffer([16384], dtype="float32", data=conv2d_transpose_nhwc.data) + inputs_flat = T.decl_tensor([8192], dtype="float32", data=inputs.data) + weight_flat = T.decl_tensor([2097152], dtype="float32", data=weight.data) + conv2d_transpose_nhwc_flat = T.decl_tensor([16384], dtype="float32", data=conv2d_transpose_nhwc.data) # body T.launch_thread(blockIdx_x, 64) - conv2d_transpose_nhwc_local = T.decl_buffer([8], "float32", scope="local") - PadInput_shared = T.decl_buffer([768], "float32", scope="shared") - weight_shared = T.decl_buffer([4096], "float32", scope="shared") + conv2d_transpose_nhwc_local = T.decl_tensor([8], "float32", scope="local") + PadInput_shared = T.decl_tensor([768], "float32", scope="shared") + weight_shared = T.decl_tensor([4096], "float32", scope="shared") T.launch_thread(threadIdx_x, 32) for i2_3_init, i1_4_init, i2_4_init in T.grid(2, 2, 2): conv2d_transpose_nhwc_local[i1_4_init * 4 + i2_3_init * 2 + i2_4_init] = T.float32(0) diff --git a/tests/python/s_tir/transform/test_s_tir_transform_rewrite_unsafe_select.py b/tests/python/s_tir/transform/test_s_tir_transform_rewrite_unsafe_select.py index 220e1c3d84e7..651a862aefe5 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_rewrite_unsafe_select.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_rewrite_unsafe_select.py @@ -27,7 +27,7 @@ def test_rewrite_Select(): class ModuleY: @Ts.prim_func def main(i: T.int32): - A = T.alloc_buffer((100,)) + A = T.alloc_tensor((100,)) T.evaluate(T.Select(i > 1, A[i - 1], T.float32(1.0))) yy = tvm.s_tir.transform.RewriteUnsafeSelect()(ModuleY)["main"].body.seq[-1].value @@ -36,7 +36,7 @@ def main(i: T.int32): class ModuleZ: @Ts.prim_func def main(i: T.int32): - A = T.alloc_buffer((100,)) + A = T.alloc_tensor((100,)) T.evaluate( T.Select( T.Select(i > 1, A[i - 1], T.float32(1.0)) > T.float32(0.0), A[i], T.float32(0.1) @@ -49,7 +49,7 @@ def main(i: T.int32): class ModuleA: @Ts.prim_func def main(i: T.int32): - A = T.alloc_buffer((100,)) + A = T.alloc_tensor((100,)) # Inline y and z to avoid Let bindings - outer Select condition is safe (no buffer access) T.evaluate( T.Select( diff --git a/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py b/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py index 5d5a52e5f629..bb0291295416 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_thread_sync.py @@ -42,7 +42,7 @@ def run_passes(func: tvm.tirx.PrimFunc): @pytest.mark.skipif(not env.has_cuda(), reason="need cuda") def test_sync_read_thread_id_independent_location(): @Ts.prim_func(check_well_formed=False) - def func(p0_arg: T.Buffer((1, 2, 1, 1), "float32"), p1: T.Buffer(2, "float32")) -> None: + def func(p0_arg: T.Tensor((1, 2, 1, 1), "float32"), p1: T.Tensor(2, "float32")) -> None: threadIdx_x = T.env_thread("threadIdx.x") blockIdx_x = T.env_thread("blockIdx.x") p0 = T.buffer([2], dtype="float32", data=p0_arg.data) @@ -64,8 +64,8 @@ def func(p0_arg: T.Buffer((1, 2, 1, 1), "float32"), p1: T.Buffer(2, "float32")) def test_sync_inside_condition(): @Ts.prim_func - def func1(A: T.Buffer((4, 4), "float32")) -> None: - A_shared = T.alloc_buffer((4, 4), "float32", scope="shared") + def func1(A: T.Tensor((4, 4), "float32")) -> None: + A_shared = T.alloc_tensor((4, 4), "float32", scope="shared") bx = T.launch_thread("blockIdx.x", 1) tx = T.launch_thread("threadIdx.x", 32) if A[0, 0] > 1.0: @@ -75,8 +75,8 @@ def func1(A: T.Buffer((4, 4), "float32")) -> None: A[i, j] = A_shared[i, j] + 1.0 @Ts.prim_func - def func2(A: T.Buffer((4, 4), "float32")) -> None: - A_shared = T.alloc_buffer((4, 4), "float32", scope="shared") + def func2(A: T.Tensor((4, 4), "float32")) -> None: + A_shared = T.alloc_tensor((4, 4), "float32", scope="shared") bx = T.launch_thread("blockIdx.x", 1) tx = T.launch_thread("threadIdx.x", 32) if T.tvm_thread_invariant(A[0, 0] > 1.0): @@ -86,8 +86,8 @@ def func2(A: T.Buffer((4, 4), "float32")) -> None: A[i, j] = A_shared[i, j] + 1.0 @Ts.prim_func - def func3(A: T.Buffer((4, 4), "float32")) -> None: - A_shared = T.alloc_buffer((4, 4), "float32", scope="shared") + def func3(A: T.Tensor((4, 4), "float32")) -> None: + A_shared = T.alloc_tensor((4, 4), "float32", scope="shared") bx = T.launch_thread("blockIdx.x", 1) tx = T.launch_thread("threadIdx.x", 32) while T.tvm_thread_invariant(A[0, 0] > 1.0): @@ -106,38 +106,38 @@ def func3(A: T.Buffer((4, 4), "float32")) -> None: def test_sync_shared_dyn(): @Ts.prim_func(private=True) - def func(A: T.Buffer((4, 4), "float32"), E: T.Buffer((4, 4), "float32")): + def func(A: T.Tensor((4, 4), "float32"), E: T.Tensor((4, 4), "float32")): blockIdx_x = T.launch_thread("blockIdx.x", 1) - B = T.alloc_buffer((24,), "float32", scope="shared.dyn") - C = T.alloc_buffer((1,), "float32", scope="local") - D = T.alloc_buffer((16,), "float32", scope="shared.dyn") + B = T.alloc_tensor((24,), "float32", scope="shared.dyn") + C = T.alloc_tensor((1,), "float32", scope="local") + D = T.alloc_tensor((16,), "float32", scope="shared.dyn") threadIdx_x = T.launch_thread("threadIdx.x", 16) - B_1 = T.decl_buffer((24,), data=B.data, scope="shared.dyn") - A_1 = T.decl_buffer((16,), data=A.data) + B_1 = T.decl_tensor((24,), data=B.data, scope="shared.dyn") + A_1 = T.decl_tensor((16,), data=A.data) B_1[threadIdx_x // 4 * 6 + threadIdx_x % 4] = A_1[threadIdx_x] - C_1 = T.decl_buffer((1,), data=C.data, scope="local") + C_1 = T.decl_tensor((1,), data=C.data, scope="local") C_1[0] = B_1[threadIdx_x // 4 * 6 + threadIdx_x % 4] - D_1 = T.decl_buffer((16,), data=D.data, scope="shared.dyn") + D_1 = T.decl_tensor((16,), data=D.data, scope="shared.dyn") D_1[threadIdx_x] = C_1[0] - E_1 = T.decl_buffer((16,), data=E.data) + E_1 = T.decl_tensor((16,), data=E.data) E_1[threadIdx_x] = D_1[threadIdx_x] @Ts.prim_func(private=True) - def expected(A: T.Buffer((4, 4), "float32"), E: T.Buffer((4, 4), "float32")): + def expected(A: T.Tensor((4, 4), "float32"), E: T.Tensor((4, 4), "float32")): blockIdx_x = T.launch_thread("blockIdx.x", 1) - B_1 = T.alloc_buffer((24,), "float32", scope="shared.dyn") - C_1 = T.alloc_buffer((1,), "float32", scope="local") - D_1 = T.alloc_buffer((16,), "float32", scope="shared.dyn") + B_1 = T.alloc_tensor((24,), "float32", scope="shared.dyn") + C_1 = T.alloc_tensor((1,), "float32", scope="local") + D_1 = T.alloc_tensor((16,), "float32", scope="shared.dyn") threadIdx_x = T.launch_thread("threadIdx.x", 16) - B_1_1 = T.decl_buffer((24,), data=B_1.data, scope="shared.dyn") - A_1 = T.decl_buffer((16,), data=A.data) + B_1_1 = T.decl_tensor((24,), data=B_1.data, scope="shared.dyn") + A_1 = T.decl_tensor((16,), data=A.data) B_1_1[threadIdx_x // 4 * 6 + threadIdx_x % 4] = A_1[threadIdx_x] - C_1_1 = T.decl_buffer((1,), data=C_1.data, scope="local") + C_1_1 = T.decl_tensor((1,), data=C_1.data, scope="local") C_1_1[0] = B_1_1[threadIdx_x // 4 * 6 + threadIdx_x % 4] - D_1_1 = T.decl_buffer((16,), data=D_1.data, scope="shared.dyn") + D_1_1 = T.decl_tensor((16,), data=D_1.data, scope="shared.dyn") T.evaluate(T.call_intrin("int32", "tirx.tvm_storage_sync", "shared.dyn")) D_1_1[threadIdx_x] = C_1_1[0] - E_1 = T.decl_buffer((16,), data=E.data) + E_1 = T.decl_tensor((16,), data=E.data) E_1[threadIdx_x] = D_1_1[threadIdx_x] mod = tvm.IRModule({"main": func}) @@ -147,13 +147,13 @@ def expected(A: T.Buffer((4, 4), "float32"), E: T.Buffer((4, 4), "float32")): def test_sync_shared_aliasing_buffer_views(): @Ts.prim_func(private=True) - def func(A: T.Buffer((64,), "float32")): + def func(A: T.Tensor((64,), "float32")): blockIdx_x = T.launch_thread("blockIdx.x", 1) - shared_storage = T.alloc_buffer((32,), "float16", scope="shared") - local = T.alloc_buffer((1,), "float32", scope="local") + shared_storage = T.alloc_tensor((32,), "float16", scope="shared") + local = T.alloc_tensor((1,), "float32", scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 32) - shared_half = T.decl_buffer((32,), "float16", data=shared_storage.data, scope="shared") - shared_float = T.decl_buffer((16,), "float32", data=shared_storage.data, scope="shared") + shared_half = T.decl_tensor((32,), "float16", data=shared_storage.data, scope="shared") + shared_float = T.decl_tensor((16,), "float32", data=shared_storage.data, scope="shared") for i in range(2): shared_half[threadIdx_x] = T.Cast("float16", A[i * 32 + threadIdx_x]) T.tvm_storage_sync("shared") @@ -172,16 +172,16 @@ def func(A: T.Buffer((64,), "float32")): @pytest.mark.skipif(not env.has_cuda(), reason="need cuda") def test_sync_bind(): @Ts.prim_func(private=True) - def func(A: T.Buffer((16 * 512), "float32")): + def func(A: T.Tensor((16 * 512), "float32")): blockIdx_x = T.launch_thread("blockIdx.x", 16) - A_shared = T.alloc_buffer((512,), "float32", scope="shared") - in_thread_A_temp = T.alloc_buffer((1,), "float32", scope="local") - cross_thread_A_temp = T.alloc_buffer((1,), "float32", scope="local") + A_shared = T.alloc_tensor((512,), "float32", scope="shared") + in_thread_A_temp = T.alloc_tensor((1,), "float32", scope="local") + cross_thread_A_temp = T.alloc_tensor((1,), "float32", scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 128) - A_shared_1 = T.decl_buffer((512,), data=A_shared.data, scope="shared") + A_shared_1 = T.decl_tensor((512,), data=A_shared.data, scope="shared") for ax0 in range(512): A_shared_1[ax0] = A[blockIdx_x * 512 + ax0] - in_thread_A_temp_1 = T.decl_buffer((1,), data=in_thread_A_temp.data, scope="local") + in_thread_A_temp_1 = T.decl_tensor((1,), data=in_thread_A_temp.data, scope="local") in_thread_A_temp_1[0] = T.float32(0) A_temp_1 = T.bind(in_thread_A_temp_1[0] + A_shared_1[threadIdx_x]) in_thread_A_temp_1[0] = A_temp_1 @@ -191,7 +191,7 @@ def func(A: T.Buffer((16 * 512), "float32")): in_thread_A_temp_1[0] = A_temp_3 A_temp_4 = T.bind(in_thread_A_temp_1[0] + A_shared_1[threadIdx_x + 384]) in_thread_A_temp_1[0] = A_temp_4 - cross_thread_A_temp_1 = T.decl_buffer((1,), data=cross_thread_A_temp.data, scope="local") + cross_thread_A_temp_1 = T.decl_tensor((1,), data=cross_thread_A_temp.data, scope="local") with T.attr( T.comm_reducer(lambda x0, y0: x0 + y0, [T.float32(0)]), "reduce_scope", @@ -206,16 +206,16 @@ def func(A: T.Buffer((16 * 512), "float32")): ) @Ts.prim_func(private=True) - def expected(A: T.Buffer((8192,), "float32")): + def expected(A: T.Tensor((8192,), "float32")): blockIdx_x = T.launch_thread("blockIdx.x", 16) - A_shared_1 = T.alloc_buffer((512,), "float32", scope="shared") - in_thread_A_temp_1 = T.alloc_buffer((1,), "float32", scope="local") - cross_thread_A_temp_1 = T.alloc_buffer((1,), "float32", scope="local") + A_shared_1 = T.alloc_tensor((512,), "float32", scope="shared") + in_thread_A_temp_1 = T.alloc_tensor((1,), "float32", scope="local") + cross_thread_A_temp_1 = T.alloc_tensor((1,), "float32", scope="local") threadIdx_x = T.launch_thread("threadIdx.x", 128) - A_shared_1_1 = T.decl_buffer((512,), data=A_shared_1.data, scope="shared") + A_shared_1_1 = T.decl_tensor((512,), data=A_shared_1.data, scope="shared") for ax0 in range(512): A_shared_1_1[ax0] = A[blockIdx_x * 512 + ax0] - in_thread_A_temp_1_1 = T.decl_buffer((1,), data=in_thread_A_temp_1.data, scope="local") + in_thread_A_temp_1_1 = T.decl_tensor((1,), data=in_thread_A_temp_1.data, scope="local") in_thread_A_temp_1_1[0] = T.float32(0) T.evaluate(T.call_intrin("int32", "tirx.tvm_storage_sync", "shared")) A_temp_1 = T.bind(in_thread_A_temp_1_1[0] + A_shared_1_1[threadIdx_x]) @@ -226,7 +226,7 @@ def expected(A: T.Buffer((8192,), "float32")): in_thread_A_temp_1_1[0] = A_temp_3 A_temp_4 = T.bind(in_thread_A_temp_1_1[0] + A_shared_1_1[threadIdx_x + 384]) in_thread_A_temp_1_1[0] = A_temp_4 - cross_thread_A_temp_1_1 = T.decl_buffer( + cross_thread_A_temp_1_1 = T.decl_tensor( (1,), data=cross_thread_A_temp_1.data, scope="local" ) T.attr( diff --git a/tests/python/s_tir/transform/test_s_tir_transform_unify_thread_binding.py b/tests/python/s_tir/transform/test_s_tir_transform_unify_thread_binding.py index b03703231180..ae3938153d61 100644 --- a/tests/python/s_tir/transform/test_s_tir_transform_unify_thread_binding.py +++ b/tests/python/s_tir/transform/test_s_tir_transform_unify_thread_binding.py @@ -43,7 +43,7 @@ def _check_fail(original): @Ts.prim_func def element_wise_thread_x( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i in T.thread_binding(0, 128, "blockIdx.x"): for j0_0 in T.thread_binding(0, 4, "threadIdx.x"): @@ -58,7 +58,7 @@ def element_wise_thread_x( @Ts.prim_func def unified_element_wise_thread_x( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for blockIdx_x in T.thread_binding(0, 128, "blockIdx.x"): for threadIdx_x in T.thread_binding(0, 4, "threadIdx.x"): @@ -76,9 +76,9 @@ def unified_element_wise_thread_x( @Ts.prim_func def element_wise_thread_x_different_dtype( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: for i in T.thread_binding(128, "blockIdx.x"): for j0_0 in T.thread_binding(4, "threadIdx.x"): @@ -93,9 +93,9 @@ def element_wise_thread_x_different_dtype( @Ts.prim_func def unified_element_wise_thread_x_different_dtype( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ) -> None: for blockIdx_x in T.thread_binding(128, "blockIdx.x"): for threadIdx_x in T.thread_binding(4, "threadIdx.x"): @@ -113,7 +113,7 @@ def unified_element_wise_thread_x_different_dtype( @Ts.prim_func def element_wise_env_thread_x( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: j1_0 = T.env_thread("threadIdx.x") j0_0 = T.env_thread("threadIdx.x") @@ -133,7 +133,7 @@ def element_wise_env_thread_x( @Ts.prim_func def unified_element_wise_env_thread_x( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for blockIdx_x in T.thread_binding(0, 128, "blockIdx.x"): for threadIdx_x in T.thread_binding(0, 4, "threadIdx.x"): @@ -150,7 +150,7 @@ def unified_element_wise_env_thread_x( @Ts.prim_func -def element_wise_vthread_x(A: T.Buffer([128, 128]), B: T.Buffer([128, 128])) -> None: +def element_wise_vthread_x(A: T.Tensor([128, 128]), B: T.Tensor([128, 128])) -> None: for i_0 in T.thread_binding(0, 2, "vthread.x"): for i_1 in T.thread_binding(0, 64, "threadIdx.x"): for j_0 in T.thread_binding(0, 2, "vthread.x"): @@ -160,7 +160,7 @@ def element_wise_vthread_x(A: T.Buffer([128, 128]), B: T.Buffer([128, 128])) -> @Ts.prim_func -def unified_element_wise_vthread_x(A: T.Buffer([128, 128]), B: T.Buffer([128, 128])) -> None: +def unified_element_wise_vthread_x(A: T.Tensor([128, 128]), B: T.Tensor([128, 128])) -> None: for vthread_x in T.thread_binding(0, 2, "vthread.x"): for threadIdx_x in T.thread_binding(0, 64, "threadIdx.x"): for j_1 in T.serial(0, 64): @@ -172,7 +172,7 @@ def unified_element_wise_vthread_x(A: T.Buffer([128, 128]), B: T.Buffer([128, 12 @Ts.prim_func def element_wise_two_thread_x_in_same_kernel_not_equal( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 64]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 64]) ) -> None: for i in T.thread_binding(0, 128, "blockIdx.x"): for j0 in T.thread_binding(0, 128, "threadIdx.x"): @@ -183,10 +183,10 @@ def element_wise_two_thread_x_in_same_kernel_not_equal( @Ts.prim_func def element_wise_kernels_with_different_size( - A: T.Buffer([128, 128]), - B: T.Buffer([128, 128]), - C: T.Buffer([256, 256]), - D: T.Buffer([256, 256]), + A: T.Tensor([128, 128]), + B: T.Tensor([128, 128]), + C: T.Tensor([256, 256]), + D: T.Tensor([256, 256]), ) -> None: for i0 in T.thread_binding(0, 128, "blockIdx.x"): for j0 in T.thread_binding(0, 128, "threadIdx.x"): @@ -198,10 +198,10 @@ def element_wise_kernels_with_different_size( @Ts.prim_func def unified_element_wise_kernels_with_different_size( - A: T.Buffer([128, 128]), - B: T.Buffer([128, 128]), - C: T.Buffer([256, 256]), - D: T.Buffer([256, 256]), + A: T.Tensor([128, 128]), + B: T.Tensor([128, 128]), + C: T.Tensor([256, 256]), + D: T.Tensor([256, 256]), ) -> None: for blockIdx_x in T.thread_binding(0, 128, "blockIdx.x"): for threadIdx_x in T.thread_binding(0, 128, "threadIdx.x"): @@ -213,7 +213,7 @@ def unified_element_wise_kernels_with_different_size( @Ts.prim_func def element_wise_implicit_block( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for i in T.thread_binding(0, 128, "threadIdx.y"): for j0_0 in T.thread_binding(0, 4, "threadIdx.x"): @@ -228,7 +228,7 @@ def element_wise_implicit_block( @Ts.prim_func def unified_element_wise_implicit_block( - A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128]) + A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128]) ) -> None: for blockIdx_x in T.thread_binding(0, 128, "threadIdx.y"): for threadIdx_x in T.thread_binding(0, 4, "threadIdx.x"): @@ -276,7 +276,7 @@ def test_implicit_block(): def test_inner_binding_with_annotation(): @Ts.prim_func - def inner_binding_with_annotation(A: T.Buffer((64,), "float32"), B: T.Buffer((64,), "float32")): + def inner_binding_with_annotation(A: T.Tensor((64,), "float32"), B: T.Tensor((64,), "float32")): for bx in T.thread_binding(32, "blockIdx.x"): for tx in T.thread_binding(2, "threadIdx.x", annotations={"my_annotation": 1}): with Ts.sblock("block"): @@ -285,7 +285,7 @@ def inner_binding_with_annotation(A: T.Buffer((64,), "float32"), B: T.Buffer((64 @Ts.prim_func def unified_inner_binding_with_annotation( - A: T.Buffer((64,), "float32"), B: T.Buffer((64,), "float32") + A: T.Tensor((64,), "float32"), B: T.Tensor((64,), "float32") ): for blockIdx_x in T.thread_binding(32, thread="blockIdx.x"): for threadIdx_x in T.thread_binding(2, thread="threadIdx.x"): diff --git a/tests/python/s_tir/transform/test_tirx_force_narrow_reject_blocks.py b/tests/python/s_tir/transform/test_tirx_force_narrow_reject_blocks.py index 0a47c67cc291..54d0fe15dffa 100644 --- a/tests/python/s_tir/transform/test_tirx_force_narrow_reject_blocks.py +++ b/tests/python/s_tir/transform/test_tirx_force_narrow_reject_blocks.py @@ -23,7 +23,7 @@ def test_reject_blocks(): @Ts.prim_func(private=True) - def before(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def before(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): for i in T.serial(0, T.int64(16)): for j in T.serial(0, T.int64(8)): with Ts.sblock(): diff --git a/tests/python/script/test_jit_specialization.py b/tests/python/script/test_jit_specialization.py index 1818b128cebc..870a76879e8f 100644 --- a/tests/python/script/test_jit_specialization.py +++ b/tests/python/script/test_jit_specialization.py @@ -197,12 +197,12 @@ def test_tirx_jit_specializes_captured_shape_and_value(): width = 4 @T.jit(private=True) - def fill(output: T.Buffer((width,), "int32"), *, value: T.constexpr): + def fill(output: T.Tensor((width,), "int32"), *, value: T.constexpr): for index in range(width): output[index] = value @T.prim_func(private=True) - def expected(output: T.Buffer((4,), "int32")): + def expected(output: T.Tensor((4,), "int32")): for index in range(4): output[index] = 3 @@ -262,8 +262,8 @@ def function(): @pytest.mark.parametrize("postponed", [False, True]) def test_jit_initial_alias_preserves_annotation_forms(tmp_path, postponed): source = """@Script.jit(private=True) -def kernel(output: Script.Buffer((1,), "int32"), *, value: Script.constexpr, - optional: Script.Optional(Script.Buffer((1,), "int32"))): +def kernel(output: Script.Tensor((1,), "int32"), *, value: Script.constexpr, + optional: Script.Optional(Script.Tensor((1,), "int32"))): output[0] = value """ if postponed: diff --git a/tests/python/sym/test_sym_deduce_bound.py b/tests/python/sym/test_sym_deduce_bound.py index e76c8551209f..e5b99f35c23a 100644 --- a/tests/python/sym/test_sym_deduce_bound.py +++ b/tests/python/sym/test_sym_deduce_bound.py @@ -19,7 +19,7 @@ import tvm import tvm.testing -from tvm.tirx.buffer import decl_buffer +from tvm.tirx.buffer import decl_tensor def test_deduce(): @@ -228,7 +228,7 @@ def test_non_support(lhs): test_non_support(tvm.tirx.EQ(a, 16)) test_non_support(tvm.tirx.NE(a, 16)) test_non_support(tvm.tirx.log(a)) - test_non_support(tvm.tirx.BufferLoad(decl_buffer([16], "int32"), [a])) + test_non_support(tvm.tirx.BufferLoad(decl_tensor([16], "int32"), [a])) def test_deduce_floordiv(): diff --git a/tests/python/sym/test_sym_iter_affine_map.py b/tests/python/sym/test_sym_iter_affine_map.py index 61a22fde093e..ff11396ca5b6 100644 --- a/tests/python/sym/test_sym_iter_affine_map.py +++ b/tests/python/sym/test_sym_iter_affine_map.py @@ -1444,7 +1444,7 @@ def test_detect_iter_map_with_bufferload_recursion(): i = tvm.tirx.Var("i", "int32") j = tvm.tirx.Var("j", "int32") - buffer = tvm.tirx.decl_buffer((n,), "int32", name="seqlen") + buffer = tvm.tirx.decl_tensor((n,), "int32", name="seqlen") indices = [(buffer[i] + j) // divisor] iter_vars = { diff --git a/tests/python/sym/test_sym_rewrite_simplify.py b/tests/python/sym/test_sym_rewrite_simplify.py index 61fd4084f0ce..49645818ab2f 100644 --- a/tests/python/sym/test_sym_rewrite_simplify.py +++ b/tests/python/sym/test_sym_rewrite_simplify.py @@ -1297,7 +1297,7 @@ class TestDivZero(BaseCompare): class TestSubBufferload(BaseCompare): - buf = tvm.tirx.decl_buffer([1], dtype="float32") + buf = tvm.tirx.decl_tensor([1], dtype="float32") load = tvm.tirx.BufferLoad(buf, [0]) test_case = tvm.testing.parameter( diff --git a/tests/python/sym/test_sym_z3.py b/tests/python/sym/test_sym_z3.py index 373f6477686e..d355d5ae9316 100644 --- a/tests/python/sym/test_sym_z3.py +++ b/tests/python/sym/test_sym_z3.py @@ -691,7 +691,7 @@ def test_z3_memo_pool_reuse_survives_clone(): # BufferLoad is read-state, so these expressions are memoized only for the # duration of the query and then erased. This leaves holes in the Z3 memo # pool that subsequent pure expressions must be able to reuse safely. - buffer = tirx.decl_buffer((16,), "int32") + buffer = tirx.decl_tensor((16,), "int32") read_expr = tirx.all(*(buffer[i] >= 0 for i in range(16))) assert not analyzer.can_prove(read_expr, SB) diff --git a/tests/python/target/test_arm_target.py b/tests/python/target/test_arm_target.py index 30ac6662b859..d08b915390c2 100644 --- a/tests/python/target/test_arm_target.py +++ b/tests/python/target/test_arm_target.py @@ -59,7 +59,7 @@ def test_scalable_div(sve_device_vector_length): dev = tvm.cpu(0) @Ts.prim_func - def my_func(A: T.Buffer((1,), "int32")): + def my_func(A: T.Tensor((1,), "int32")): T.func_attr({"global_symbol": "my_module", "tirx.noalias": True}) A[0] = T.Div(10000, 4 * T.vscale()) @@ -80,7 +80,7 @@ def test_scalable_buffer_load_store(sve_device_vector_length): dev = tvm.cpu(0) @Ts.prim_func - def my_func(A: T.Buffer((num_elements,), "float32"), B: T.Buffer((num_elements,), "float32")): + def my_func(A: T.Tensor((num_elements,), "float32"), B: T.Tensor((num_elements,), "float32")): T.func_attr({"global_symbol": "my_module", "tirx.noalias": True}) B[T.ramp(0, 1, 4 * T.vscale())] = A[T.ramp(0, 1, 4 * T.vscale())] @@ -105,7 +105,7 @@ def test_scalable_loop_bound(sve_device_vector_length): dev = tvm.cpu(0) @Ts.prim_func - def my_func(A: T.Buffer((num_elements,), "float32"), B: T.Buffer((num_elements,), "float32")): + def my_func(A: T.Tensor((num_elements,), "float32"), B: T.Tensor((num_elements,), "float32")): T.func_attr({"global_symbol": "my_module", "tirx.noalias": True}) for i in T.serial(0, 4 * T.vscale()): B[i] = A[i] @@ -128,7 +128,7 @@ def test_scalable_broadcast(sve_device_vector_length): dev = tvm.cpu(0) @Ts.prim_func - def my_func(A: T.Buffer((num_elements,), "float32")): + def my_func(A: T.Tensor((num_elements,), "float32")): T.func_attr({"global_symbol": "my_module", "tirx.noalias": True}) A[T.ramp(0, 1, 4 * T.vscale())] = T.broadcast(1, 4 * T.vscale()) diff --git a/tests/python/te/test_te_create_primfunc.py b/tests/python/te/test_te_create_primfunc.py index a703736b0029..5ba19e49e624 100644 --- a/tests/python/te/test_te_create_primfunc.py +++ b/tests/python/te/test_te_create_primfunc.py @@ -71,7 +71,7 @@ def te_matmul(): @Ts.prim_func -def tir_matmul(A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def tir_matmul(A: T.Tensor((128, 128)), B: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0, j0, k0 in T.grid(128, 128, 128): @@ -84,9 +84,9 @@ def tir_matmul(A: T.Buffer((128, 128)), B: T.Buffer((128, 128)), C: T.Buffer((12 @Ts.prim_func def tir_matmul_int64( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(128)), "float32"), - C: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(128)), "float32"), + C: T.Tensor((T.int64(128), T.int64(128)), "float32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0, j0, k0 in T.grid(T.int64(128), T.int64(128), T.int64(128)): @@ -113,7 +113,7 @@ def te_element_wise(): @Ts.prim_func -def tir_element_wise(A: T.Buffer((128, 128)), C: T.Buffer((128, 128))) -> None: +def tir_element_wise(A: T.Tensor((128, 128)), C: T.Tensor((128, 128))) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) B = Ts.sblock_alloc_buffer((128, 128)) @@ -165,7 +165,7 @@ def te_conv2d(): @Ts.prim_func def tir_conv2d( - A: T.Buffer([16, 16, 14, 14]), W: T.Buffer([16, 3, 3, 32]), B: T.Buffer([16, 32, 14, 14]) + A: T.Tensor([16, 16, 14, 14]), W: T.Tensor([16, 3, 3, 32]), B: T.Tensor([16, 32, 14, 14]) ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -206,7 +206,7 @@ def te_multi_output(): @Ts.prim_func def tir_multi_output( - A0: T.Buffer((m, n)), A1: T.Buffer((m, n)), B0: T.Buffer((m, n)), B1: T.Buffer((m, n)) + A0: T.Tensor((m, n)), A1: T.Tensor((m, n)), B0: T.Tensor((m, n)), B1: T.Tensor((m, n)) ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -244,9 +244,9 @@ def te_extern(): @Ts.prim_func def tir_extern( - A: T.Buffer((128, 128), elem_offset=off1), - B: T.Buffer((128, 128), elem_offset=off2), - C: T.Buffer((128, 128), elem_offset=off3), + A: T.Tensor((128, 128), elem_offset=off1), + B: T.Tensor((128, 128), elem_offset=off2), + C: T.Tensor((128, 128), elem_offset=off3), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -306,9 +306,9 @@ def te_extern_epilogue(): @Ts.prim_func def tir_extern_epilogue( - A: T.Buffer((4, 3), offset_factor=1), - B: T.Buffer((3, 2), offset_factor=1), - D: T.Buffer((4, 2), "float32"), + A: T.Tensor((4, 3), offset_factor=1), + B: T.Tensor((3, 2), offset_factor=1), + D: T.Tensor((4, 2), "float32"), ): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -341,7 +341,7 @@ def te_reordered_matmul(): @Ts.prim_func def tir_reordered_matmul( - C: T.Buffer((128, 128)), A: T.Buffer((128, 128)), B: T.Buffer((128, 128)) + C: T.Tensor((128, 128)), A: T.Tensor((128, 128)), B: T.Tensor((128, 128)) ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -464,9 +464,9 @@ def test_tensor_attr(): @Ts.prim_func def expected_layout_attr( - A: T.Buffer((128, 128), "float32"), - B: T.Buffer((128, 128), "float32"), - D: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + B: T.Tensor((128, 128), "float32"), + D: T.Tensor((128, 128), "float32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True, "layout_free_buffers": [1]}) C = Ts.sblock_alloc_buffer([128, 128], dtype="float32") @@ -485,9 +485,9 @@ def expected_layout_attr( @Ts.prim_func def expected_layout_attr_int64( - A: T.Buffer((T.int64(128), T.int64(128)), "float32"), - B: T.Buffer((T.int64(128), T.int64(128)), "float32"), - D: T.Buffer((T.int64(128), T.int64(128)), "float32"), + A: T.Tensor((T.int64(128), T.int64(128)), "float32"), + B: T.Tensor((T.int64(128), T.int64(128)), "float32"), + D: T.Tensor((T.int64(128), T.int64(128)), "float32"), ): T.func_attr({"global_symbol": "main", "tirx.noalias": True, "layout_free_buffers": [1]}) C = Ts.sblock_alloc_buffer([T.int64(128), T.int64(128)], dtype="float32") @@ -560,10 +560,10 @@ def f_identity(dtype0: tvm.DataType, dtype1: tvm.DataType): @Ts.prim_func def tir_argmax_idx_val( - idx: T.Buffer([m, n], dtype="int32"), - val: T.Buffer([m, n], dtype="float32"), - argmax_v0: T.Buffer([m], dtype="int32"), - argmax_v1: T.Buffer([m], dtype="float32"), + idx: T.Tensor([m, n], dtype="int32"), + val: T.Tensor([m, n], dtype="float32"), + argmax_v0: T.Tensor([m], dtype="int32"), + argmax_v1: T.Tensor([m], dtype="float32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -613,10 +613,10 @@ def f_identity(dtype0: tvm.DataType, dtype1: tvm.DataType): @Ts.prim_func def tir_argmax_val_idx( - val: T.Buffer([m, n], dtype="float32"), - idx: T.Buffer([m, n], dtype="int32"), - argmax_v0: T.Buffer([m], dtype="float32"), - argmax_v1: T.Buffer([m], dtype="int32"), + val: T.Tensor([m, n], dtype="float32"), + idx: T.Tensor([m, n], dtype="int32"), + argmax_v0: T.Tensor([m], dtype="float32"), + argmax_v1: T.Tensor([m], dtype="int32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -666,9 +666,9 @@ def te_func(): @Ts.prim_func def expected( - a: T.Buffer((), "int32"), - b: T.Buffer((), "int32"), - c: T.Buffer((), "int32"), + a: T.Tensor((), "int32"), + b: T.Tensor((), "int32"), + c: T.Tensor((), "int32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with Ts.sblock("root"): @@ -692,8 +692,8 @@ def te_reshape(): @Ts.prim_func def tir_reshape( - A: T.Buffer((T.int64(2), T.int64(4)), "float32"), - T_reshape: T.Buffer((T.int64(4), T.int64(2)), "float32"), + A: T.Tensor((T.int64(2), T.int64(4)), "float32"), + T_reshape: T.Tensor((T.int64(4), T.int64(2)), "float32"), ): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0, i1 in T.grid(T.int64(4), T.int64(2)): @@ -738,8 +738,8 @@ def te_resize2d_symbolic(): @Ts.prim_func def tir_resize2d_symbolic( - A: T.Buffer((T.int64(2), T.int64(3), T.int64(128), T.int64(128)), "float32"), - resize: T.Buffer([T.int64(2), T.int64(3), oh, ow], dtype="float32"), + A: T.Tensor((T.int64(2), T.int64(3), T.int64(128), T.int64(128)), "float32"), + resize: T.Tensor([T.int64(2), T.int64(3), oh, ow], dtype="float32"), ): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -799,10 +799,10 @@ def te_extern(): @Ts.prim_func def tir_extern( - A: T.Buffer([128, 128], dtype="float32", offset_factor=1), - B: T.Buffer([128, 128], dtype="float32", offset_factor=1), - P: T.Buffer([1], dtype="float32", offset_factor=1), - C: T.Buffer([128, 128], dtype="float32", offset_factor=1), + A: T.Tensor([128, 128], dtype="float32", offset_factor=1), + B: T.Tensor([128, 128], dtype="float32", offset_factor=1), + P: T.Tensor([1], dtype="float32", offset_factor=1), + C: T.Tensor([128, 128], dtype="float32", offset_factor=1), ): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) @@ -828,7 +828,7 @@ def te_slice_with_var_input(): @Ts.prim_func -def tir_slice_with_var_input(tensor: T.Buffer((m, n)), idx: T.int64, slice: T.Buffer((idx, n))): # noqa: F821 +def tir_slice_with_var_input(tensor: T.Tensor((m, n)), idx: T.int64, slice: T.Tensor((idx, n))): # noqa: F821 T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) # with Ts.sblock("root"): @@ -849,7 +849,7 @@ def test_loop_aware_initial_value(): """Test initial value aware of spatial iter position""" @Ts.prim_func - def tir_workload(a: T.Buffer((5, 5)), b: T.Buffer((5,)), sum_red: T.Buffer((5,))): + def tir_workload(a: T.Tensor((5, 5)), b: T.Tensor((5,)), sum_red: T.Tensor((5,))): T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) for i, ax in T.grid(5, 5): @@ -882,7 +882,7 @@ def test_loop_aware_reducer_combiner(): """Test combiner aware of spatial iter position""" @Ts.prim_func - def tir_workload(a: T.Buffer((5, 5)), b: T.Buffer((5,)), sum_red: T.Buffer((5,))): + def tir_workload(a: T.Tensor((5, 5)), b: T.Tensor((5,)), sum_red: T.Tensor((5,))): T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) for i, ax in T.grid(5, 5): @@ -917,8 +917,8 @@ def te_workload(): def test_adaptive_pooling_window(): @Ts.prim_func def tir_workload( - x: T.Buffer((1, 1024, 16, 40), "float32"), - adaptive_pool_avg: T.Buffer((1, 1024, 12, 30), "float32"), + x: T.Tensor((1, 1024, 16, 40), "float32"), + adaptive_pool_avg: T.Tensor((1, 1024, 12, 30), "float32"), ): T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) # fmt: off @@ -991,7 +991,7 @@ def test_global_pool(): def test_nested_reduce_domain_dependency(): @Ts.prim_func def tir_workload( - x: T.Buffer((8, 8, 8, 8, 8), "float32"), compute: T.Buffer((8, 8, 8), "float32") + x: T.Tensor((8, 8, 8, 8, 8), "float32"), compute: T.Tensor((8, 8, 8), "float32") ): T.func_attr({"tirx.noalias": True, "global_symbol": "main"}) for i0, i1, i2 in T.grid(8, 8, 8): diff --git a/tests/python/tirx-analysis/test_tir_analysis_undefined_vars.py b/tests/python/tirx-analysis/test_tir_analysis_undefined_vars.py index 82343213a886..836bb263c68a 100644 --- a/tests/python/tirx-analysis/test_tir_analysis_undefined_vars.py +++ b/tests/python/tirx-analysis/test_tir_analysis_undefined_vars.py @@ -22,9 +22,9 @@ def test_decl_buffer_data_is_use(): - """DeclBuffer's data var should be reported as undefined (USE), not defined. + """DeclTensor's data var should be reported as undefined (USE), not defined. - When UndefinedVars encounters a DeclBuffer, the data pointer references + When UndefinedVars encounters a DeclTensor, the data pointer references an existing variable from the enclosing scope. It must appear in the undefined list so that callers (e.g., CreateComputeScope) capture it. """ @@ -32,13 +32,13 @@ def test_decl_buffer_data_is_use(): from tvm.ir import PointerType, PrimType data_ptr = tirx.Var("buf_data", PointerType(PrimType("float32"))) - buf = tirx.decl_buffer((n,), "float32", "buf", data=data_ptr) + buf = tirx.decl_tensor((n,), "float32", "buf", data=data_ptr) body = tirx.Evaluate(tirx.BufferLoad(buf, [0])) decl = tirx.Bind( buf, tvm.ir.Call( - "tirx.decl_buffer", + "tirx.decl_tensor", [ data_ptr, tvm.ir.Tuple(buf.shape), @@ -52,14 +52,14 @@ def test_decl_buffer_data_is_use(): undef = tvm.tirx.analysis.undefined_vars(stmt, []) undef_names = {v.name for v in undef} - # data_ptr must be undefined (it comes from outside the DeclBuffer) + # data_ptr must be undefined (it comes from outside the DeclTensor) assert "buf_data" in undef_names, f"Expected buf_data in undefined vars, got {undef_names}" def test_decl_buffer_elem_offset_is_use(): - """DeclBuffer's elem_offset var should be reported as undefined (USE). + """DeclTensor's elem_offset var should be reported as undefined (USE). - After FlattenBuffer, DeclBuffer nodes carry elem_offset vars from + After FlattenBuffer, DeclTensor nodes carry elem_offset vars from match_buffer entries. These must appear in the undefined list. """ from tvm.ir import PointerType, PrimType @@ -67,13 +67,13 @@ def test_decl_buffer_elem_offset_is_use(): n = tirx.Var("n", "int32") data_ptr = tirx.Var("buf_data", PointerType(PrimType("float32"))) elem_off = tirx.Var("buf_elem_offset", "int32") - buf = tirx.decl_buffer((n,), "float32", "buf", data=data_ptr, elem_offset=elem_off) + buf = tirx.decl_tensor((n,), "float32", "buf", data=data_ptr, elem_offset=elem_off) body = tirx.Evaluate(tirx.BufferLoad(buf, [0])) decl = tirx.Bind( buf, tvm.ir.Call( - "tirx.decl_buffer", + "tirx.decl_tensor", [ data_ptr, tvm.ir.Tuple(buf.shape), @@ -94,19 +94,19 @@ def test_decl_buffer_elem_offset_is_use(): def test_alloc_buffer_data_is_def(): - """AllocBuffer's data var should NOT be reported as undefined (it's a DEF). + """AllocTensor's data var should NOT be reported as undefined (it's a DEF). - AllocBuffer allocates new storage — the data pointer is a new definition, + AllocTensor allocates new storage — the data pointer is a new definition, not a reference to an external variable. """ n = tirx.Var("n", "int32") - buf = tirx.decl_buffer((n,), "float32", "buf") + buf = tirx.decl_tensor((n,), "float32", "buf") body = tirx.Evaluate(tirx.BufferLoad(buf, [0])) alloc = tvm.tirx.Bind( buf, tvm.ir.Call( - "tirx.alloc_buffer", + "tirx.alloc_tensor", [ tvm.ir.Tuple(buf.shape), tvm.ir.DataTypeImm(tvm.DataType(buf.dtype)), @@ -120,7 +120,7 @@ def test_alloc_buffer_data_is_def(): undef = tvm.tirx.analysis.undefined_vars(stmt, []) undef_names = {v.name for v in undef} - # The buffer Var itself is defined by AllocBuffer. + # The buffer Var itself is defined by AllocTensor. assert buf.name not in undef_names # shape var n should be undefined (comes from enclosing scope) assert "n" in undef_names, f"Expected shape var 'n' in undefined vars, got {undef_names}" @@ -128,7 +128,7 @@ def test_alloc_buffer_data_is_def(): def test_buffer_data_projection_is_buffer_use(): """An opaque data projection must retain the BufferVar identity.""" - buf = tirx.decl_buffer((16,), "float32", "buf") + buf = tirx.decl_tensor((16,), "float32", "buf") stmt = tirx.Evaluate(buf.data) undef = tvm.tirx.analysis.undefined_vars(stmt, []) 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 9fef5089e6d3..eb7be9a52ba1 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 @@ -31,10 +31,10 @@ def test_pass_simple(): @T.prim_func def element_wise( - A: T.Buffer((128, 128), "float32"), - C: T.Buffer((128, 128), "float32"), + A: T.Tensor((128, 128), "float32"), + C: T.Tensor((128, 128), "float32"), ): - B = T.alloc_buffer((128, 128), "float32") + B = T.alloc_tensor((128, 128), "float32") for i, j in T.grid(128, 128): B[i, j] = A[i, j] * 2.0 for i, j in T.grid(128, 128): @@ -144,7 +144,7 @@ def test_reuse_of_env_thread_in_function_is_well_formed(): """ @T.prim_func - def func(A: T.Buffer([256], "float32")): + def func(A: T.Tensor([256], "float32")): threadIdx_x = T.env_thread("threadIdx.x") with T.launch_thread(threadIdx_x, 256): A[threadIdx_x] = A[threadIdx_x] + 1.0 @@ -166,7 +166,7 @@ def test_reuse_of_env_thread_in_function_is_mandatory(): """ @T.prim_func - def func(A: T.Buffer([256], "float32")): + def func(A: T.Tensor([256], "float32")): with T.launch_thread("threadIdx.x", 256) as threadIdx_x: A[threadIdx_x] = A[threadIdx_x] + 1.0 @@ -189,7 +189,7 @@ def test_reuse_of_env_thread_across_functions_is_ill_formed(): @I.ir_module(check_well_formed=False) class mod: @T.prim_func - def kernel_1(A: T.Buffer([256], "float32")): + def kernel_1(A: T.Tensor([256], "float32")): T.attr( T.iter_var(threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", @@ -198,7 +198,7 @@ def kernel_1(A: T.Buffer([256], "float32")): A[threadIdx_x] = A[threadIdx_x] + T.float32(1) @T.prim_func - def kernel_2(A: T.Buffer([256], "float32")): + def kernel_2(A: T.Tensor([256], "float32")): T.attr( T.iter_var(threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", @@ -223,7 +223,7 @@ def test_multiple_buffer_arguments_may_share_allocation(): @I.ir_module class mod: @T.prim_func - def func(A: T.Buffer([256], "float32"), B: T.Buffer([256], "float32", data=A.data)): # noqa: F821 + def func(A: T.Tensor([256], "float32"), B: T.Tensor([256], "float32", data=A.data)): # noqa: F821 pass tvm.tirx.analysis.verify_well_formed(mod) @@ -319,10 +319,10 @@ def func(): def test_buffer_param_is_well_formed(): - """BufferType-annotated parameters are in scope for the body.""" + """TensorType-annotated parameters are in scope for the body.""" @T.prim_func - def func(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def func(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): for i in T.grid(128): B[i] = A[i] * 2.0 @@ -330,11 +330,11 @@ def func(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): def test_decl_buffer_is_well_formed(): - """A DeclBuffer statement introduces a buffer into scope for its body.""" + """A DeclTensor statement introduces a buffer into scope for its body.""" @T.prim_func - def func(A: T.Buffer((128,), "float32")): - B = T.alloc_buffer((128,), "float32") + def func(A: T.Tensor((128,), "float32")): + B = T.alloc_tensor((128,), "float32") for i in T.grid(128): B[i] = A[i] * 2.0 @@ -347,8 +347,8 @@ def test_alloc_buffer_is_well_formed(): @I.ir_module class mod: @T.prim_func - def func(A: T.Buffer((128,), "float32")): - B = T.alloc_buffer([128], "float32") + def func(A: T.Tensor((128,), "float32")): + B = T.alloc_tensor([128], "float32") for i in T.grid(128): B[i] = A[i] * 2.0 @@ -358,7 +358,7 @@ def func(A: T.Buffer((128,), "float32")): def test_tensor_load_asserted_type_matches_source_and_indices(): @T.prim_func def func(): - buffer = T.alloc_buffer((4,), "float32") + buffer = T.alloc_tensor((4,), "float32") T.evaluate(buffer[0]) serialized = tvm.ir.save_json(func) @@ -383,7 +383,7 @@ def func(): def test_tensor_load_malformed_indices_return_false_without_asserting(): - buffer = tvm.tirx.decl_buffer((4, 4), "float32") + buffer = tvm.tirx.decl_tensor((4, 4), "float32") vector_index = tvm.tirx.Ramp(0, 1, 4) load = tvm.tirx.BufferLoad(buffer, [0, vector_index]) func = tvm.tirx.PrimFunc([buffer], tvm.tirx.Evaluate(load)) diff --git a/tests/python/tirx-base/test_tir_base.py b/tests/python/tirx-base/test_tir_base.py index 7a8dcd09abe1..2f9ed629b044 100644 --- a/tests/python/tirx-base/test_tir_base.py +++ b/tests/python/tirx-base/test_tir_base.py @@ -162,7 +162,7 @@ def func(a: T.float32, b: T.float32): def test_break_statement(): @T.prim_func - def func(In: T.Buffer((2,), "int32"), Out: T.Buffer((2,), "int32")): + def func(In: T.Tensor((2,), "int32"), Out: T.Tensor((2,), "int32")): Out[0] = 0 Out[1] = 1 for i in range(10): @@ -189,7 +189,7 @@ def func(In: T.Buffer((2,), "int32"), Out: T.Buffer((2,), "int32")): def test_continue_statement(): @T.prim_func - def func(Out: T.Buffer((2,), "int32")): + def func(Out: T.Tensor((2,), "int32")): T.func_attr({"global_symbol": "main"}) Out[0] = 0 Out[1] = 0 @@ -198,7 +198,7 @@ def func(Out: T.Buffer((2,), "int32")): if (i * 10 + j) % 3 != 0: continue Out[0] = Out[0] + 1 - k = T.decl_buffer([], "int32") + k = T.decl_tensor([], "int32") k[()] = 0 while k[()] < Out[0]: k[()] = k[()] + 1 diff --git a/tests/python/tirx-base/test_tir_buffer.py b/tests/python/tirx-base/test_tir_buffer.py index 87066ef0ec29..c1b242ed28f0 100644 --- a/tests/python/tirx-base/test_tir_buffer.py +++ b/tests/python/tirx-base/test_tir_buffer.py @@ -29,20 +29,26 @@ def test_buffer(): m = tvm.tirx.Var("m", "int32") n = tvm.tirx.Var("n", "int32") l = tvm.tirx.Var("l", "int32") - Ab = tvm.tirx.decl_buffer((m, n), "float32") - Bb = tvm.tirx.decl_buffer((n, l), "float32") + Ab = tvm.tirx.decl_tensor((m, n), "float32") + Bb = tvm.tirx.decl_tensor((n, l), "float32") assert type(Ab) is tvm.ir.Var assert tvm.tirx.is_buffer_var(Ab) - assert isinstance(Ab.ty, tvm.tirx.BufferType) + assert isinstance(Ab.ty, tvm.tirx.TensorType) assert Ab.ty.dtype == tvm.ir.PrimType("float32") assert tuple(Ab.ty.shape) == (m, n) assert not tvm.tirx.is_buffer_var(m) + serialized = tvm.ir.save_json(Ab.ty) + assert '"tirx.TensorType"' in serialized + restored = tvm.ir.load_json(serialized) + assert isinstance(restored, tvm.tirx.TensorType) + tvm.ir.assert_structural_equal(restored, Ab.ty, map_free_vars=True) + def test_buffer_compatibility_alias_and_global_var_properties(): scalar = tvm.ir.Var("scalar", tvm.ir.PrimType("int32")) - buffer = tvm.tirx.decl_buffer((8,), "float32") + buffer = tvm.tirx.decl_tensor((8,), "float32") assert tvm.tirx.Buffer is tvm.ir.Var assert isinstance(scalar, tvm.tirx.Buffer) @@ -55,12 +61,12 @@ def test_buffer_compatibility_alias_and_global_var_properties(): for name in ("shape", "dtype", "data"): assert not hasattr(scalar, name) - with pytest.raises(AttributeError, match="only available on a Var with BufferType"): + with pytest.raises(AttributeError, match="only available on a Var with TensorType"): getattr(scalar, name) def test_buffer_data_is_typed_projection(): - buffer = tvm.tirx.decl_buffer((8,), "bool", scope="shared") + buffer = tvm.tirx.decl_tensor((8,), "bool", scope="shared") assert buffer.ty.dtype == tvm.ir.PrimType("bool") assert tvm.tirx.buffer_data_pointer_type(buffer) == tvm.ir.PointerType( @@ -73,7 +79,7 @@ def test_buffer_data_is_typed_projection(): def test_buffer_pointer_type_derived_from_dtype_and_scope(): data = tvm.ir.Var("storage", tvm.ir.PointerType(tvm.ir.PrimType("uint8"), "local")) - buffer = tvm.tirx.decl_buffer((8,), "float16", data=data) + buffer = tvm.tirx.decl_tensor((8,), "float16", data=data) assert buffer.ty.dtype == tvm.ir.PrimType("float16") assert buffer.ty.storage_scope == "local" @@ -81,13 +87,13 @@ def test_buffer_pointer_type_derived_from_dtype_and_scope(): def test_decl_buffer_physical_data_binding(): - buffer = tvm.tirx.decl_buffer((8,), "float32") + buffer = tvm.tirx.decl_tensor((8,), "float32") data = tvm.tirx.Var("data", buffer.data.ty) decl = tvm.tirx.Bind( buffer, tvm.ir.Call( - "tirx.decl_buffer", + "tirx.decl_tensor", [ data, tvm.ir.Tuple(buffer.shape), @@ -104,7 +110,7 @@ def test_decl_buffer_physical_data_binding(): def test_buffer_access_ptr(): m = tvm.tirx.Var("m", "int32") n = tvm.tirx.Var("n", "int32") - Ab = tvm.tirx.decl_buffer((m, n), "float32", strides=[n + 1, 1]) + Ab = tvm.tirx.decl_tensor((m, n), "float32", strides=[n + 1, 1]) aptr = Ab.access_ptr("rw") assert isinstance(aptr.ty, tvm.ir.PointerType) assert aptr.ty.element_type == tvm.ir.PrimType("void") @@ -113,7 +119,7 @@ def test_buffer_access_ptr(): assert aptr.args[4].value == BufferAccessKind.READ | BufferAccessKind.WRITE typed_ptr = Ab.access_ptr("r", ptr_type="uint8") assert typed_ptr.ty == tvm.ir.PointerType(tvm.ir.PrimType("uint8")) - shared = tvm.tirx.decl_buffer((m, n), "float32", scope="shared") + shared = tvm.tirx.decl_tensor((m, n), "float32", scope="shared") assert shared.access_ptr("r").ty == tvm.ir.PointerType(tvm.ir.PrimType("void"), "shared") assert shared.access_ptr("r", ptr_type="uint8").ty == tvm.ir.PointerType( tvm.ir.PrimType("uint8"), "shared" @@ -125,7 +131,7 @@ def test_buffer_access_ptr(): def test_buffer_access_ptr_offset(): m = tvm.tirx.Var("m", "int32") n = tvm.tirx.Var("n", "int32") - Ab = tvm.tirx.decl_buffer((m, n), "float32") + Ab = tvm.tirx.decl_tensor((m, n), "float32") aptr = Ab.access_ptr("rw", offset=100) tvm.testing.assert_prim_expr_equal(aptr.args[2], 100) assert aptr.args[4].value == BufferAccessKind.READ | BufferAccessKind.WRITE @@ -143,12 +149,12 @@ def test_buffer_access_ptr_offset(): def test_buffer_access_ptr_extent(): m = tvm.tirx.Var("m", "int32") n = tvm.tirx.Var("n", "int32") - Ab = tvm.tirx.decl_buffer((m, n), "float32") + Ab = tvm.tirx.decl_tensor((m, n), "float32") aptr = Ab.access_ptr("rw") tvm.ir.assert_structural_equal(aptr.args[3], m * n) aptr = Ab.access_ptr("rw", offset=100) tvm.ir.assert_structural_equal(aptr.args[3], m * n - 100) - Ab = tvm.tirx.decl_buffer((m, n), "float32", strides=[n + 1, 1]) + Ab = tvm.tirx.decl_tensor((m, n), "float32", strides=[n + 1, 1]) aptr = Ab.access_ptr("rw", offset=100) tvm.ir.assert_structural_equal(aptr.args[3], Ab.ty.strides[0] * m - 100) @@ -162,7 +168,7 @@ def test_buffer_access_ptr_extent(): def test_buffer_vload(): m = tvm.tirx.Var("m", "int32") n = tvm.tirx.Var("n", "int32") - Ab = tvm.tirx.decl_buffer((m, n), "float32", elem_offset=100) + Ab = tvm.tirx.decl_tensor((m, n), "float32", elem_offset=100) load = Ab.vload([2, 3]) tvm.ir.assert_structural_equal(load.indices, [T.int32(2), T.int32(3)]) @@ -170,7 +176,7 @@ def test_buffer_vload(): def test_buffer_offset_of(): m = tvm.tirx.Var("m", "int32") n = tvm.tirx.Var("n", "int32") - Ab = tvm.tirx.decl_buffer((m, n), "float32", elem_offset=100) + Ab = tvm.tirx.decl_tensor((m, n), "float32", elem_offset=100) offset = Ab.offset_of([2, 3]) tvm.ir.assert_structural_equal(offset, [n * 2 + 103]) @@ -181,8 +187,8 @@ def test_buffer_index_merge_mult_mod(): s = tvm.tirx.Var("s", "int32") k0 = tvm.tirx.Var("k0", "int32") k1 = tvm.tirx.Var("k1", "int32") - A = tvm.tirx.decl_buffer((m, n), "float32") - A_stride = tvm.tirx.decl_buffer((m, n), "float32", strides=(s, 1)) + A = tvm.tirx.decl_tensor((m, n), "float32") + A_stride = tvm.tirx.decl_tensor((m, n), "float32", strides=(s, 1)) def assert_simplified_equal(index_simplified, index_direct): ( @@ -225,7 +231,7 @@ def assert_simplified_equal(index_simplified, index_direct): assert_simplified_equal(index_simplified, index_direct) # Test Case5 - B = tvm.tirx.decl_buffer((1, 14, 14, 1024)) + B = tvm.tirx.decl_tensor((1, 14, 14, 1024)) i = tvm.tirx.Var("i", "int32") j = tvm.tirx.Var("j", "int32") k = tvm.tirx.Var("k", "int32") @@ -253,7 +259,7 @@ def assert_simplified_equal(index_simplified, index_direct): def test_buffer_flatten(): """A buffer should flatten to a 1-d shape""" - buf = tvm.tirx.decl_buffer([16, 32]) + buf = tvm.tirx.decl_tensor([16, 32]) flat = buf.get_flattened_buffer() # A metadata-changing rewrite creates a fresh typed Var. The physical # pointer is always derived from that Var instead of being stored as a @@ -266,7 +272,7 @@ def test_buffer_flatten(): def test_buffer_flatten_preserves_identity(): """Flattening a 1-d buffer should return the original""" - buf = tvm.tirx.decl_buffer([16]) + buf = tvm.tirx.decl_tensor([16]) flat = buf.get_flattened_buffer() assert buf.same_as(flat) diff --git a/tests/python/tirx-base/test_tir_constructor.py b/tests/python/tirx-base/test_tir_constructor.py index 3c24f75f89cd..20c019652fc6 100644 --- a/tests/python/tirx-base/test_tir_constructor.py +++ b/tests/python/tirx-base/test_tir_constructor.py @@ -104,7 +104,7 @@ def test_expr_constructor(): assert x.condition == a buffer_var = tvm.tirx.Var("buf", tvm.ir.PointerType(tvm.ir.PrimType("float32"))) - buffer = tvm.tirx.decl_buffer([16], "float32", data=buffer_var) + buffer = tvm.tirx.decl_tensor([16], "float32", data=buffer_var) x = tvm.tirx.BufferLoad(buffer, [1]) assert isinstance(x, tvm.ir.TensorLoad) assert x.ty == tvm.ir.PrimType("float32") @@ -222,7 +222,7 @@ def call_with(arg): def test_buffer_region_call_wrappers_reject(): - buffer = tvm.tirx.decl_buffer([4], "int32") + buffer = tvm.tirx.decl_tensor([4], "int32") region = buffer[0:4] calls = [ lambda: tvm.tirx.call_intrin("int32", "tirx.reinterpret", region), @@ -242,22 +242,22 @@ def test_buffer_region_call_wrappers_reject(): def test_buffer_region_type_is_singleton(): - lhs = tvm.tirx.decl_buffer([1], "int32")[0:1] - rhs = tvm.tirx.decl_buffer([2], "float32")[0:2] + lhs = tvm.tirx.decl_tensor([1], "int32")[0:1] + rhs = tvm.tirx.decl_tensor([2], "float32")[0:2] assert isinstance(lhs, TensorRegion) assert isinstance(rhs, TensorRegion) assert lhs.ty.same_as(rhs.ty) def test_buffer_region_is_not_arithmetic_operand(): - int_region = tvm.tirx.decl_buffer([4], "int32")[0:4] + int_region = tvm.tirx.decl_tensor([4], "int32")[0:4] with pytest.raises(TypeError, match="construct a TensorLoad explicitly"): tvm.tirx.IterVar((0, 4), "i", tvm.tirx.IterVar.DataPar) + int_region def test_operator_base_categories_have_primitive_type(): var = tvm.tirx.Var("x", "int32") - buffer = tvm.tirx.decl_buffer([4], "float32") + buffer = tvm.tirx.decl_tensor([4], "float32") expressions = [ tvm.tirx.IntImm("int32", 1), tvm.tirx.Add(var, 1), @@ -313,7 +313,7 @@ def test_stmt_constructor(): assert x.body == nop buffer_var = tvm.tirx.Var("buf", tvm.ir.PointerType(tvm.ir.PrimType("bool"))) - buffer = tvm.tirx.decl_buffer([16], "bool", data=buffer_var) + buffer = tvm.tirx.decl_tensor([16], "bool", data=buffer_var) x = tvm.tirx.BufferStore(buffer, tvm.tirx.IntImm("bool", 1), [10]) assert isinstance(x, tvm.tirx.BufferStore) assert x.buffer == buffer @@ -322,11 +322,11 @@ def test_stmt_constructor(): assert list(x.indices) == [10] assert x.value.value == 1 - buf = tvm.tirx.decl_buffer([10], "float32") + buf = tvm.tirx.decl_tensor([10], "float32") x = tvm.tirx.Bind( buf, tvm.ir.Call( - "tirx.alloc_buffer", + "tirx.alloc_tensor", [ tvm.ir.Tuple(buf.shape), tvm.ir.DataTypeImm(tvm.DataType(buf.dtype)), @@ -336,7 +336,7 @@ def test_stmt_constructor(): ty=buf.ty, ), ) - assert _is_buffer_binding(x, "tirx.alloc_buffer") + assert _is_buffer_binding(x, "tirx.alloc_tensor") assert x.var == buf x = tvm.tirx.AttrStmt(buffer_var, "xyz", 1, nop) diff --git a/tests/python/tirx-base/test_tir_host_func.py b/tests/python/tirx-base/test_tir_host_func.py index 126d01b28f9d..5fc83eb6670c 100644 --- a/tests/python/tirx-base/test_tir_host_func.py +++ b/tests/python/tirx-base/test_tir_host_func.py @@ -26,9 +26,9 @@ class Module: @T.prim_func def main( - A: T.Buffer((729, 729), "float32"), - B: T.Buffer((729, 729), "float32"), - C: T.Buffer((729, 729), "float32"), + A: T.Tensor((729, 729), "float32"), + B: T.Tensor((729, 729), "float32"), + C: T.Tensor((729, 729), "float32"), ): T.func_attr( { diff --git a/tests/python/tirx-base/test_tir_imm_values.py b/tests/python/tirx-base/test_tir_imm_values.py index d1dc449e15ea..5333b6c3659b 100644 --- a/tests/python/tirx-base/test_tir_imm_values.py +++ b/tests/python/tirx-base/test_tir_imm_values.py @@ -257,19 +257,19 @@ def test_tir_floatimm_const_fold(): """Behavior check: folding fp32 match platform f32 arithmetic""" @T.prim_func - def float_imm_multiply(x: T.float32, y: T.float32, z: T.Buffer((), "float32")): + def float_imm_multiply(x: T.float32, y: T.float32, z: T.Tensor((), "float32")): z[()] = x * y @T.prim_func - def float_imm_add(x: T.float32, y: T.float32, z: T.Buffer((), "float32")): + def float_imm_add(x: T.float32, y: T.float32, z: T.Tensor((), "float32")): z[()] = x + y @T.prim_func - def float_imm_sub(x: T.float32, y: T.float32, z: T.Buffer((), "float32")): + def float_imm_sub(x: T.float32, y: T.float32, z: T.Tensor((), "float32")): z[()] = x - y @T.prim_func - def float_imm_div(x: T.float32, y: T.float32, z: T.Buffer((), "float32")): + def float_imm_div(x: T.float32, y: T.float32, z: T.Tensor((), "float32")): z[()] = x / y def __wrap_build(f): diff --git a/tests/python/tirx-base/test_tir_intrin.py b/tests/python/tirx-base/test_tir_intrin.py index 5de95d2f51a6..45d31eaa2a77 100644 --- a/tests/python/tirx-base/test_tir_intrin.py +++ b/tests/python/tirx-base/test_tir_intrin.py @@ -39,7 +39,7 @@ def _unary_kernel(op, dtype="float32", out_dtype=None, gpu=False): n = T.int32() @T.prim_func - def kernel(A: T.Buffer((n,), dtype), B: T.Buffer((n,), out_dtype)): + def kernel(A: T.Tensor((n,), dtype), B: T.Tensor((n,), out_dtype)): if I.constexpr(gpu): for bx in T.thread_binding(T.ceildiv(n, 64), thread="blockIdx.x"): for tx in T.thread_binding(64, thread="threadIdx.x"): @@ -57,7 +57,7 @@ def _binary_kernel(op, rhs_dtype="float32"): @T.prim_func def kernel( - A: T.Buffer((n,), "float32"), B: T.Buffer((n,), rhs_dtype), C: T.Buffer((n,), "float32") + A: T.Tensor((n,), "float32"), B: T.Tensor((n,), rhs_dtype), C: T.Tensor((n,), "float32") ): for i in range(n): C[i] = op(A[i], B[i]) @@ -318,10 +318,10 @@ def run_and_check(): class Module: @T.prim_func def test_tir_fma( - A_1: T.Buffer([n], strides=[stride], elem_offset=0, align=64, offset_factor=1), - B_1: T.Buffer([n], strides=[stride_1], elem_offset=0, align=64, offset_factor=1), - C_1: T.Buffer([n], strides=[stride_2], elem_offset=0, align=64, offset_factor=1), - d_1: T.Buffer([n], strides=[stride_3], elem_offset=0, align=64, offset_factor=1), + A_1: T.Tensor([n], strides=[stride], elem_offset=0, align=64, offset_factor=1), + B_1: T.Tensor([n], strides=[stride_1], elem_offset=0, align=64, offset_factor=1), + C_1: T.Tensor([n], strides=[stride_2], elem_offset=0, align=64, offset_factor=1), + d_1: T.Tensor([n], strides=[stride_3], elem_offset=0, align=64, offset_factor=1), ) -> None: # function attr dict T.func_attr({"global_symbol": "test_fma", "tirx.noalias": True}) diff --git a/tests/python/tirx-base/test_tir_nodes.py b/tests/python/tirx-base/test_tir_nodes.py index e79f01fb556c..601da970c52d 100644 --- a/tests/python/tirx-base/test_tir_nodes.py +++ b/tests/python/tirx-base/test_tir_nodes.py @@ -80,7 +80,7 @@ def test_ir2(): storage_type = ir.PrimType("int32") handle_type = ir.PointerType(storage_type) array = tvm.tirx.Var("array", handle_type) - buf = tvm.tirx.decl_buffer([buf_size], "int32", data=array) + buf = tvm.tirx.decl_tensor([buf_size], "int32", data=array) st = tvm.tirx.BufferStore(buf, x + 1, [1]) assert isinstance(st, tvm.tirx.BufferStore) @@ -306,7 +306,7 @@ def test_equality_string_imm(): def test_prim_func(): x = tvm.tirx.Var("x", "int32") y = tvm.tirx.Var("y", "int32") - b = tvm.tirx.decl_buffer((x,), "float32") + b = tvm.tirx.decl_tensor((x,), "float32") stmt = tvm.tirx.SeqStmt([tvm.tirx.Bind(x, 10), tvm.tirx.Evaluate(x + 1)]) func = tvm.tirx.PrimFunc([x, y, b], stmt) @@ -341,7 +341,7 @@ def test_scoped_storage_vars(): def test_buffer_load_store(): - b = tvm.tirx.decl_buffer((10,), "float32") + b = tvm.tirx.decl_tensor((10,), "float32") x = tvm.tirx.BufferLoad(b, [0]) assert isinstance(x, tvm.ir.TensorLoad) assert callable(tvm.tirx.BufferLoad) @@ -415,7 +415,7 @@ def test_broadcast_to_scalable_vec(): def test_buffer_load_scalable_vec(): - buf = tvm.tirx.decl_buffer((24,), "float32") + buf = tvm.tirx.decl_tensor((24,), "float32") index = tvm.tirx.expr.Ramp(1, 1, 8 * tvm.tirx.vscale()) load = tvm.tirx.BufferLoad(buf, [index]) @@ -424,7 +424,7 @@ def test_buffer_load_scalable_vec(): def test_buffer_store_scalable_vec(): - b = tvm.tirx.decl_buffer((24,), "int32") + b = tvm.tirx.decl_tensor((24,), "int32") value = tvm.tirx.expr.Broadcast(1, 4 * tvm.tirx.vscale()) index = tvm.tirx.expr.Ramp(0, 1, 4 * tvm.tirx.vscale()) store = tvm.tirx.BufferStore(b, value, [index]) @@ -434,7 +434,7 @@ def test_buffer_store_scalable_vec(): def test_scalable_vec_cast(): - b = tvm.tirx.decl_buffer((24,), "float32") + b = tvm.tirx.decl_tensor((24,), "float32") value = tvm.tirx.expr.Broadcast(1, 12 * tvm.tirx.vscale()).astype("float32xvscalex12") index = tvm.tirx.expr.Ramp(0, 1, 12 * tvm.tirx.vscale()) diff --git a/tests/python/tirx-base/test_tir_op_types.py b/tests/python/tirx-base/test_tir_op_types.py index 085b0c69f619..38c5d83e3681 100644 --- a/tests/python/tirx-base/test_tir_op_types.py +++ b/tests/python/tirx-base/test_tir_op_types.py @@ -47,11 +47,11 @@ def test_tir_op_tvm_struct_set(): def test_tir_op_address_of(): - buffer = tirx.decl_buffer((128), "float32") + buffer = tirx.decl_tensor((128), "float32") expr = tirx.address_of(buffer[0]) assert expr.op.name == "tirx.address_of" storage = tirx.Var("storage", tvm.ir.PointerType(tvm.ir.PrimType("uint8"), "shared.dyn")) - pooled_buffer = tirx.decl_buffer((128), "float32", data=storage, scope="shared.dyn") + pooled_buffer = tirx.decl_tensor((128), "float32", data=storage, scope="shared.dyn") expected_ty = tvm.ir.PointerType(tvm.ir.PrimType("float32"), "shared.dyn") assert tirx.address_of(pooled_buffer).ty == expected_ty assert tirx.address_of(pooled_buffer[0]).ty == expected_ty @@ -94,7 +94,7 @@ def test_tir_op_call_likely(): def test_tir_op_tvm_thread_allreduce(): x = tirx.Var("x", "int32") - buffer = tirx.decl_buffer((128), "float32") + buffer = tirx.decl_tensor((128), "float32") y = tirx.Var("y", "handle") z = tirx.Var("z", "int32") expr = tirx.tvm_thread_allreduce(x, buffer[0], True, y, z) @@ -107,7 +107,7 @@ def test_tir_op_type_annotation(): def test_tir_op_tvm_access_ptr(): - buffer = tirx.decl_buffer((128), "float32") + buffer = tirx.decl_tensor((128), "float32") for ptype in ("float32", tvm.ir.PrimType("float32")): expr = tirx.tvm_access_ptr(ptype, buffer.data, 0, 1, 2) assert expr.op.name == "tirx.tvm_access_ptr" @@ -124,33 +124,33 @@ def test_tir_op_tvm_throw_last_error(): def test_tir_op_tvm_load_matrix_sync(): - buffer = tirx.decl_buffer((16, 16), "float32") + buffer = tirx.decl_tensor((16, 16), "float32") x = tirx.Var("x", "handle") expr = tirx.tvm_load_matrix_sync(buffer.data, 16, 16, 16, 0, x, 128, "row_major") assert expr.op.name == "tirx.tvm_load_matrix_sync" def test_tir_op_tvm_store_matrix_sync(): - buffer = tirx.decl_buffer((16, 16), "float32") + buffer = tirx.decl_tensor((16, 16), "float32") x = tirx.Var("x", "handle") expr = tirx.tvm_store_matrix_sync(buffer.data, 16, 16, 16, 0, x, 128, "row_major") assert expr.op.name == "tirx.tvm_store_matrix_sync" def test_tir_op_tvm_mma_sync(): - buffer_0 = tirx.decl_buffer((16, 16), "float32") - buffer_1 = tirx.decl_buffer((16, 16), "float32") - buffer_2 = tirx.decl_buffer((16, 16), "float32") - buffer_3 = tirx.decl_buffer((16, 16), "float32") + buffer_0 = tirx.decl_tensor((16, 16), "float32") + buffer_1 = tirx.decl_tensor((16, 16), "float32") + buffer_2 = tirx.decl_tensor((16, 16), "float32") + buffer_3 = tirx.decl_tensor((16, 16), "float32") expr = tirx.tvm_mma_sync(buffer_0.data, 0, buffer_1.data, 0, buffer_2.data, 0, buffer_3.data, 0) assert expr.op.name == "tirx.tvm_mma_sync" def test_tir_op_tvm_bmma_sync(): - buffer_0 = tirx.decl_buffer((16, 16), "float32") - buffer_1 = tirx.decl_buffer((16, 16), "float32") - buffer_2 = tirx.decl_buffer((16, 16), "float32") - buffer_3 = tirx.decl_buffer((16, 16), "float32") + buffer_0 = tirx.decl_tensor((16, 16), "float32") + buffer_1 = tirx.decl_tensor((16, 16), "float32") + buffer_2 = tirx.decl_tensor((16, 16), "float32") + buffer_3 = tirx.decl_tensor((16, 16), "float32") expr = tirx.tvm_bmma_sync( buffer_0.data, 0, buffer_1.data, 0, buffer_2.data, 0, buffer_3.data, 0 ) @@ -158,15 +158,15 @@ def test_tir_op_tvm_bmma_sync(): def test_tir_op_tvm_fill_fragment(): - buffer = tirx.decl_buffer((16, 16), "float32") + buffer = tirx.decl_tensor((16, 16), "float32") expr = tirx.tvm_fill_fragment(buffer.data, 16, 16, 16, 0, 0) assert expr.op.name == "tirx.tvm_fill_fragment" def test_tir_op_ptx_mma(): - buffer_a = tirx.decl_buffer([32], "int4", scope="local") - buffer_b = tirx.decl_buffer([16], "uint4", scope="local") - buffer_c = tirx.decl_buffer([4], "int32", scope="local") + buffer_a = tirx.decl_tensor([32], "int4", scope="local") + buffer_b = tirx.decl_tensor([16], "uint4", scope="local") + buffer_c = tirx.decl_tensor([4], "int32", scope="local") expr = _cuda_op.ptx_legacy_mma( "m8n8k32", "row", @@ -188,8 +188,8 @@ def test_tir_op_ptx_mma(): def test_tir_op_mma_store(): x = tirx.Var("x", ty="int32") y = tirx.Var("y", ty="int32") - buffer_w = tirx.decl_buffer([16, 8], dtype="int32", scope="warp", offset_factor=1) - buffer = tirx.decl_buffer( + buffer_w = tirx.decl_tensor([16, 8], dtype="int32", scope="warp", offset_factor=1) + buffer = tirx.decl_tensor( [16, 16], dtype="int32", scope="global", offset_factor=1, strides=[x, y] ) expr = _cuda_op.mma_store( @@ -205,14 +205,14 @@ def test_tir_op_mma_store(): def test_tir_op_mma_fill(): - buffer_w = tirx.decl_buffer([16, 8], dtype="int32", scope="warp", offset_factor=1) + buffer_w = tirx.decl_tensor([16, 8], dtype="int32", scope="warp", offset_factor=1) expr = _cuda_op.mma_fill("int32", 8, buffer_w.data, buffer_w.ty.elem_offset) assert expr.op.name == "tirx.mma_fill" def test_op_ptx_cp_async(): - buffer_shared = tirx.decl_buffer([16, 16], "float16", scope="shared") - buffer_local = tirx.decl_buffer([8], "float16", scope="local") + buffer_shared = tirx.decl_tensor([16, 16], "float16", scope="shared") + buffer_local = tirx.decl_tensor([8], "float16", scope="local") expr = _cuda_op.ptx_cp_async_legacy(buffer_shared.data, 0, buffer_local.data, 0, 16) assert expr.op.name == "tirx.s_tir.cp_async_raw" @@ -229,14 +229,14 @@ def test_op_ptx_cp_async(): def test_tir_op_vectorlow(): - buffer = tirx.decl_buffer((4, 4), "int8", offset_factor=1) + buffer = tirx.decl_tensor((4, 4), "int8", offset_factor=1) vec = buffer.vload([0, 0], dtype="int8x16") expr = tirx.vectorlow("int8x8", vec) assert expr.op.name == "tirx.vectorlow" def test_tir_op_vectorhigh(): - buffer = tirx.decl_buffer((4, 4), "int8", offset_factor=1) + buffer = tirx.decl_tensor((4, 4), "int8", offset_factor=1) vec = buffer.vload([0, 0], dtype="int8x16") expr = tirx.vectorhigh("int8x8", vec) assert expr.op.name == "tirx.vectorhigh" @@ -251,7 +251,7 @@ def test_tir_op_dp4a(): def test_tir_op_vectorcombine(): - buffer = tirx.decl_buffer((4, 4), "int8", offset_factor=1) + buffer = tirx.decl_tensor((4, 4), "int8", offset_factor=1) vec = buffer.vload([0, 0], dtype="int8x16") expr = tirx.vectorcombine("int8x8", vec, vec) assert expr.op.name == "tirx.vectorcombine" @@ -290,7 +290,7 @@ def test_tir_op_TVMBackendAllocWorkspace(): def test_tir_op_TVMBackendFreeWorkspace(): - buffer = tirx.decl_buffer((128), "float32") + buffer = tirx.decl_tensor((128), "float32") expr = tirx.TVMBackendFreeWorkspace(0, 1, buffer.data) assert expr.op.name == "tirx.TVMBackendFreeWorkspace" diff --git a/tests/python/tirx-base/test_tir_ptx_cp_async.py b/tests/python/tirx-base/test_tir_ptx_cp_async.py index 72562db1816d..3371a5d2e97a 100644 --- a/tests/python/tirx-base/test_tir_ptx_cp_async.py +++ b/tests/python/tirx-base/test_tir_ptx_cp_async.py @@ -25,13 +25,13 @@ @T.prim_func -def ptx_cp_async(A: T.Buffer((32, 128), "float16"), B: T.Buffer((32, 128), "float16")) -> None: +def ptx_cp_async(A: T.Tensor((32, 128), "float16"), B: T.Tensor((32, 128), "float16")) -> None: T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) bx = T.env_thread("blockIdx.x") tx = T.env_thread("threadIdx.x") T.launch_thread(bx, 1) T.launch_thread(tx, 32) - A_shared = T.alloc_buffer([32, 128], "float16", scope="shared") + A_shared = T.alloc_tensor([32, 128], "float16", scope="shared") for i in range(16): T.evaluate( T.s_tir.cp_async_raw.legacy( diff --git a/tests/python/tirx-base/test_tir_ptx_griddepcontrol.py b/tests/python/tirx-base/test_tir_ptx_griddepcontrol.py index 7ea8ac6171d1..7b2722ef3ef2 100644 --- a/tests/python/tirx-base/test_tir_ptx_griddepcontrol.py +++ b/tests/python/tirx-base/test_tir_ptx_griddepcontrol.py @@ -25,7 +25,7 @@ @T.prim_func -def ptx_griddepcontrol(A: T.Buffer((32,), "float32"), B: T.Buffer((32,), "float32")) -> None: +def ptx_griddepcontrol(A: T.Tensor((32,), "float32"), B: T.Tensor((32,), "float32")) -> None: T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) bx = T.env_thread("blockIdx.x") tx = T.env_thread("threadIdx.x") diff --git a/tests/python/tirx-base/test_tir_ptx_ldmatrix.py b/tests/python/tirx-base/test_tir_ptx_ldmatrix.py index 5654c0153ced..ff4405e7bb1a 100644 --- a/tests/python/tirx-base/test_tir_ptx_ldmatrix.py +++ b/tests/python/tirx-base/test_tir_ptx_ldmatrix.py @@ -26,15 +26,15 @@ @T.prim_func def ptx_ldmatrix( - A: T.Buffer((16, 16), "float16"), B: T.Buffer((16, 16), "float16"), num: T.int32, trans: T.uint8 + A: T.Tensor((16, 16), "float16"), B: T.Tensor((16, 16), "float16"), num: T.int32, trans: T.uint8 ) -> None: T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) bx = T.env_thread("blockIdx.x") tx = T.env_thread("threadIdx.x") T.launch_thread(bx, 1) T.launch_thread(tx, 32) - A_shared = T.alloc_buffer([16, 16], "float16", scope="shared") - A_local = T.alloc_buffer([8], "float16", scope="local") + A_shared = T.alloc_tensor([16, 16], "float16", scope="shared") + A_local = T.alloc_tensor([8], "float16", scope="local") for i in range(8): A_shared[i * 2 + tx // 16, tx % 16] = A[i * 2 + tx // 16, tx % 16] T.evaluate( diff --git a/tests/python/tirx-base/test_tir_ptx_mma.py b/tests/python/tirx-base/test_tir_ptx_mma.py index 84101b3e60a2..21f8c86ed47f 100644 --- a/tests/python/tirx-base/test_tir_ptx_mma.py +++ b/tests/python/tirx-base/test_tir_ptx_mma.py @@ -26,9 +26,9 @@ @T.prim_func def gemm_mma_m8n8k4_row_col_fp64pf64fp64( - A: T.Buffer([8, 4], dtype="float64"), - B: T.Buffer([8, 4], dtype="float64"), - C: T.Buffer([8, 8], dtype="float64"), + A: T.Tensor([8, 4], dtype="float64"), + B: T.Tensor([8, 4], dtype="float64"), + C: T.Tensor([8, 8], dtype="float64"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -38,9 +38,9 @@ def gemm_mma_m8n8k4_row_col_fp64pf64fp64( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([1], "float64", scope="local") - MultiB = T.decl_buffer([1], "float64", scope="local") - Accum = T.decl_buffer([2], "float64", scope="local") + MultiA = T.decl_tensor([1], "float64", scope="local") + MultiB = T.decl_tensor([1], "float64", scope="local") + Accum = T.decl_tensor([2], "float64", scope="local") for i in range(2): Accum[i] = T.float64(0) @@ -94,9 +94,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m8n8k4_row_row_fp16fp16fp16( - A: T.Buffer([16, 4], dtype="float16"), - B: T.Buffer([4, 16], dtype="float16"), - C: T.Buffer([16, 16], dtype="float16"), + A: T.Tensor([16, 4], dtype="float16"), + B: T.Tensor([4, 16], dtype="float16"), + C: T.Tensor([16, 16], dtype="float16"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -106,9 +106,9 @@ def gemm_mma_m8n8k4_row_row_fp16fp16fp16( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([4], "float16", scope="local") - MultiB = T.decl_buffer([4], "float16", scope="local") - Accum = T.decl_buffer([8], "float16", scope="local") + MultiA = T.decl_tensor([4], "float16", scope="local") + MultiB = T.decl_tensor([4], "float16", scope="local") + Accum = T.decl_tensor([8], "float16", scope="local") for i in range(8): Accum[i] = T.float32(0) @@ -173,9 +173,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m8n8k4_row_row_fp16fp16fp32( - A: T.Buffer([16, 4], dtype="float16"), - B: T.Buffer([4, 16], dtype="float16"), - C: T.Buffer([16, 16], dtype="float32"), + A: T.Tensor([16, 4], dtype="float16"), + B: T.Tensor([4, 16], dtype="float16"), + C: T.Tensor([16, 16], dtype="float32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -185,9 +185,9 @@ def gemm_mma_m8n8k4_row_row_fp16fp16fp32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([4], "float16", scope="local") - MultiB = T.decl_buffer([4], "float16", scope="local") - Accum = T.decl_buffer([8], "float32", scope="local") + MultiA = T.decl_tensor([4], "float16", scope="local") + MultiB = T.decl_tensor([4], "float16", scope="local") + Accum = T.decl_tensor([8], "float32", scope="local") for i in range(8): Accum[i] = T.float32(0) @@ -259,9 +259,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m8n8k16_row_col_s8s8s32( - A: T.Buffer([8, 16], dtype="int8"), - B: T.Buffer([8, 16], dtype="int8"), - C: T.Buffer([8, 8], dtype="int32"), + A: T.Tensor([8, 16], dtype="int8"), + B: T.Tensor([8, 16], dtype="int8"), + C: T.Tensor([8, 8], dtype="int32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -271,9 +271,9 @@ def gemm_mma_m8n8k16_row_col_s8s8s32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([4], "int8", scope="local") - MultiB = T.decl_buffer([4], "int8", scope="local") - Accum = T.decl_buffer([2], "int32", scope="local") + MultiA = T.decl_tensor([4], "int8", scope="local") + MultiB = T.decl_tensor([4], "int8", scope="local") + Accum = T.decl_tensor([2], "int32", scope="local") for i in range(2): Accum[i] = T.int32(0) @@ -333,9 +333,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m8n8k16_row_col_s8u8s32( - A: T.Buffer([8, 16], dtype="int8"), - B: T.Buffer([8, 16], dtype="uint8"), - C: T.Buffer([8, 8], dtype="int32"), + A: T.Tensor([8, 16], dtype="int8"), + B: T.Tensor([8, 16], dtype="uint8"), + C: T.Tensor([8, 8], dtype="int32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -345,9 +345,9 @@ def gemm_mma_m8n8k16_row_col_s8u8s32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([4], "int8", scope="local") - MultiB = T.decl_buffer([4], "uint8", scope="local") - Accum = T.decl_buffer([2], "int32", scope="local") + MultiA = T.decl_tensor([4], "int8", scope="local") + MultiB = T.decl_tensor([4], "uint8", scope="local") + Accum = T.decl_tensor([2], "int32", scope="local") for i in range(2): Accum[i] = T.int32(0) @@ -407,9 +407,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m8n8k32_row_col_s4s4s32( - A: T.Buffer([8, 32], dtype="int4"), - B: T.Buffer([8, 32], dtype="int4"), - C: T.Buffer([8, 8], dtype="int32"), + A: T.Tensor([8, 32], dtype="int4"), + B: T.Tensor([8, 32], dtype="int4"), + C: T.Tensor([8, 8], dtype="int32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -419,9 +419,9 @@ def gemm_mma_m8n8k32_row_col_s4s4s32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([8], "int4", scope="local") - MultiB = T.decl_buffer([8], "int4", scope="local") - Accum = T.decl_buffer([2], "int32", scope="local") + MultiA = T.decl_tensor([8], "int4", scope="local") + MultiB = T.decl_tensor([8], "int4", scope="local") + Accum = T.decl_tensor([2], "int32", scope="local") for i in range(2): Accum[i] = T.int32(0) @@ -475,9 +475,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m8n8k32_row_col_s4u4s32( - A: T.Buffer([8, 32], dtype="int4"), - B: T.Buffer([8, 32], dtype="uint4"), - C: T.Buffer([8, 8], dtype="int32"), + A: T.Tensor([8, 32], dtype="int4"), + B: T.Tensor([8, 32], dtype="uint4"), + C: T.Tensor([8, 8], dtype="int32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -487,9 +487,9 @@ def gemm_mma_m8n8k32_row_col_s4u4s32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([8], "int4", scope="local") - MultiB = T.decl_buffer([8], "uint4", scope="local") - Accum = T.decl_buffer([2], "int32", scope="local") + MultiA = T.decl_tensor([8], "int4", scope="local") + MultiB = T.decl_tensor([8], "uint4", scope="local") + Accum = T.decl_tensor([2], "int32", scope="local") for i in range(2): Accum[i] = T.int32(0) @@ -543,9 +543,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m16n8k8_row_col_fp16fp16fp32( - A: T.Buffer([16, 8], dtype="float16"), - B: T.Buffer([8, 8], dtype="float16"), - C: T.Buffer([16, 8], dtype="float32"), + A: T.Tensor([16, 8], dtype="float16"), + B: T.Tensor([8, 8], dtype="float16"), + C: T.Tensor([16, 8], dtype="float32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -555,9 +555,9 @@ def gemm_mma_m16n8k8_row_col_fp16fp16fp32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([4], "float16", scope="local") - MultiB = T.decl_buffer([2], "float16", scope="local") - Accum = T.decl_buffer([4], "float32", scope="local") + MultiA = T.decl_tensor([4], "float16", scope="local") + MultiB = T.decl_tensor([2], "float16", scope="local") + Accum = T.decl_tensor([4], "float32", scope="local") for i in range(4): Accum[i] = T.float32(0) @@ -619,9 +619,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m16n8k16_row_col_fp16fp16fp16( - A: T.Buffer([16, 16], dtype="float16"), - B: T.Buffer([8, 16], dtype="float16"), - C: T.Buffer([16, 8], dtype="float16"), + A: T.Tensor([16, 16], dtype="float16"), + B: T.Tensor([8, 16], dtype="float16"), + C: T.Tensor([16, 8], dtype="float16"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -631,9 +631,9 @@ def gemm_mma_m16n8k16_row_col_fp16fp16fp16( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([8], "float16", scope="local") - MultiB = T.decl_buffer([4], "float16", scope="local") - Accum = T.decl_buffer([4], "float16", scope="local") + MultiA = T.decl_tensor([8], "float16", scope="local") + MultiB = T.decl_tensor([4], "float16", scope="local") + Accum = T.decl_tensor([4], "float16", scope="local") for i in range(4): Accum[i] = T.float32(0) @@ -698,9 +698,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m16n8k16_row_col_fp16fp16fp32( - A: T.Buffer([16, 16], dtype="float16"), - B: T.Buffer([8, 16], dtype="float16"), - C: T.Buffer([16, 8], dtype="float32"), + A: T.Tensor([16, 16], dtype="float16"), + B: T.Tensor([8, 16], dtype="float16"), + C: T.Tensor([16, 8], dtype="float32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -710,9 +710,9 @@ def gemm_mma_m16n8k16_row_col_fp16fp16fp32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([8], "float16", scope="local") - MultiB = T.decl_buffer([4], "float16", scope="local") - Accum = T.decl_buffer([4], "float32", scope="local") + MultiA = T.decl_tensor([8], "float16", scope="local") + MultiB = T.decl_tensor([4], "float16", scope="local") + Accum = T.decl_tensor([4], "float32", scope="local") for i in range(4): Accum[i] = T.float32(0) @@ -777,9 +777,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m16n8k16_row_col_s8s8s32( - A: T.Buffer([16, 16], dtype="int8"), - B: T.Buffer([8, 16], dtype="int8"), - C: T.Buffer([16, 8], dtype="int32"), + A: T.Tensor([16, 16], dtype="int8"), + B: T.Tensor([8, 16], dtype="int8"), + C: T.Tensor([16, 8], dtype="int32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -789,9 +789,9 @@ def gemm_mma_m16n8k16_row_col_s8s8s32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([8], "int8", scope="local") - MultiB = T.decl_buffer([4], "int8", scope="local") - Accum = T.decl_buffer([4], "int32", scope="local") + MultiA = T.decl_tensor([8], "int8", scope="local") + MultiB = T.decl_tensor([4], "int8", scope="local") + Accum = T.decl_tensor([4], "int32", scope="local") for i in range(4): Accum[i] = T.int32(0) @@ -856,9 +856,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m16n8k16_row_col_s8u8s32( - A: T.Buffer([16, 16], dtype="int8"), - B: T.Buffer([8, 16], dtype="uint8"), - C: T.Buffer([16, 8], dtype="int32"), + A: T.Tensor([16, 16], dtype="int8"), + B: T.Tensor([8, 16], dtype="uint8"), + C: T.Tensor([16, 8], dtype="int32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -868,9 +868,9 @@ def gemm_mma_m16n8k16_row_col_s8u8s32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([8], "int8", scope="local") - MultiB = T.decl_buffer([4], "uint8", scope="local") - Accum = T.decl_buffer([4], "int32", scope="local") + MultiA = T.decl_tensor([8], "int8", scope="local") + MultiB = T.decl_tensor([4], "uint8", scope="local") + Accum = T.decl_tensor([4], "int32", scope="local") for i in range(4): Accum[i] = T.int32(0) @@ -935,9 +935,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m16n8k32_row_col_s8s8s32( - A: T.Buffer([16, 32], dtype="int8"), - B: T.Buffer([8, 32], dtype="int8"), - C: T.Buffer([16, 8], dtype="int32"), + A: T.Tensor([16, 32], dtype="int8"), + B: T.Tensor([8, 32], dtype="int8"), + C: T.Tensor([16, 8], dtype="int32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -947,9 +947,9 @@ def gemm_mma_m16n8k32_row_col_s8s8s32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([16], "int8", scope="local") - MultiB = T.decl_buffer([8], "int8", scope="local") - Accum = T.decl_buffer([4], "int32", scope="local") + MultiA = T.decl_tensor([16], "int8", scope="local") + MultiB = T.decl_tensor([8], "int8", scope="local") + Accum = T.decl_tensor([4], "int32", scope="local") for i in range(4): Accum[i] = T.int32(0) @@ -1014,9 +1014,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m16n8k32_row_col_s8u8s32( - A: T.Buffer([16, 32], dtype="int8"), - B: T.Buffer([8, 32], dtype="uint8"), - C: T.Buffer([16, 8], dtype="int32"), + A: T.Tensor([16, 32], dtype="int8"), + B: T.Tensor([8, 32], dtype="uint8"), + C: T.Tensor([16, 8], dtype="int32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -1026,9 +1026,9 @@ def gemm_mma_m16n8k32_row_col_s8u8s32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([16], "int8", scope="local") - MultiB = T.decl_buffer([8], "uint8", scope="local") - Accum = T.decl_buffer([4], "int32", scope="local") + MultiA = T.decl_tensor([16], "int8", scope="local") + MultiB = T.decl_tensor([8], "uint8", scope="local") + Accum = T.decl_tensor([4], "int32", scope="local") for i in range(4): Accum[i] = T.int32(0) @@ -1093,9 +1093,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m16n8k64_row_col_s4s4s32( - A: T.Buffer([16, 64], dtype="int4"), - B: T.Buffer([8, 64], dtype="int4"), - C: T.Buffer([16, 8], dtype="int32"), + A: T.Tensor([16, 64], dtype="int4"), + B: T.Tensor([8, 64], dtype="int4"), + C: T.Tensor([16, 8], dtype="int32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -1105,9 +1105,9 @@ def gemm_mma_m16n8k64_row_col_s4s4s32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([32], "int4", scope="local") - MultiB = T.decl_buffer([16], "int4", scope="local") - Accum = T.decl_buffer([4], "int32", scope="local") + MultiA = T.decl_tensor([32], "int4", scope="local") + MultiB = T.decl_tensor([16], "int4", scope="local") + Accum = T.decl_tensor([4], "int32", scope="local") for i in range(4): Accum[i] = T.int32(0) @@ -1166,9 +1166,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m16n8k64_row_col_s4u4s32( - A: T.Buffer([16, 64], dtype="int4"), - B: T.Buffer([8, 64], dtype="uint4"), - C: T.Buffer([16, 8], dtype="int32"), + A: T.Tensor([16, 64], dtype="int4"), + B: T.Tensor([8, 64], dtype="uint4"), + C: T.Tensor([16, 8], dtype="int32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -1178,9 +1178,9 @@ def gemm_mma_m16n8k64_row_col_s4u4s32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([32], "int4", scope="local") - MultiB = T.decl_buffer([16], "uint4", scope="local") - Accum = T.decl_buffer([4], "int32", scope="local") + MultiA = T.decl_tensor([32], "int4", scope="local") + MultiB = T.decl_tensor([16], "uint4", scope="local") + Accum = T.decl_tensor([4], "int32", scope="local") for i in range(4): Accum[i] = T.int32(0) @@ -1239,9 +1239,9 @@ def run_and_check(): @T.prim_func def gemm_mma_m16n8k256_row_col_b1b1s32( - A: T.Buffer([16, 256], dtype="int1"), - B: T.Buffer([8, 256], dtype="int1"), - C: T.Buffer([16, 8], dtype="int32"), + A: T.Tensor([16, 256], dtype="int1"), + B: T.Tensor([8, 256], dtype="int1"), + C: T.Tensor([16, 8], dtype="int32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -1251,9 +1251,9 @@ def gemm_mma_m16n8k256_row_col_b1b1s32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - MultiA = T.decl_buffer([128], "int1", scope="local") - MultiB = T.decl_buffer([64], "int1", scope="local") - Accum = T.decl_buffer([4], "int32", scope="local") + MultiA = T.decl_tensor([128], "int1", scope="local") + MultiB = T.decl_tensor([64], "int1", scope="local") + Accum = T.decl_tensor([4], "int32", scope="local") for i in range(4): Accum[i] = T.int32(0) diff --git a/tests/python/tirx-base/test_tir_ptx_mma_sp.py b/tests/python/tirx-base/test_tir_ptx_mma_sp.py index ec0fcf6cf026..2fa4947844f4 100644 --- a/tests/python/tirx-base/test_tir_ptx_mma_sp.py +++ b/tests/python/tirx-base/test_tir_ptx_mma_sp.py @@ -44,10 +44,10 @@ def get_dense_mat_by_mask(val, mask): @T.prim_func def mma_sp_m16n8k16_f16f16f16( - A: T.Buffer([16, 8], dtype="float16"), - B: T.Buffer([16, 8], dtype="float16"), - C: T.Buffer([16, 8], dtype="float16"), - metadata: T.Buffer([8], dtype="uint32"), + A: T.Tensor([16, 8], dtype="float16"), + B: T.Tensor([16, 8], dtype="float16"), + C: T.Tensor([16, 8], dtype="float16"), + metadata: T.Tensor([8], dtype="uint32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -57,10 +57,10 @@ def mma_sp_m16n8k16_f16f16f16( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - multi_a = T.decl_buffer([4], "float16", scope="local") - multi_b = T.decl_buffer([4], "float16", scope="local") - accum = T.decl_buffer([4], "float16", scope="local") - meta_local = T.decl_buffer([1], "uint32", scope="local") + multi_a = T.decl_tensor([4], "float16", scope="local") + multi_b = T.decl_tensor([4], "float16", scope="local") + accum = T.decl_tensor([4], "float16", scope="local") + meta_local = T.decl_tensor([1], "uint32", scope="local") for i in range(4): accum[i] = T.float16(0) @@ -72,9 +72,9 @@ def mma_sp_m16n8k16_f16f16f16( meta_local[0] = metadata[tx // 4] - a_words = T.decl_buffer([2], "uint32", data=multi_a.data, scope="local") - b_words = T.decl_buffer([2], "uint32", data=multi_b.data, scope="local") - acc_words = T.decl_buffer([2], "uint32", data=accum.data, scope="local") + a_words = T.decl_tensor([2], "uint32", data=multi_a.data, scope="local") + b_words = T.decl_tensor([2], "uint32", data=multi_b.data, scope="local") + acc_words = T.decl_tensor([2], "uint32", data=accum.data, scope="local") T.ptx.mma.sp.sync.aligned.m16n8k16.row.col.f16.f16.f16.f16( acc_words[0], acc_words[1], @@ -94,10 +94,10 @@ def mma_sp_m16n8k16_f16f16f16( @T.prim_func def mma_sp_m16n8k16_f16f16f32( - A: T.Buffer([16, 8], dtype="float16"), - B: T.Buffer([16, 8], dtype="float16"), - C: T.Buffer([16, 8], dtype="float32"), - metadata: T.Buffer([8], dtype="uint32"), + A: T.Tensor([16, 8], dtype="float16"), + B: T.Tensor([16, 8], dtype="float16"), + C: T.Tensor([16, 8], dtype="float32"), + metadata: T.Tensor([8], dtype="uint32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -107,10 +107,10 @@ def mma_sp_m16n8k16_f16f16f32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - multi_a = T.decl_buffer([4], "float16", scope="local") - multi_b = T.decl_buffer([4], "float16", scope="local") - accum = T.decl_buffer([4], "float32", scope="local") - meta_local = T.decl_buffer([1], "uint32", scope="local") + multi_a = T.decl_tensor([4], "float16", scope="local") + multi_b = T.decl_tensor([4], "float16", scope="local") + accum = T.decl_tensor([4], "float32", scope="local") + meta_local = T.decl_tensor([1], "uint32", scope="local") for i in range(4): accum[i] = T.float16(0) @@ -122,8 +122,8 @@ def mma_sp_m16n8k16_f16f16f32( meta_local[0] = metadata[tx // 4] - a_words = T.decl_buffer([2], "uint32", data=multi_a.data, scope="local") - b_words = T.decl_buffer([2], "uint32", data=multi_b.data, scope="local") + a_words = T.decl_tensor([2], "uint32", data=multi_a.data, scope="local") + b_words = T.decl_tensor([2], "uint32", data=multi_b.data, scope="local") T.ptx.mma.sp.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32( accum[0], accum[1], @@ -147,10 +147,10 @@ def mma_sp_m16n8k16_f16f16f32( @T.prim_func def mma_sp_m16n8k32_f16f16f16( - A: T.Buffer([16, 16], dtype="float16"), - B: T.Buffer([32, 8], dtype="float16"), - C: T.Buffer([16, 8], dtype="float16"), - metadata: T.Buffer([16], dtype="uint32"), + A: T.Tensor([16, 16], dtype="float16"), + B: T.Tensor([32, 8], dtype="float16"), + C: T.Tensor([16, 8], dtype="float16"), + metadata: T.Tensor([16], dtype="uint32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -160,10 +160,10 @@ def mma_sp_m16n8k32_f16f16f16( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - multi_a = T.decl_buffer([8], "float16", scope="local") - multi_b = T.decl_buffer([8], "float16", scope="local") - accum = T.decl_buffer([4], "float16", scope="local") - meta_local = T.decl_buffer([1], "uint32", scope="local") + multi_a = T.decl_tensor([8], "float16", scope="local") + multi_b = T.decl_tensor([8], "float16", scope="local") + accum = T.decl_tensor([4], "float16", scope="local") + meta_local = T.decl_tensor([1], "uint32", scope="local") for i in range(4): accum[i] = T.float16(0) @@ -175,9 +175,9 @@ def mma_sp_m16n8k32_f16f16f16( meta_local[0] = metadata[tx // 4 * 2 + tx % 2] - a_words = T.decl_buffer([4], "uint32", data=multi_a.data, scope="local") - b_words = T.decl_buffer([4], "uint32", data=multi_b.data, scope="local") - acc_words = T.decl_buffer([2], "uint32", data=accum.data, scope="local") + a_words = T.decl_tensor([4], "uint32", data=multi_a.data, scope="local") + b_words = T.decl_tensor([4], "uint32", data=multi_b.data, scope="local") + acc_words = T.decl_tensor([2], "uint32", data=accum.data, scope="local") T.ptx.mma.sp.sync.aligned.m16n8k32.row.col.f16.f16.f16.f16( acc_words[0], acc_words[1], @@ -201,10 +201,10 @@ def mma_sp_m16n8k32_f16f16f16( @T.prim_func def mma_sp_m16n8k32_f16f16f32( - A: T.Buffer([16, 16], dtype="float16"), - B: T.Buffer([32, 8], dtype="float16"), - C: T.Buffer([16, 8], dtype="float32"), - metadata: T.Buffer([16], dtype="uint32"), + A: T.Tensor([16, 16], dtype="float16"), + B: T.Tensor([32, 8], dtype="float16"), + C: T.Tensor([16, 8], dtype="float32"), + metadata: T.Tensor([16], dtype="uint32"), ): T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) @@ -214,10 +214,10 @@ def mma_sp_m16n8k32_f16f16f32( T.launch_thread(brow, 1) T.launch_thread(bcol, 1) T.launch_thread(tx, 32) - multi_a = T.decl_buffer([8], "float16", scope="local") - multi_b = T.decl_buffer([8], "float16", scope="local") - accum = T.decl_buffer([4], "float32", scope="local") - meta_local = T.decl_buffer([1], "uint32", scope="local") + multi_a = T.decl_tensor([8], "float16", scope="local") + multi_b = T.decl_tensor([8], "float16", scope="local") + accum = T.decl_tensor([4], "float32", scope="local") + meta_local = T.decl_tensor([1], "uint32", scope="local") for i in range(4): accum[i] = T.float16(0) @@ -229,8 +229,8 @@ def mma_sp_m16n8k32_f16f16f32( meta_local[0] = metadata[tx // 4 * 2 + tx % 2] - a_words = T.decl_buffer([4], "uint32", data=multi_a.data, scope="local") - b_words = T.decl_buffer([4], "uint32", data=multi_b.data, scope="local") + a_words = T.decl_tensor([4], "uint32", data=multi_a.data, scope="local") + b_words = T.decl_tensor([4], "uint32", data=multi_b.data, scope="local") T.ptx.mma.sp.sync.aligned.m16n8k32.row.col.f32.f16.f16.f32( accum[0], accum[1], diff --git a/tests/python/tirx-base/test_tir_ptx_scalar_f32_math.py b/tests/python/tirx-base/test_tir_ptx_scalar_f32_math.py index a08356862be7..4e62b402a191 100644 --- a/tests/python/tirx-base/test_tir_ptx_scalar_f32_math.py +++ b/tests/python/tirx-base/test_tir_ptx_scalar_f32_math.py @@ -26,11 +26,11 @@ @T.prim_func def ptx_scalar_f32_math( - A: T.Buffer((32,), "float32"), - B: T.Buffer((32,), "float32"), - C_add: T.Buffer((32,), "float32"), - C_mul: T.Buffer((32,), "float32"), - C_max: T.Buffer((32,), "float32"), + A: T.Tensor((32,), "float32"), + B: T.Tensor((32,), "float32"), + C_add: T.Tensor((32,), "float32"), + C_mul: T.Tensor((32,), "float32"), + C_max: T.Tensor((32,), "float32"), ) -> None: T.func_attr({"global_symbol": "default_function", "tirx.noalias": True}) bx = T.env_thread("blockIdx.x") diff --git a/tests/python/tirx-base/test_tir_specialize.py b/tests/python/tirx-base/test_tir_specialize.py index 2e559080eb25..b410192c7f50 100644 --- a/tests/python/tirx-base/test_tir_specialize.py +++ b/tests/python/tirx-base/test_tir_specialize.py @@ -34,7 +34,7 @@ def assert_structural_equal_ignore_global_symbol(lhs, rhs): @T.prim_func -def matmul(A: T.Buffer([m, n]), B: T.Buffer([m, n]), C: T.Buffer([m, m]), n: T.int32) -> None: +def matmul(A: T.Tensor([m, n]), B: T.Tensor([m, n]), C: T.Tensor([m, m]), n: T.int32) -> None: for i, j, k in T.grid(m, m, n): if k == 0: C[i, j] = 0.0 @@ -42,7 +42,7 @@ def matmul(A: T.Buffer([m, n]), B: T.Buffer([m, n]), C: T.Buffer([m, m]), n: T.i @T.prim_func -def matmul_128(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([128, 128])) -> None: +def matmul_128(A: T.Tensor([128, 128]), B: T.Tensor([128, 128]), C: T.Tensor([128, 128])) -> None: for i, j, k in T.grid(128, 128, 128): if k == 0: C[i, j] = 0.0 @@ -53,7 +53,7 @@ def matmul_128(A: T.Buffer([128, 128]), B: T.Buffer([128, 128]), C: T.Buffer([12 @T.prim_func -def matmul_m_128(A: T.Buffer([m, 128]), B: T.Buffer([m, 128]), C: T.Buffer([m, m])) -> None: +def matmul_m_128(A: T.Tensor([m, 128]), B: T.Tensor([m, 128]), C: T.Tensor([m, m])) -> None: for i, j, k in T.grid(m, m, 128): if k == 0: C[i, j] = 0.0 @@ -67,7 +67,7 @@ def matmul_m_128(A: T.Buffer([m, 128]), B: T.Buffer([m, 128]), C: T.Buffer([m, m @T.prim_func(check_well_formed=False) -def matmul_m_8x(A: T.Buffer([m, x * 8]), B: T.Buffer([m, x * 8]), C: T.Buffer([m, m])) -> None: +def matmul_m_8x(A: T.Tensor([m, x * 8]), B: T.Tensor([m, x * 8]), C: T.Tensor([m, m])) -> None: for i, j, k in T.grid(m, m, x * 8): if k == 0: C[i, j] = 0.0 @@ -79,8 +79,8 @@ def matmul_m_8x(A: T.Buffer([m, x * 8]), B: T.Buffer([m, x * 8]), C: T.Buffer([m @T.prim_func -def element_wise(A: T.Buffer((m, n), "float32"), C: T.Buffer((m, n), "float32")) -> None: - B = T.alloc_buffer((m, n), "float32") +def element_wise(A: T.Tensor((m, n), "float32"), C: T.Tensor((m, n), "float32")) -> None: + B = T.alloc_tensor((m, n), "float32") for i, j in T.grid(m, n): B[i, j] = A[i, j] * 2.0 @@ -91,9 +91,9 @@ def element_wise(A: T.Buffer((m, n), "float32"), C: T.Buffer((m, n), "float32")) @T.prim_func def element_wise_128_64( - A: T.Buffer((128, 64), "float32"), C: T.Buffer((128, 64), "float32") + A: T.Tensor((128, 64), "float32"), C: T.Tensor((128, 64), "float32") ) -> None: - B = T.alloc_buffer((128, 64), "float32") + B = T.alloc_tensor((128, 64), "float32") for i, j in T.grid(128, 64): B[i, j] = A[i, j] * 2.0 @@ -106,8 +106,8 @@ def element_wise_128_64( @T.prim_func -def element_wise_128_n(A: T.Buffer((128, n), "float32"), C: T.Buffer((128, n), "float32")) -> None: - B = T.alloc_buffer((128, n), "float32") +def element_wise_128_n(A: T.Tensor((128, n), "float32"), C: T.Tensor((128, n), "float32")) -> None: + B = T.alloc_tensor((128, n), "float32") for i, j in T.grid(128, n): B[i, j] = A[i, j] * 2.0 @@ -123,8 +123,8 @@ def element_wise_128_n(A: T.Buffer((128, n), "float32"), C: T.Buffer((128, n), " @T.prim_func def mem_copy( - A: T.Buffer((mem_copy_m, mem_copy_n), "float32", strides=[p, 1], elem_offset=q), # noqa: F821 - B: T.Buffer((mem_copy_m, mem_copy_n), "float32", strides=[p, 1], elem_offset=q), # noqa: F821 + A: T.Tensor((mem_copy_m, mem_copy_n), "float32", strides=[p, 1], elem_offset=q), # noqa: F821 + B: T.Tensor((mem_copy_m, mem_copy_n), "float32", strides=[p, 1], elem_offset=q), # noqa: F821 m: mem_copy_m, n: mem_copy_n, p: T.int32, @@ -136,8 +136,8 @@ def mem_copy( @T.prim_func def mem_copy_16_16_8_4( - A: T.Buffer((16, 16), "float32", strides=[8, 1], elem_offset=4), - B: T.Buffer((16, 16), "float32", strides=[8, 1], elem_offset=4), + A: T.Tensor((16, 16), "float32", strides=[8, 1], elem_offset=4), + B: T.Tensor((16, 16), "float32", strides=[8, 1], elem_offset=4), ) -> None: for i, j in T.grid(16, 16): B[i, j] = A[i, j] @@ -150,13 +150,13 @@ def mem_copy_16_16_8_4( @T.prim_func def mem_copy_m_n_p_n( - A: T.Buffer( + A: T.Tensor( (mem_copy_m_n_p_n_m, mem_copy_m_n_p_n_n), "float32", strides=[p, 1], # noqa: F821 elem_offset=mem_copy_m_n_p_n_n, ), - B: T.Buffer( + B: T.Tensor( (mem_copy_m_n_p_n_m, mem_copy_m_n_p_n_n), "float32", strides=[p, 1], # noqa: F821 @@ -178,7 +178,7 @@ def test_specialize_nothing(): def test_specialize_matmul(): a, _, _, n = matmul.params # fully specialized - func = matmul.specialize({a: tvm.tirx.decl_buffer((128, 128))}) + func = matmul.specialize({a: tvm.tirx.decl_tensor((128, 128))}) assert_structural_equal_ignore_global_symbol(func, matmul_128) # partially specialized func = matmul.specialize({n: 128}) @@ -192,17 +192,17 @@ def test_specialize_elemwise(): a, c = element_wise.params C = c # fully specialized - func = element_wise.specialize({a: tvm.tirx.decl_buffer((128, 64))}) + func = element_wise.specialize({a: tvm.tirx.decl_tensor((128, 64))}) assert_structural_equal_ignore_global_symbol(func, element_wise_128_64) # partially specialized - func = element_wise.specialize({c: tvm.tirx.decl_buffer((128, C.ty.shape[1]))}) + func = element_wise.specialize({c: tvm.tirx.decl_tensor((128, C.ty.shape[1]))}) assert_structural_equal_ignore_global_symbol(func, element_wise_128_n) def test_specialize_mem_copy(): a, _, m, n, p, q = mem_copy.params # fully specialized - func = mem_copy.specialize({a: tvm.tirx.decl_buffer((16, 16), strides=[8, 1], elem_offset=4)}) + func = mem_copy.specialize({a: tvm.tirx.decl_tensor((16, 16), strides=[8, 1], elem_offset=4)}) assert_structural_equal_ignore_global_symbol(func, mem_copy_16_16_8_4) func = mem_copy.specialize({n: 16, m: 16, p: 8, q: 4}) assert_structural_equal_ignore_global_symbol(func, mem_copy_16_16_8_4) @@ -220,32 +220,32 @@ def test_specialize_with_const_folding(): n = T.dynamic("n", "int32") @T.prim_func - def before(A: T.Buffer([n // 8, 8], "int32"), B: T.Buffer([n], "int32")): + def before(A: T.Tensor([n // 8, 8], "int32"), B: T.Tensor([n], "int32")): for i in range(n - 1): B[i] = A[i // 8, i % 8] + (n + 1) * 42 @T.prim_func - def expected(A: T.Buffer([2, 8], "int32"), B: T.Buffer([16], "int32")): + def expected(A: T.Tensor([2, 8], "int32"), B: T.Tensor([16], "int32")): for i in range(15): B[i] = A[i // 8, i % 8] + 714 b = before.params[1] - after = before.specialize({b: tvm.tirx.decl_buffer([16], dtype="int32")}) + after = before.specialize({b: tvm.tirx.decl_tensor([16], dtype="int32")}) assert_structural_equal_ignore_global_symbol(expected, after) def test_specialize_decl_buffer(): - """Buffers occurring in a DeclBuffer statement should be updated""" + """Buffers occurring in a DeclTensor statement should be updated""" @T.prim_func(private=True) def before(A_data: T.handle("float32"), A_size: T.int32): - A_buf = T.decl_buffer(A_size, "float32", data=A_data) + A_buf = T.decl_tensor(A_size, "float32", data=A_data) for i in range(A_size): A_buf[i] = A_buf[i] * 2.0 @T.prim_func(private=True) def expected(A_data: T.handle("float32")): - A_buf = T.decl_buffer(16, "float32", data=A_data) + A_buf = T.decl_tensor(16, "float32", data=A_data) for i in range(16): A_buf[i] = A_buf[i] * 2.0 @@ -259,13 +259,13 @@ def test_specialize_preserves_decl_buffer_alias(): before_n = T.int32() @T.prim_func(private=True) - def before(A: T.Buffer((before_n,), "int32"), n: before_n): - A_flat = T.decl_buffer((n,), "int32", data=A.data) + def before(A: T.Tensor((before_n,), "int32"), n: before_n): + A_flat = T.decl_tensor((n,), "int32", data=A.data) A_flat[n - 1] = 42 @T.prim_func(private=True) - def expected(A: T.Buffer((8,), "int32")): - A_flat = T.decl_buffer((8,), "int32", data=A.data) + def expected(A: T.Tensor((8,), "int32")): + A_flat = T.decl_tensor((8,), "int32", data=A.data) A_flat[7] = 42 after = before.specialize({before.params[1]: 8}) @@ -281,16 +281,16 @@ def test_specialize_buffer_var_to_var(): """ @T.prim_func(private=True) - def before(A: T.Buffer([16, 16], "float32"), B: T.Buffer([16, 16], "float32")): - A_flat = T.decl_buffer([256], "float32", data=A.data) - B_flat = T.decl_buffer([256], "float32", data=B.data) + def before(A: T.Tensor([16, 16], "float32"), B: T.Tensor([16, 16], "float32")): + A_flat = T.decl_tensor([256], "float32", data=A.data) + B_flat = T.decl_tensor([256], "float32", data=B.data) for i in range(256): B_flat[i] = A_flat[i] * 2.0 @T.prim_func(private=True) - def expected(A: T.Buffer([16, 16], "float32")): - A_flat = T.decl_buffer([256], "float32", data=A.data) - B_flat = T.decl_buffer([256], "float32", data=A.data) + def expected(A: T.Tensor([16, 16], "float32")): + A_flat = T.decl_tensor([256], "float32", data=A.data) + B_flat = T.decl_tensor([256], "float32", data=A.data) for i in range(256): B_flat[i] = A_flat[i] * 2.0 @@ -303,24 +303,24 @@ def expected(A: T.Buffer([16, 16], "float32")): def test_specialize_buffer_var_to_expr(): - """A DeclBuffer source expression may be specialized directly.""" + """A DeclTensor source expression may be specialized directly.""" @T.prim_func(private=True) def before(A_data: T.handle("float32"), B_data: T.handle("float32")): - A_buf = T.decl_buffer(32, "float32", data=A_data) - B_buf = T.decl_buffer(16, "float32", data=B_data) + A_buf = T.decl_tensor(32, "float32", data=A_data) + B_buf = T.decl_tensor(16, "float32", data=B_data) for i in range(16): B_buf[i] = A_buf[i] * 2.0 @T.prim_func(private=True) def expected(A_data: T.handle("float32")): - A_buf = T.decl_buffer(32, "float32", data=A_data) - B_buf = T.decl_buffer(16, "float32", data=T.address_of(A_buf[16])) + A_buf = T.decl_tensor(32, "float32", data=A_data) + B_buf = T.decl_tensor(16, "float32", data=T.address_of(A_buf[16])) for i in range(16): B_buf[i] = A_buf[i] * 2.0 B_data = before.params[1] - # body is a SeqStmt; the first statement is DeclBuffer for A_buf + # body is a SeqStmt; the first statement is DeclTensor for A_buf A_buf = before.body[0].var param_map = {B_data: tvm.tirx.address_of(A_buf[16])} after = before.specialize(param_map) diff --git a/tests/python/tirx-base/test_tir_stmt_functor_substitute.py b/tests/python/tirx-base/test_tir_stmt_functor_substitute.py index 9bcbadb0e2a8..19738014e5e4 100644 --- a/tests/python/tirx-base/test_tir_stmt_functor_substitute.py +++ b/tests/python/tirx-base/test_tir_stmt_functor_substitute.py @@ -62,14 +62,14 @@ def test_substitute_allocate(): class Before: @T.prim_func def main(n: T.int32): - A = T.alloc_buffer((n,), "float32") + A = T.alloc_tensor((n,), "float32") T.evaluate(A.data) @I.ir_module class Expected: @T.prim_func def main(): - A = T.alloc_buffer((16,), "float32") + A = T.alloc_tensor((16,), "float32") T.evaluate(A.data) After = _apply_substitute(Before) @@ -81,7 +81,7 @@ def test_substitute_buffer_load(): class Before: @T.prim_func def main(n: T.int32): - A = T.alloc_buffer((n,), "float32") + A = T.alloc_tensor((n,), "float32") for i in range(n): T.evaluate(A[i]) @@ -89,7 +89,7 @@ def main(n: T.int32): class Expected: @T.prim_func def main(): - A = T.alloc_buffer((16,), "float32") + A = T.alloc_tensor((16,), "float32") for i in range(16): T.evaluate(A[i]) @@ -102,14 +102,14 @@ def test_substitute_decl_buffer(): class Before: @T.prim_func def main(n: T.int32): - A = T.alloc_buffer((n,), "float32") + A = T.alloc_tensor((n,), "float32") T.evaluate(A.data) @I.ir_module class Expected: @T.prim_func def main(): - A = T.alloc_buffer((16,), "float32") + A = T.alloc_tensor((16,), "float32") T.evaluate(A.data) After = _apply_substitute(Before) diff --git a/tests/python/tirx-transform/test_tir_inline_private_functions.py b/tests/python/tirx-transform/test_tir_inline_private_functions.py index 6f0c24ea7bd5..54a381aefcbb 100644 --- a/tests/python/tirx-transform/test_tir_inline_private_functions.py +++ b/tests/python/tirx-transform/test_tir_inline_private_functions.py @@ -39,14 +39,14 @@ class TestSimple(BaseTestCase): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer([80, 16], "float32"), B: T.Buffer([64, 16], "float32")): + def main(A: T.Tensor([80, 16], "float32"), B: T.Tensor([64, 16], "float32")): for i in range(64): Before.subroutine(T.address_of(A[i, 0]), T.address_of(B[i, 0])) @T.prim_func(private=True) def subroutine(A_data: T.handle("float32"), B_data: T.handle("float32")): - A = T.decl_buffer([16, 16], "float32", data=A_data) - B = T.decl_buffer([16], "float32", data=B_data) + A = T.decl_tensor([16, 16], "float32", data=A_data) + B = T.decl_tensor([16], "float32", data=B_data) for i in range(16): B[i] = 0.0 for j in range(16): @@ -55,10 +55,10 @@ def subroutine(A_data: T.handle("float32"), B_data: T.handle("float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer([80, 16], "float32"), B: T.Buffer([64, 16], "float32")): + def main(A: T.Tensor([80, 16], "float32"), B: T.Tensor([64, 16], "float32")): for i in range(64): - Aview = T.decl_buffer([16, 16], "float32", data=T.address_of(A[i, 0])) - Bview = T.decl_buffer([16], "float32", data=T.address_of(B[i, 0])) + Aview = T.decl_tensor([16, 16], "float32", data=T.address_of(A[i, 0])) + Bview = T.decl_tensor([16], "float32", data=T.address_of(B[i, 0])) for j in range(16): Bview[j] = 0.0 for k in range(16): @@ -78,7 +78,7 @@ class TestRetainCrossFunctionSubroutines(BaseTestCase): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer([80, 16], "float32"), B: T.Buffer([64, 16], "float32")): + def main(A: T.Tensor([80, 16], "float32"), B: T.Tensor([64, 16], "float32")): T.func_attr({"target": T.target("llvm")}) for i in range(64): Before.subroutine(T.address_of(A[i, 0]), T.address_of(B[i, 0])) @@ -86,8 +86,8 @@ def main(A: T.Buffer([80, 16], "float32"), B: T.Buffer([64, 16], "float32")): @T.prim_func(private=True) def subroutine(A_data: T.handle("float32"), B_data: T.handle("float32")): T.func_attr({"target": T.target({"kind": "cuda", "arch": "sm_80"})}) - A = T.decl_buffer([16, 16], "float32", data=A_data) - B = T.decl_buffer([16], "float32", data=B_data) + A = T.decl_tensor([16, 16], "float32", data=A_data) + B = T.decl_tensor([16], "float32", data=B_data) for i in range(16): B[i] = 0.0 for j in range(16): @@ -108,12 +108,12 @@ class TestRetainRecursiveSubroutines(BaseTestCase): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(16, "float32")): + def main(A: T.Tensor(16, "float32")): Before.subroutine(T.address_of(A[0]), 16) @T.prim_func(private=True) def subroutine(A_data: T.handle("float32"), A_size: T.int32): - A = T.decl_buffer(A_size, "float32", data=A_data) + A = T.decl_tensor(A_size, "float32", data=A_data) A[1] = A[0] + A[1] if A_size > 1: @@ -132,32 +132,32 @@ def test_produces_expected(self): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer([2, 16], "float32"), B: T.Buffer([2, 16], "float32")): + def main(A: T.Tensor([2, 16], "float32"), B: T.Tensor([2, 16], "float32")): Before.subroutine(T.address_of(A[0, 0]), T.address_of(B[0, 0])) Before.subroutine(T.address_of(A[1, 0]), T.address_of(B[1, 0])) @T.prim_func(private=True) def subroutine(A_data: T.handle("float32"), B_data: T.handle("float32")): - A = T.decl_buffer(16, "float32", data=A_data) - B = T.decl_buffer(16, "float32", data=B_data) + A = T.decl_tensor(16, "float32", data=A_data) + B = T.decl_tensor(16, "float32", data=B_data) for i in range(16): B[i] = A[i] * 2.0 @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer([80, 16], "float32"), B: T.Buffer([64, 16], "float32")): + def main(A: T.Tensor([80, 16], "float32"), B: T.Tensor([64, 16], "float32")): A_data_1 = T.bind(T.address_of(A[0, 0]), T.handle("float32")) - A_1 = T.decl_buffer(16, "float32", data=A_data_1) + A_1 = T.decl_tensor(16, "float32", data=A_data_1) B_data_1: T.let[T.handle("float32")] = T.address_of(B[0, 0]) - B_1 = T.decl_buffer(16, "float32", data=B_data_1) + B_1 = T.decl_tensor(16, "float32", data=B_data_1) for i in range(16): B_1[i] = A_1[i] * 2.0 A_data_2 = T.bind(T.address_of(A[1, 0]), T.handle("float32")) - A_2 = T.decl_buffer(16, "float32", data=A_data_2) + A_2 = T.decl_tensor(16, "float32", data=A_data_2) B_data_2: T.let[T.handle("float32")] = T.address_of(B[1, 0]) - B_2 = T.decl_buffer(16, "float32", data=B_data_2) + B_2 = T.decl_tensor(16, "float32", data=B_data_2) for i in range(16): B_2[i] = A_2[i] * 2.0 @@ -181,7 +181,7 @@ def test_produces_expected(self): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(16, "float32")): + def main(A: T.Tensor(16, "float32")): for i in range(16): A[i] = Before.subroutine(i) @@ -195,7 +195,7 @@ def subroutine(i: T.int32) -> T.float32: @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(16, "float32")): + def main(A: T.Tensor(16, "float32")): for i in range(16): cos = T.cos(T.cast(i, "float32")) sin = T.sin(T.cast(i, "float32")) @@ -220,7 +220,7 @@ def test_produces_expected(self): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(16, "float32")): + def main(A: T.Tensor(16, "float32")): Before.subroutine( T.tvm_stack_make_array( A.data, @@ -233,14 +233,14 @@ def main(A: T.Buffer(16, "float32")): ) @T.prim_func(private=True) - def subroutine(A: T.Buffer(16, "float32")): + def subroutine(A: T.Tensor(16, "float32")): for i in range(16): A[i] = A[i] * 2.0 @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(16, "float32")): + def main(A: T.Tensor(16, "float32")): for i in range(16): A[i] = A[i] * 2.0 diff --git a/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py b/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py index cb1fd5f2adcb..d11d11f2ac6f 100644 --- a/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py +++ b/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py @@ -64,9 +64,9 @@ def main( Cptr: T.handle("bfloat16"), ): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), "bfloat16", data=Aptr) - B = T.decl_buffer((100,), "bfloat16") - C = T.decl_buffer((100,), "bfloat16", data=Cptr) + A = T.decl_tensor((100,), "bfloat16", data=Aptr) + B = T.decl_tensor((100,), "bfloat16") + C = T.decl_tensor((100,), "bfloat16", data=Cptr) for i in T.grid(100): B[i] = A[i] C[i] = T.exp(B[i]) @@ -82,9 +82,9 @@ def main( Cptr: T.handle("bfloat16"), ): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), "bfloat16", data=Aptr) - B = T.decl_buffer((100,), "float32") - C = T.decl_buffer((100,), "bfloat16", data=Cptr) + A = T.decl_tensor((100,), "bfloat16", data=Aptr) + B = T.decl_tensor((100,), "float32") + C = T.decl_tensor((100,), "bfloat16", data=Cptr) for i in T.grid(100): B[i] = bf16tof32(A[i]) C[i] = f32tobf16(T.exp(B[i])) @@ -100,9 +100,9 @@ def main( Cptr: T.handle("uint16"), ): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), "uint16", data=Aptr) - B = T.decl_buffer((100,), "float32") - C = T.decl_buffer((100,), "uint16", data=Cptr) + A = T.decl_tensor((100,), "uint16", data=Aptr) + B = T.decl_tensor((100,), "float32") + C = T.decl_tensor((100,), "uint16", data=Cptr) for i in T.grid(100): B[i] = u16tof32(A[i]) C[i] = f32tou16(T.exp(B[i])) @@ -124,9 +124,9 @@ class Before: @T.prim_func def main(Aptr: T.handle("bfloat16"), Cptr: T.handle("bfloat16")): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((16,), "bfloat16", data=Aptr) - B = T.decl_buffer((16,), "bfloat16") - C = T.decl_buffer((16,), "bfloat16", data=Cptr) + A = T.decl_tensor((16,), "bfloat16", data=Aptr) + B = T.decl_tensor((16,), "bfloat16") + C = T.decl_tensor((16,), "bfloat16", data=Cptr) mask = T.local_scalar("boolx4") mask = T.Broadcast(T.bool(True), 4) T.evaluate( @@ -171,7 +171,7 @@ def collect_once(node): buffers = { node.var.name: str(node.var.dtype) for node in nodes - if _is_buffer_binding(node, "tirx.alloc_buffer", "tirx.decl_buffer") + if _is_buffer_binding(node, "tirx.alloc_tensor", "tirx.decl_tensor") } masked_loads = [ node @@ -217,10 +217,10 @@ def main( Dptr: T.handle("bfloat16"), ): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), "bfloat16", data=Aptr) - B = T.decl_buffer((100,), "bfloat16", data=Bptr) - D = T.decl_buffer((100,), "bfloat16", data=Dptr) - C = T.decl_buffer((100,), "bfloat16") + A = T.decl_tensor((100,), "bfloat16", data=Aptr) + B = T.decl_tensor((100,), "bfloat16", data=Bptr) + D = T.decl_tensor((100,), "bfloat16", data=Dptr) + C = T.decl_tensor((100,), "bfloat16") for i in T.grid(100): C[i] = A[i] + B[i] D[i] = T.exp(C[i]) @@ -237,10 +237,10 @@ def main( Dptr: T.handle("bfloat16"), ): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), "bfloat16", data=Aptr) - B = T.decl_buffer((100,), "bfloat16", data=Bptr) - D = T.decl_buffer((100,), "bfloat16", data=Dptr) - C = T.decl_buffer((100,), "float32") + A = T.decl_tensor((100,), "bfloat16", data=Aptr) + B = T.decl_tensor((100,), "bfloat16", data=Bptr) + D = T.decl_tensor((100,), "bfloat16", data=Dptr) + C = T.decl_tensor((100,), "float32") for i in T.grid(100): C[i] = bf16tof32(A[i]) + bf16tof32(B[i]) D[i] = f32tobf16(T.exp(C[i])) @@ -257,10 +257,10 @@ def main( Dptr: T.handle("uint16"), ): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), "uint16", data=Aptr) - B = T.decl_buffer((100,), "uint16", data=Bptr) - D = T.decl_buffer((100,), "uint16", data=Dptr) - C = T.decl_buffer((100,), "float32") + A = T.decl_tensor((100,), "uint16", data=Aptr) + B = T.decl_tensor((100,), "uint16", data=Bptr) + D = T.decl_tensor((100,), "uint16", data=Dptr) + C = T.decl_tensor((100,), "float32") for i in T.grid(100): C[i] = u16tof32(A[i]) + u16tof32(B[i]) D[i] = f32tou16(T.exp(C[i])) @@ -286,10 +286,10 @@ def main( Dptr: T.handle("bfloat16"), ): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), "bfloat16", data=Aptr) - B = T.decl_buffer((100,), "bfloat16", data=Bptr) - D = T.decl_buffer((100,), "bfloat16", data=Dptr) - C = T.decl_buffer((100,), "bfloat16") + A = T.decl_tensor((100,), "bfloat16", data=Aptr) + B = T.decl_tensor((100,), "bfloat16", data=Bptr) + D = T.decl_tensor((100,), "bfloat16", data=Dptr) + C = T.decl_tensor((100,), "bfloat16") for i in T.grid(100): C[i] = A[i] + B[i] D[i] = T.exp(C[i]) @@ -306,10 +306,10 @@ def main( Dptr: T.handle("bfloat16"), ): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), "bfloat16", data=Aptr) - B = T.decl_buffer((100,), "bfloat16", data=Bptr) - D = T.decl_buffer((100,), "bfloat16", data=Dptr) - C = T.decl_buffer((100,), "bfloat16") + A = T.decl_tensor((100,), "bfloat16", data=Aptr) + B = T.decl_tensor((100,), "bfloat16", data=Bptr) + D = T.decl_tensor((100,), "bfloat16", data=Dptr) + C = T.decl_tensor((100,), "bfloat16") for i in T.grid(100): C[i] = A[i] + B[i] D[i] = T.exp(C[i]) @@ -326,10 +326,10 @@ def main( Dptr: T.handle("bfloat16"), ): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), "bfloat16", data=Aptr) - B = T.decl_buffer((100,), "bfloat16", data=Bptr) - D = T.decl_buffer((100,), "bfloat16", data=Dptr) - C = T.decl_buffer((100,), "bfloat16") + A = T.decl_tensor((100,), "bfloat16", data=Aptr) + B = T.decl_tensor((100,), "bfloat16", data=Bptr) + D = T.decl_tensor((100,), "bfloat16", data=Dptr) + C = T.decl_tensor((100,), "bfloat16") for i in T.grid(100): C[i] = A[i] + B[i] D[i] = T.exp(C[i]) @@ -352,12 +352,12 @@ class Before: def main( Aptr: T.handle("bfloat16", storage_scope="shared"), ): - A_flat = T.decl_buffer(4096, "bfloat16", data=Aptr) + A_flat = T.decl_tensor(4096, "bfloat16", data=Aptr) for i in range(128): threadIdx_x = T.launch_thread("threadIdx.x", 32) - reduce = T.decl_buffer(1, dtype="bfloat16", scope="local") + reduce = T.decl_tensor(1, dtype="bfloat16", scope="local") with T.attr( T.comm_reducer(lambda x, y: x + y, [T.bfloat16(0)]), @@ -381,12 +381,12 @@ class After: def main( Aptr: T.handle("bfloat16", storage_scope="shared"), ): - A_flat_1 = T.decl_buffer(4096, "bfloat16", data=Aptr) + A_flat_1 = T.decl_tensor(4096, "bfloat16", data=Aptr) for i in range(128): threadIdx_x = T.launch_thread("threadIdx.x", 32) - reduce = T.decl_buffer(1, dtype="float32", scope="local") + reduce = T.decl_tensor(1, dtype="float32", scope="local") with T.attr( T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), @@ -416,12 +416,12 @@ class After: def main( Aptr: T.handle("uint16", storage_scope="shared"), ): - A_flat = T.decl_buffer(4096, "uint16", data=Aptr) + A_flat = T.decl_tensor(4096, "uint16", data=Aptr) for i in range(128): threadIdx_x = T.launch_thread("threadIdx.x", 32) - reduce = T.decl_buffer(1, dtype="float32", scope="local") + reduce = T.decl_tensor(1, dtype="float32", scope="local") with T.attr( T.comm_reducer(lambda x, y: x + y, [T.float32(0)]), @@ -460,12 +460,12 @@ class Before: def main( Aptr: T.handle("bfloat16", storage_scope="shared"), ): - A_flat = T.decl_buffer(4096, "bfloat16", data=Aptr) + A_flat = T.decl_tensor(4096, "bfloat16", data=Aptr) for i in range(128): threadIdx_x = T.launch_thread("threadIdx.x", 32) - reduce = T.decl_buffer(1, dtype="bfloat16", scope="local") + reduce = T.decl_tensor(1, dtype="bfloat16", scope="local") with T.attr( T.comm_reducer(lambda x, y: x + y, [T.bfloat16(0)]), @@ -489,12 +489,12 @@ class After: def main( Aptr: T.handle("bfloat16", storage_scope="shared"), ): - A_flat = T.decl_buffer(4096, "bfloat16", data=Aptr) + A_flat = T.decl_tensor(4096, "bfloat16", data=Aptr) for i in range(128): threadIdx_x = T.launch_thread("threadIdx.x", 32) - reduce = T.decl_buffer(1, dtype="bfloat16", scope="local") + reduce = T.decl_tensor(1, dtype="bfloat16", scope="local") with T.attr( T.comm_reducer(lambda x, y: x + y, [T.bfloat16(0)]), @@ -518,12 +518,12 @@ class After: def main( Aptr: T.handle("bfloat16", storage_scope="shared"), ): - A_flat = T.decl_buffer(4096, "bfloat16", data=Aptr) + A_flat = T.decl_tensor(4096, "bfloat16", data=Aptr) for i in range(128): threadIdx_x = T.launch_thread("threadIdx.x", 32) - reduce = T.decl_buffer(1, dtype="bfloat16", scope="local") + reduce = T.decl_tensor(1, dtype="bfloat16", scope="local") with T.attr( T.comm_reducer(lambda x, y: x + y, [T.bfloat16(0)]), diff --git a/tests/python/tirx-transform/test_tir_transform_common_subexpr_elim.py b/tests/python/tirx-transform/test_tir_transform_common_subexpr_elim.py index 8fe88fd895ef..84fb94078cb9 100644 --- a/tests/python/tirx-transform/test_tir_transform_common_subexpr_elim.py +++ b/tests/python/tirx-transform/test_tir_transform_common_subexpr_elim.py @@ -29,7 +29,7 @@ def test_basic(): @tvm.script.ir_module class Before: @T.prim_func - def main(B: T.Buffer((50,), "int32"), i1: T.int32, i2: T.int32, z3: T.int32): + def main(B: T.Tensor((50,), "int32"), i1: T.int32, i2: T.int32, z3: T.int32): z1 = T.bind(1) z2 = T.bind(2) B[i1] = z1 + z2 @@ -42,7 +42,7 @@ def main(B: T.Buffer((50,), "int32"), i1: T.int32, i2: T.int32, z3: T.int32): @tvm.script.ir_module class Expected: @T.prim_func - def main(B: T.Buffer((50,), "int32"), i1: T.int32, i2: T.int32, z3: T.int32): + def main(B: T.Tensor((50,), "int32"), i1: T.int32, i2: T.int32, z3: T.int32): z1 = T.bind(1) z2 = T.bind(2) cse_v1 = T.bind(z1 + z2) @@ -67,7 +67,7 @@ def test_if_single_branch(): class Before: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), i1: T.int32, i2: T.int32, i3: T.int32, @@ -85,7 +85,7 @@ def main( class Expected: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), i1: T.int32, i2: T.int32, i3: T.int32, @@ -113,7 +113,7 @@ def test_if_both_branches(): class Before: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), i1: T.int32, i2: T.int32, i3: T.int32, @@ -131,7 +131,7 @@ def main( class Expected: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), i1: T.int32, i2: T.int32, i3: T.int32, @@ -159,7 +159,7 @@ def test_cascade(): class Before: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), i1: T.int32, i2: T.int32, i3: T.int32, @@ -175,7 +175,7 @@ def main( class Expected: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), i1: T.int32, i2: T.int32, i3: T.int32, @@ -257,7 +257,7 @@ def test_for_loop(): @tvm.script.ir_module class Before: @T.prim_func - def main(B: T.Buffer((50,), "int32"), y: T.int32, z: T.int32): + def main(B: T.Tensor((50,), "int32"), y: T.int32, z: T.int32): for i in range(10): B[i] = y + z B[i + 10] = y + z @@ -265,7 +265,7 @@ def main(B: T.Buffer((50,), "int32"), y: T.int32, z: T.int32): @tvm.script.ir_module class Expected: @T.prim_func - def main(B: T.Buffer((50,), "int32"), y: T.int32, z: T.int32): + def main(B: T.Tensor((50,), "int32"), y: T.int32, z: T.int32): for i in range(10): cse_v1 = T.bind(y + z) B[i] = cse_v1 @@ -284,7 +284,7 @@ def test_for_hoist(): @tvm.script.ir_module class Before: @T.prim_func - def main(B: T.Buffer((50,), "int32"), y: T.int32, z: T.int32): + def main(B: T.Tensor((50,), "int32"), y: T.int32, z: T.int32): B[0] = y + z for i in range(10): B[i + 1] = y + z @@ -292,7 +292,7 @@ def main(B: T.Buffer((50,), "int32"), y: T.int32, z: T.int32): @tvm.script.ir_module class Expected: @T.prim_func - def main(B: T.Buffer((50,), "int32"), y: T.int32, z: T.int32): + def main(B: T.Tensor((50,), "int32"), y: T.int32, z: T.int32): cse_v1 = T.bind(y + z) B[0] = cse_v1 for i in range(10): @@ -311,14 +311,14 @@ def test_cannot_lift_bufferload(): @tvm.script.ir_module class Before: @T.prim_func - def main(A: T.Buffer((50,), "int32"), B: T.Buffer((50,), "int32")): + def main(A: T.Tensor((50,), "int32"), B: T.Tensor((50,), "int32")): B[0] = A[0] + A[0] B[1] = A[0] + A[0] @tvm.script.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((50,), "int32"), B: T.Buffer((50,), "int32")): + def main(A: T.Tensor((50,), "int32"), B: T.Tensor((50,), "int32")): B[0] = A[0] + A[0] B[1] = A[0] + A[0] @@ -336,7 +336,7 @@ def test_nested_if(): class Before: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), c1: T.int32, c2: T.int32, y: T.int32, @@ -354,7 +354,7 @@ def main( class Expected: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), c1: T.int32, c2: T.int32, y: T.int32, @@ -382,7 +382,7 @@ def test_multi_independent(): class Before: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), a: T.int32, b: T.int32, c: T.int32, @@ -397,7 +397,7 @@ def main( class Expected: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), a: T.int32, b: T.int32, c: T.int32, @@ -423,14 +423,14 @@ def test_if_condition(): @tvm.script.ir_module class Before: @T.prim_func - def main(B: T.Buffer((50,), "int32"), y: T.int32, z: T.int32): + def main(B: T.Tensor((50,), "int32"), y: T.int32, z: T.int32): if y + z > 0: B[0] = y + z @tvm.script.ir_module class Expected: @T.prim_func - def main(B: T.Buffer((50,), "int32"), y: T.int32, z: T.int32): + def main(B: T.Tensor((50,), "int32"), y: T.int32, z: T.int32): cse_v1 = T.bind(y + z) if cse_v1 > 0: B[0] = cse_v1 @@ -447,14 +447,14 @@ def test_cannot_lift_call(): @tvm.script.ir_module class Before: @T.prim_func - def main(B: T.Buffer((50,), "int32"), x: T.int32): + def main(B: T.Tensor((50,), "int32"), x: T.int32): B[0] = T.call_extern("my_func", x, dtype="int32") + 1 B[1] = T.call_extern("my_func", x, dtype="int32") + 1 @tvm.script.ir_module class Expected: @T.prim_func - def main(B: T.Buffer((50,), "int32"), x: T.int32): + def main(B: T.Tensor((50,), "int32"), x: T.int32): B[0] = T.call_extern("my_func", x, dtype="int32") + 1 B[1] = T.call_extern("my_func", x, dtype="int32") + 1 @@ -473,7 +473,7 @@ def test_no_single_use_binding(): class Before: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), x: T.int32, y: T.int32, z: T.int32, @@ -485,7 +485,7 @@ def main( class Expected: @T.prim_func def main( - B: T.Buffer((50,), "int32"), + B: T.Tensor((50,), "int32"), x: T.int32, y: T.int32, z: T.int32, @@ -507,14 +507,14 @@ def test_for_extent_lift(): @tvm.script.ir_module class Before: @T.prim_func - def main(B: T.Buffer((50,), "int32"), y: T.int32, z: T.int32): + def main(B: T.Tensor((50,), "int32"), y: T.int32, z: T.int32): for i in range(y + z): B[i] = y + z @tvm.script.ir_module class Expected: @T.prim_func - def main(B: T.Buffer((50,), "int32"), y: T.int32, z: T.int32): + def main(B: T.Tensor((50,), "int32"), y: T.int32, z: T.int32): cse_v1 = T.bind(y + z) for i in range(cse_v1): B[i] = cse_v1 @@ -533,8 +533,8 @@ def test_loop_var_expr_stays_inside(): class Before: @T.prim_func def main( - A: T.Buffer((50,), "int32"), - B: T.Buffer((50,), "int32"), + A: T.Tensor((50,), "int32"), + B: T.Tensor((50,), "int32"), ): for i in range(10): A[i * 4] = B[i * 4] @@ -543,8 +543,8 @@ def main( class Expected: @T.prim_func def main( - A: T.Buffer((50,), "int32"), - B: T.Buffer((50,), "int32"), + A: T.Tensor((50,), "int32"), + B: T.Tensor((50,), "int32"), ): for i in range(10): cse_v1 = T.bind(i * 4) @@ -589,7 +589,7 @@ def test_let_body_no_extraction(): y = tvm.tirx.Var("y", "int32") # Let(x, 1, (x+y) + (x+y)) -- x+y appears twice but x is Let-bound let_expr = tvm.tirx.Let(x, tvm.tirx.IntImm("int32", 1), (x + y) + (x + y)) - buf = tvm.tirx.decl_buffer((10,), "int32", name="B") + buf = tvm.tirx.decl_tensor((10,), "int32", name="B") i = tvm.tirx.Var("i", "int32") store = tvm.tirx.BufferStore(buf, let_expr, [i]) loop = tvm.tirx.For( @@ -619,7 +619,7 @@ def test_let_value_cse(): z = tvm.tirx.Var("z", "int32") # Let(x, y+z, x+1) with y+z also appearing outside the Let let_expr = tvm.tirx.Let(x, y + z, x + 1) - buf = tvm.tirx.decl_buffer((10,), "int32", name="B") + buf = tvm.tirx.decl_tensor((10,), "int32", name="B") i = tvm.tirx.Var("i", "int32") store = tvm.tirx.BufferStore(buf, (y + z) + let_expr, [i]) loop = tvm.tirx.For( @@ -652,7 +652,7 @@ def test_nested_let_no_extraction(): nested_let = tvm.tirx.Let( x, tvm.tirx.IntImm("int32", 1), tvm.tirx.Let(y, tvm.tirx.IntImm("int32", 2), inner) ) - buf = tvm.tirx.decl_buffer((10,), "int32", name="B") + buf = tvm.tirx.decl_tensor((10,), "int32", name="B") i = tvm.tirx.Var("i", "int32") store = tvm.tirx.BufferStore(buf, nested_let, [i]) loop = tvm.tirx.For( @@ -688,9 +688,9 @@ def test_let_floordiv_pattern(): inner_let = tvm.tirx.Let(rdiv, tvm.tirx.Div(x, y), select_expr) outer_let = tvm.tirx.Let(rmod, tvm.tirx.Mod(x, y), inner_let) # Wrap in Let(x, load, Let(y, load, ...)) - buf_a = tvm.tirx.decl_buffer((10,), "int32", name="A") - buf_b = tvm.tirx.decl_buffer((10,), "int32", name="B") - buf_c = tvm.tirx.decl_buffer((10,), "int32", name="C") + buf_a = tvm.tirx.decl_tensor((10,), "int32", name="A") + buf_b = tvm.tirx.decl_tensor((10,), "int32", name="B") + buf_c = tvm.tirx.decl_tensor((10,), "int32", name="C") i = tvm.tirx.Var("i", "int32") full_expr = tvm.tirx.Let( x, @@ -722,7 +722,7 @@ def test_no_lift_bool_predicate(): @tvm.script.ir_module class Before: @T.prim_func - def main(B: T.Buffer((50,), "int32"), n: T.int32, x: T.int32): + def main(B: T.Tensor((50,), "int32"), n: T.int32, x: T.int32): for i in range(50): if i < n: B[i] = x @@ -743,7 +743,7 @@ def test_no_lift_bool_logical(): @tvm.script.ir_module class Before: @T.prim_func - def main(B: T.Buffer((50,), "int32"), a: T.bool, b: T.bool, x: T.int32): + def main(B: T.Tensor((50,), "int32"), a: T.bool, b: T.bool, x: T.int32): if T.And(a, b): B[0] = x if T.And(a, b): @@ -760,7 +760,7 @@ def test_shared_subtree_stays_ssa(): @tvm.script.ir_module class Payload: @T.prim_func - def main(B: T.Buffer((50,), "int32"), i1: T.int32, i2: T.int32): + def main(B: T.Tensor((50,), "int32"), i1: T.int32, i2: T.int32): B[i1] = (i1 + i2) * 2 B[i2] = (i1 + i2) * 3 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 e27ab33afe24..143bee7985fc 100644 --- a/tests/python/tirx-transform/test_tir_transform_convert_ssa.py +++ b/tests/python/tirx-transform/test_tir_transform_convert_ssa.py @@ -171,7 +171,7 @@ def test_reused_buffer_obj(): @T.prim_func(private=True) def func(a: T.handle("float32")): - A = T.decl_buffer(shape=1, dtype="float32", data=a) + A = T.decl_tensor(shape=1, dtype="float32", data=a) T.evaluate(A[0]) before = tvm.IRModule( @@ -185,12 +185,12 @@ def func(a: T.handle("float32")): class expected: @T.prim_func def func_a(a: T.handle("float32")): - A = T.decl_buffer(shape=1, dtype="float32", data=a) + A = T.decl_tensor(shape=1, dtype="float32", data=a) T.evaluate(A[0]) @T.prim_func def func_b(a: T.handle("float32")): - A = T.decl_buffer(shape=1, dtype="float32", data=a) + A = T.decl_tensor(shape=1, dtype="float32", data=a) T.evaluate(A[0]) after = tvm.tirx.transform.ConvertSSA()(before) @@ -201,7 +201,7 @@ def test_reused_buffer_parameter(): """De-duplicate buffer parameters across the entire module.""" @T.prim_func(private=True) - def func(A: T.Buffer(1, "float32")): + def func(A: T.Tensor(1, "float32")): T.evaluate(A[0]) before = tvm.IRModule( @@ -214,11 +214,11 @@ def func(A: T.Buffer(1, "float32")): @I.ir_module class expected: @T.prim_func - def func_a(A: T.Buffer(1, "float32")): + def func_a(A: T.Tensor(1, "float32")): T.evaluate(A[0]) @T.prim_func - def func_b(A: T.Buffer(1, "float32")): + def func_b(A: T.Tensor(1, "float32")): T.evaluate(A[0]) after = tvm.tirx.transform.ConvertSSA()(before) @@ -228,7 +228,7 @@ def func_b(A: T.Buffer(1, "float32")): def test_reused_compound_buffer_shape_var(): """De-duplicate implicit Vars nested in buffer parameter shapes.""" n = tirx.Var("n", "int32") - A = tirx.decl_buffer((tirx.max(n, 1),), layout=None) + A = tirx.decl_tensor((tirx.max(n, 1),), layout=None) func = tirx.PrimFunc([A], tirx.Evaluate(n)) before = tvm.IRModule( { @@ -253,7 +253,7 @@ def test_no_change_if_already_ssa(): @I.ir_module class before: @T.prim_func - def func(A: T.Buffer(1, "float32")): + def func(A: T.Tensor(1, "float32")): T.evaluate(A[0]) after = tvm.tirx.transform.ConvertSSA()(before) @@ -284,7 +284,7 @@ def test_keep_duplicate_thread_idx_in_same_function(): @I.ir_module class before: @T.prim_func - def main(A: T.Buffer([256], "float32")): + def main(A: T.Tensor([256], "float32")): threadIdx_x = T.env_thread("threadIdx.x") with T.launch_thread(threadIdx_x, 256): A[threadIdx_x] = A[threadIdx_x] + 1.0 @@ -320,7 +320,7 @@ def test_de_duplicate_thread_idx_across_multiple_functions(): @I.ir_module(check_well_formed=False) class before: @T.prim_func - def kernel_1(A: T.Buffer([256], "float32")): + def kernel_1(A: T.Tensor([256], "float32")): T.attr( T.iter_var(threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", @@ -329,7 +329,7 @@ def kernel_1(A: T.Buffer([256], "float32")): A[threadIdx_x] = A[threadIdx_x] + T.float32(1) @T.prim_func - def kernel_2(A: T.Buffer([256], "float32")): + def kernel_2(A: T.Tensor([256], "float32")): T.attr( T.iter_var(threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", @@ -343,7 +343,7 @@ def kernel_2(A: T.Buffer([256], "float32")): @I.ir_module class expected: @T.prim_func - def kernel_1(A: T.Buffer([256], "float32")): + def kernel_1(A: T.Tensor([256], "float32")): T.attr( T.iter_var(kernel_1_threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", @@ -352,7 +352,7 @@ def kernel_1(A: T.Buffer([256], "float32")): A[kernel_1_threadIdx_x] = A[kernel_1_threadIdx_x] + T.float32(1) @T.prim_func - def kernel_2(A: T.Buffer([256], "float32")): + def kernel_2(A: T.Tensor([256], "float32")): T.attr( T.iter_var(kernel_2_threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", @@ -381,12 +381,12 @@ def test_de_duplicate_thread_idx_iter_var_across_multiple_functions(): @I.ir_module(check_well_formed=False) class before: @T.prim_func - def kernel_1(A: T.Buffer([256], "float32")): + def kernel_1(A: T.Tensor([256], "float32")): T.attr(iter_var, "thread_extent", 256) A[threadIdx_x] = A[threadIdx_x] + T.float32(1) @T.prim_func - def kernel_2(A: T.Buffer([256], "float32")): + def kernel_2(A: T.Tensor([256], "float32")): T.attr(iter_var, "thread_extent", 256) A[threadIdx_x] = A[threadIdx_x] + T.float32(1) @@ -396,7 +396,7 @@ def kernel_2(A: T.Buffer([256], "float32")): @I.ir_module(check_well_formed=False) class expected: @T.prim_func - def kernel_1(A: T.Buffer([256], "float32")): + def kernel_1(A: T.Tensor([256], "float32")): T.attr( T.iter_var(kernel_1_threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", @@ -405,7 +405,7 @@ def kernel_1(A: T.Buffer([256], "float32")): A[kernel_1_threadIdx_x] = A[kernel_1_threadIdx_x] + T.float32(1) @T.prim_func - def kernel_2(A: T.Buffer([256], "float32")): + def kernel_2(A: T.Tensor([256], "float32")): T.attr( T.iter_var(kernel_2_threadIdx_x, T.Range(0, 256), "ThreadIndex", "threadIdx.x"), "thread_extent", @@ -436,14 +436,14 @@ def test_thread_idx_reused_within_and_across_functions(): @I.ir_module(check_well_formed=False) class before: @T.prim_func - def kernel_1(A: T.Buffer([256], "float32")): + def kernel_1(A: T.Tensor([256], "float32")): with T.attr(iter_var, "thread_extent", 256): A[threadIdx_x] = A[threadIdx_x] + 1.0 with T.attr(iter_var, "thread_extent", 256): A[threadIdx_x] = A[threadIdx_x] + 2.0 @T.prim_func - def kernel_2(A: T.Buffer([256], "float32")): + def kernel_2(A: T.Tensor([256], "float32")): with T.attr(iter_var, "thread_extent", 256): A[threadIdx_x] = A[threadIdx_x] + 1.0 with T.attr(iter_var, "thread_extent", 256): @@ -452,7 +452,7 @@ def kernel_2(A: T.Buffer([256], "float32")): @I.ir_module class expected: @T.prim_func - def kernel_1(A: T.Buffer([256], "float32")): + def kernel_1(A: T.Tensor([256], "float32")): threadIdx_x = T.env_thread("threadIdx.x") with T.launch_thread(threadIdx_x, 256): A[threadIdx_x] = A[threadIdx_x] + 1.0 @@ -460,7 +460,7 @@ def kernel_1(A: T.Buffer([256], "float32")): A[threadIdx_x] = A[threadIdx_x] + 2.0 @T.prim_func - def kernel_2(A: T.Buffer([256], "float32")): + def kernel_2(A: T.Tensor([256], "float32")): threadIdx_x = T.env_thread("threadIdx.x") with T.launch_thread(threadIdx_x, 256): A[threadIdx_x] = A[threadIdx_x] + 1.0 @@ -485,8 +485,8 @@ def test_track_forward_declarations_in_attr_stmt(): i0_outer_inner = tirx.Var("i0_outer_inner", "int32") i0_inner = tirx.Var("i0_inner", "int32") - A = tirx.decl_buffer(1024, "float32", "A") - B = tirx.decl_buffer(1024, "float32", "B") + A = tirx.decl_tensor(1024, "float32", "A") + B = tirx.decl_tensor(1024, "float32", "B") index = i0_outer_outer * 52 + i0_outer_inner * 4 + i0_inner @@ -522,26 +522,26 @@ def test_track_forward_declarations_in_attr_stmt(): def test_shared_shape_var_in_buffer_params_and_alloc_buffer(): - """Shape var shared across buffer params and AllocBuffer should not be renamed. + """Shape var shared across buffer params and AllocTensor should not be renamed. When the same Var (e.g., `n`) appears in multiple buffer parameter annotations (A and B both have shape [n]), ConvertSSA should not treat the second occurrence as a redefinition. All uses of `n` in the - function body (including AllocBuffer shapes) must remain the same + function body (including AllocTensor shapes) must remain the same Var object so that MakePackedAPI can bind it from the DLTensor shape. """ n = tirx.Var("n", "int32") - A = tirx.decl_buffer((n,), "float32", "A") - B = tirx.decl_buffer((n,), "float32", "B") + A = tirx.decl_tensor((n,), "float32", "A") + B = tirx.decl_tensor((n,), "float32", "B") - # AllocBuffer with shape [n] in the body (flat, no body) - C = tirx.decl_buffer((n,), "float32", "C") + # AllocTensor with shape [n] in the body (flat, no body) + C = tirx.decl_tensor((n,), "float32", "C") body = tirx.SeqStmt( [ tvm.tirx.Bind( C, tvm.ir.Call( - "tirx.alloc_buffer", + "tirx.alloc_tensor", [ tvm.ir.Tuple(C.shape), tvm.ir.DataTypeImm(tvm.DataType(C.dtype)), @@ -566,7 +566,7 @@ def test_shared_shape_var_in_buffer_params_and_alloc_buffer(): def test_reused_loop_var_in_decl_buffer_elem_offset(): """Remap a buffer whose elem_offset depends on an SSA-renamed loop var.""" loop_var = tirx.Var("loop_var", "int32") - buffer = tirx.decl_buffer( + buffer = tirx.decl_tensor( (128,), "float32", "buffer", @@ -584,7 +584,7 @@ def test_reused_loop_var_in_decl_buffer_elem_offset(): tirx.Bind( buffer, tvm.ir.Call( - "tirx.decl_buffer", + "tirx.decl_tensor", [ buffer_data, tvm.ir.Tuple(buffer.shape), diff --git a/tests/python/tirx-transform/test_tir_transform_flatten_buffer.py b/tests/python/tirx-transform/test_tir_transform_flatten_buffer.py index 6338edec7759..7d19a99b74da 100644 --- a/tests/python/tirx-transform/test_tir_transform_flatten_buffer.py +++ b/tests/python/tirx-transform/test_tir_transform_flatten_buffer.py @@ -37,9 +37,9 @@ def test_elementwise(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): + def main(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")): for i in T.serial(0, 16): - B_new = T.decl_buffer([1, 16], "float32") + B_new = T.decl_tensor([1, 16], "float32") for j in T.serial(0, 16): B_new[0, j] = A[i, j] + 1.0 for j in T.serial(0, 16): @@ -48,11 +48,11 @@ def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): - A_1 = T.decl_buffer(256, dtype="float32", data=A.data, layout=None) - C_1 = T.decl_buffer(256, dtype="float32", data=C.data, layout=None) + def main(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")): + A_1 = T.decl_tensor(256, dtype="float32", data=A.data, layout=None) + C_1 = T.decl_tensor(256, dtype="float32", data=C.data, layout=None) for i in T.serial(0, 16): - B_new = T.decl_buffer([16], "float32", layout=None) + B_new = T.decl_tensor([16], "float32", layout=None) for j in T.serial(0, 16): B_new[j] = A_1[((i * 16) + j)] + 1.0 for j in T.serial(0, 16): @@ -65,8 +65,8 @@ def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): def test_elementwise_without_decl_buffer(): """2-d buffers are flattened to 1-d - Like test_elementwise, but the TIR doesn't have the DeclBuffer - node. The T.Buffer declaration applies only during the + Like test_elementwise, but the TIR doesn't have the DeclTensor + node. The T.Tensor declaration applies only during the parsing the TVMScript, and doesn't occur in the TIR itself. In this case, the allocation should be assumed to be targeting flat memory, and should be flattened to a 1-d allocation. @@ -75,10 +75,10 @@ def test_elementwise_without_decl_buffer(): @I.ir_module(check_well_formed=False) class Before: @T.prim_func - def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): + def main(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")): for i in T.serial(0, 16): - B_new_buf = T.alloc_buffer((1, 16), "float32") - B_new = T.decl_buffer([1, 16], "float32", data=B_new_buf.data) + B_new_buf = T.alloc_tensor((1, 16), "float32") + B_new = T.decl_tensor([1, 16], "float32", data=B_new_buf.data) for j in T.serial(0, 16): B_new[0, j] = A[i, j] + 1.0 for j in T.serial(0, 16): @@ -87,12 +87,12 @@ def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): @I.ir_module(check_well_formed=False) class Expected: @T.prim_func - def main(input_A: T.Buffer((16, 16), "float32"), input_C: T.Buffer((16, 16), "float32")): - A = T.decl_buffer(256, dtype="float32", data=input_A.data, layout=None) - C = T.decl_buffer(256, dtype="float32", data=input_C.data, layout=None) + def main(input_A: T.Tensor((16, 16), "float32"), input_C: T.Tensor((16, 16), "float32")): + A = T.decl_tensor(256, dtype="float32", data=input_A.data, layout=None) + C = T.decl_tensor(256, dtype="float32", data=input_C.data, layout=None) for i in T.serial(0, 16): - B_new_buf = T.alloc_buffer((16,), "float32", layout=None) - B_new = T.decl_buffer(16, "float32", data=B_new_buf.data, layout=None) + B_new_buf = T.alloc_tensor((16,), "float32", layout=None) + B_new = T.decl_tensor(16, "float32", data=B_new_buf.data, layout=None) for j in T.serial(0, 16): B_new[j] = A[((i * 16) + j)] + 1.0 for j in T.serial(0, 16): @@ -108,7 +108,7 @@ def test_gpu(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): + def main(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")): i0 = T.env_thread("blockIdx.x") i1 = T.env_thread("threadIdx.x") i2 = T.env_thread("vthread") @@ -116,7 +116,7 @@ def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): T.launch_thread(i0, 4) T.launch_thread(i1, 2) T.launch_thread(i2, 2) - B = T.decl_buffer([1, 16], "float32", scope="local") + B = T.decl_tensor([1, 16], "float32", scope="local") for j in range(0, 16): B[0, j] = A[i0 * 4 + i1 * 2 + i2, j] + 1.0 for j in range(0, 16): @@ -125,9 +125,9 @@ def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): - A_1 = T.decl_buffer(256, dtype="float32", data=A.data, layout=None) - C_1 = T.decl_buffer(256, dtype="float32", data=C.data, layout=None) + def main(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")): + A_1 = T.decl_tensor(256, dtype="float32", data=A.data, layout=None) + C_1 = T.decl_tensor(256, dtype="float32", data=C.data, layout=None) i0 = T.env_thread("blockIdx.x") i1 = T.env_thread("threadIdx.x") @@ -136,7 +136,7 @@ def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): T.launch_thread(i0, 4) T.launch_thread(i1, 2) T.launch_thread(i2, 2) - B = T.decl_buffer([16], "float32", scope="local", layout=None) + B = T.decl_tensor([16], "float32", scope="local", layout=None) for j in range(0, 16): B[j] = A_1[i0 * 64 + i1 * 32 + i2 * 16 + j] + 1.0 for j in range(0, 16): @@ -153,13 +153,13 @@ def test_symbolic(): class Before: @T.prim_func def main( - A: T.Buffer((n, m), "float32"), # noqa: F821 - C: T.Buffer((n, m), "float32"), # noqa: F821 + A: T.Tensor((n, m), "float32"), # noqa: F821 + C: T.Tensor((n, m), "float32"), # noqa: F821 n: T.int32, m: T.int32, ) -> None: for i in range(0, n): - B = T.decl_buffer([m], "float32") + B = T.decl_tensor([m], "float32") for j in range(0, m): B[j] = A[i, j] + 1.0 for j in range(0, m): @@ -169,16 +169,16 @@ def main( class Expected: @T.prim_func def main( - A: T.Buffer((n, m), "float32"), # noqa: F821 - C: T.Buffer((n, m), "float32"), # noqa: F821 + A: T.Tensor((n, m), "float32"), # noqa: F821 + C: T.Tensor((n, m), "float32"), # noqa: F821 n: T.int32, m: T.int32, ) -> None: - A_1 = T.decl_buffer(n * m, "float32", data=A.data, layout=None) - C_1 = T.decl_buffer(n * m, "float32", data=C.data, layout=None) + A_1 = T.decl_tensor(n * m, "float32", data=A.data, layout=None) + C_1 = T.decl_tensor(n * m, "float32", data=C.data, layout=None) for i in range(0, n): - B = T.decl_buffer([m], "float32", layout=None) + B = T.decl_tensor([m], "float32", layout=None) for j in range(0, m): B[j] = A_1[i * m + j] + 1.0 for j in range(0, m): @@ -195,8 +195,8 @@ def test_fused_symbolic(): class Before: @T.prim_func def main( - A: T.Buffer((32, n, n), "float32"), # noqa: F821 - B: T.Buffer((32, n, n), "float32"), # noqa: F821 + A: T.Tensor((32, n, n), "float32"), # noqa: F821 + B: T.Tensor((32, n, n), "float32"), # noqa: F821 n: T.int32, ) -> None: for i in range(0, n * n * 32): @@ -208,12 +208,12 @@ def main( class Expected: @T.prim_func def main( - input_A: T.Buffer((32, n, n), "float32"), # noqa: F821 - input_B: T.Buffer((32, n, n), "float32"), # noqa: F821 + input_A: T.Tensor((32, n, n), "float32"), # noqa: F821 + input_B: T.Tensor((32, n, n), "float32"), # noqa: F821 n: T.int32, ) -> None: - A = T.decl_buffer(n * n * 32, "float32", data=input_A.data, layout=None) - B = T.decl_buffer(n * n * 32, "float32", data=input_B.data, layout=None) + A = T.decl_tensor(n * n * 32, "float32", data=input_A.data, layout=None) + B = T.decl_tensor(n * n * 32, "float32", data=input_B.data, layout=None) for i in range(0, n * n * 32): B[i] = A[i] @@ -229,8 +229,8 @@ def test_fused_symbolic_with_predicate(): class Before: @T.prim_func def main( - A: T.Buffer((32, n, n), "float32"), # noqa: F821 - B: T.Buffer((32, n, n), "float32"), # noqa: F821 + A: T.Tensor((32, n, n), "float32"), # noqa: F821 + B: T.Tensor((32, n, n), "float32"), # noqa: F821 n: T.int32, ) -> None: for bx, tx in T.grid((n * n + 1) // 2, 64): @@ -249,12 +249,12 @@ def main( class Expected: @T.prim_func def main( - input_A: T.Buffer((32, n, n), "float32"), # noqa: F821 - input_B: T.Buffer((32, n, n), "float32"), # noqa: F821 + input_A: T.Tensor((32, n, n), "float32"), # noqa: F821 + input_B: T.Tensor((32, n, n), "float32"), # noqa: F821 n: T.int32, ) -> None: - A = T.decl_buffer(n * n * 32, "float32", data=input_A.data, layout=None) - B = T.decl_buffer(n * n * 32, "float32", data=input_B.data, layout=None) + A = T.decl_tensor(n * n * 32, "float32", data=input_A.data, layout=None) + B = T.decl_tensor(n * n * 32, "float32", data=input_B.data, layout=None) for bx, tx in T.grid((n * n + 1) // 2, 64): if bx * 64 + tx < n * n * 32: @@ -270,10 +270,10 @@ def test_multi_alloc(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((4, 32), "float32"), D: T.Buffer((4, 32), "float32")): + def main(A: T.Tensor((4, 32), "float32"), D: T.Tensor((4, 32), "float32")): for i, j in T.grid(4, 32): - B = T.decl_buffer((4, 32), "float32", scope="global") - C = T.decl_buffer((4, 32), "float32", scope="global") + B = T.decl_tensor((4, 32), "float32", scope="global") + C = T.decl_tensor((4, 32), "float32", scope="global") B[i, j] = A[i, j] + 1.0 C[i, j] = A[i, j] + B[i, j] D[i, j] = C[i, j] * 2.0 @@ -281,13 +281,13 @@ def main(A: T.Buffer((4, 32), "float32"), D: T.Buffer((4, 32), "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((4, 32), "float32"), D: T.Buffer((4, 32), "float32")): - A_1 = T.decl_buffer(128, "float32", data=A.data, layout=None) - D_1 = T.decl_buffer(128, "float32", data=D.data, layout=None) + def main(A: T.Tensor((4, 32), "float32"), D: T.Tensor((4, 32), "float32")): + A_1 = T.decl_tensor(128, "float32", data=A.data, layout=None) + D_1 = T.decl_tensor(128, "float32", data=D.data, layout=None) for i, j in T.grid(4, 32): - B = T.decl_buffer([128], "float32", layout=None) - C = T.decl_buffer([128], "float32", layout=None) + B = T.decl_tensor([128], "float32", layout=None) + C = T.decl_tensor([128], "float32", layout=None) B[i * 32 + j] = A_1[i * 32 + j] + 1.0 C[i * 32 + j] = A_1[i * 32 + j] + B[i * 32 + j] D_1[i * 32 + j] = C[i * 32 + j] * 2.0 @@ -302,10 +302,10 @@ def test_strided(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): + def main(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")): for i0 in T.serial(4): - B = T.decl_buffer([4, 17], "float32") - B_1 = T.decl_buffer([4, 16], dtype="float32", data=B.data, strides=[17, 1]) + B = T.decl_tensor([4, 17], "float32") + B_1 = T.decl_tensor([4, 16], dtype="float32", data=B.data, strides=[17, 1]) for i1, j in T.grid(4, 16): B_1[i1, j] = A[i0 * 4 + i1, j] + 1.0 for i1, j in T.grid(4, 16): @@ -314,12 +314,12 @@ def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((16, 16), "float32"), C: T.Buffer((16, 16), "float32")): - A_1 = T.decl_buffer(256, dtype="float32", data=A.data, layout=None) - C_1 = T.decl_buffer(256, dtype="float32", data=C.data, layout=None) + def main(A: T.Tensor((16, 16), "float32"), C: T.Tensor((16, 16), "float32")): + A_1 = T.decl_tensor(256, dtype="float32", data=A.data, layout=None) + C_1 = T.decl_tensor(256, dtype="float32", data=C.data, layout=None) for i0 in T.serial(0, 4): - B = T.decl_buffer([68], "float32", layout=None) - B_1 = T.decl_buffer([68], "float32", data=B.data, layout=None) + B = T.decl_tensor([68], "float32", layout=None) + B_1 = T.decl_tensor([68], "float32", data=B.data, layout=None) for i1 in T.serial(0, 4): for j in T.serial(0, 16): B_1[i1 * 17 + j] = A_1[i0 * 64 + i1 * 16 + j] + 1.0 @@ -337,16 +337,16 @@ def test_boolean(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(10, "bool"), B: T.Buffer(10, "bool")) -> None: + def main(A: T.Tensor(10, "bool"), B: T.Tensor(10, "bool")) -> None: for i0 in T.serial(10): B[i0] = A[i0] @I.ir_module class Expected: @T.prim_func - def main(input_A: T.Buffer(10, "bool"), input_B: T.Buffer(10, "bool")) -> None: - A = T.decl_buffer(10, dtype="bool", data=input_A.data, layout=None) - B = T.decl_buffer(10, dtype="bool", data=input_B.data, layout=None) + def main(input_A: T.Tensor(10, "bool"), input_B: T.Tensor(10, "bool")) -> None: + A = T.decl_tensor(10, dtype="bool", data=input_A.data, layout=None) + B = T.decl_tensor(10, dtype="bool", data=input_B.data, layout=None) # body for i0 in T.serial(10): B[i0] = A[i0] @@ -362,7 +362,7 @@ def test_flatten_inside_block(): class Before: @T.prim_func def main(): - A = T.alloc_buffer([32, 32]) + A = T.alloc_tensor([32, 32]) for i, j in T.grid(32, 32): T.evaluate(A[i, j]) @@ -370,7 +370,7 @@ def main(): class Expected: @T.prim_func def main(): - A = T.alloc_buffer([1024], layout=None) + A = T.alloc_tensor([1024], layout=None) for i, j in T.grid(32, 32): T.evaluate(A[i * 32 + j]) @@ -383,7 +383,7 @@ def check(value): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((4, 5, 6), "int16"), B: T.Buffer((4, 5, 6), "int16")): + def main(A: T.Tensor((4, 5, 6), "int16"), B: T.Tensor((4, 5, 6), "int16")): for ax0 in T.serial(4, annotations={"pragma_unroll_explicit": value}): for ax1, ax2 in T.grid(5, 6): B[ax0, ax1, ax2] = A[ax0, ax1, ax2] 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 acecda8cb675..6905b15c607b 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 @@ -23,7 +23,7 @@ def test_thread_axis1(): @T.prim_func(private=True) - def before(A: T.Buffer((T.int64(64),), "float32"), B: T.Buffer((T.int64(64),), "float32")): + def before(A: T.Tensor((T.int64(64),), "float32"), B: T.Tensor((T.int64(64),), "float32")): blockIdx_x = T.env_thread("blockIdx.x") T.launch_thread(blockIdx_x, T.int64(2)) threadIdx_x = T.env_thread("threadIdx.x") @@ -33,7 +33,7 @@ def before(A: T.Buffer((T.int64(64),), "float32"), B: T.Buffer((T.int64(64),), " ] + T.float32(1) @T.prim_func(private=True) - def expected(A: T.Buffer((64,), "float32"), B: T.Buffer((64,), "float32")): + def expected(A: T.Tensor((64,), "float32"), B: T.Tensor((64,), "float32")): blockIdx_x = T.env_thread("blockIdx.x") T.launch_thread(blockIdx_x, 2) threadIdx_x = T.env_thread("threadIdx.x") @@ -48,9 +48,9 @@ def expected(A: T.Buffer((64,), "float32"), B: T.Buffer((64,), "float32")): def test_thread_axis2(): @T.prim_func def before( - T_reshape: T.Buffer((1, 12, 384, 384), "float32"), - placeholder_1: T.Buffer((T.int64(1), T.int64(12), T.int64(384), 384), "bool"), - T_where: T.Buffer((T.int64(1), T.int64(12), T.int64(384), 384), "float32"), + T_reshape: T.Tensor((1, 12, 384, 384), "float32"), + placeholder_1: T.Tensor((T.int64(1), T.int64(12), T.int64(384), 384), "bool"), + T_where: T.Tensor((T.int64(1), T.int64(12), T.int64(384), 384), "float32"), ) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0_i1_i2_i3_fused_1 in T.thread_binding(T.int64(256), thread="blockIdx.x"): @@ -149,9 +149,9 @@ def before( @T.prim_func def expected( - T_reshape: T.Buffer((1, 12, 384, 384), "float32"), - placeholder_1: T.Buffer((1, 12, 384, 384), "bool"), - T_where: T.Buffer((1, 12, 384, 384), "float32"), + T_reshape: T.Tensor((1, 12, 384, 384), "float32"), + placeholder_1: T.Tensor((1, 12, 384, 384), "bool"), + T_where: T.Tensor((1, 12, 384, 384), "float32"), ): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i0_i1_i2_i3_fused_1 in T.thread_binding(256, thread="blockIdx.x"): @@ -234,13 +234,13 @@ def expected( def test_block(): @T.prim_func(private=True) - def before(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def before(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): for i in T.serial(0, T.int64(16)): for j in T.serial(0, T.int64(8)): B[i * T.int64(8) + j] = A[i * T.int64(8) + j] + T.float32(1) @T.prim_func(private=True) - def expected(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def expected(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): for i in T.serial(0, T.int32(16)): for j in T.serial(0, T.int32(8)): B[i * T.int32(8) + j] = A[i * T.int32(8) + j] + T.float32(1) @@ -252,13 +252,13 @@ def expected(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): def test_i16_buffer(): @T.prim_func(private=True) - def before(A: T.Buffer((128,), "int16"), B: T.Buffer((128,), "int16")): + def before(A: T.Tensor((128,), "int16"), B: T.Tensor((128,), "int16")): for i in T.serial(0, T.int64(16)): for j in T.serial(0, T.int64(16)): B[i * 8 + j] = A[i * 8 + j] + T.int16(1) @T.prim_func(private=True) - def expected(A: T.Buffer((128,), "int16"), B: T.Buffer((128,), "int16")): + def expected(A: T.Tensor((128,), "int16"), B: T.Tensor((128,), "int16")): for i in T.serial(0, 16): for j in T.serial(0, 16): B[i * 8 + j] = A[i * 8 + j] + T.int16(1) @@ -270,7 +270,7 @@ def expected(A: T.Buffer((128,), "int16"), B: T.Buffer((128,), "int16")): def test_fail_on_buffer_param(): @T.prim_func(private=True) - def func(A: T.Buffer((128,), "int64"), B: T.Buffer((128,), "int64")): + def func(A: T.Tensor((128,), "int64"), B: T.Tensor((128,), "int64")): for i in T.serial(0, 16): for j in T.serial(0, 8): B[i * 8 + j] = A[i * 8 + j] + T.int64(1) @@ -282,8 +282,8 @@ def func(A: T.Buffer((128,), "int64"), B: T.Buffer((128,), "int64")): def test_fail_on_internal_buffer(): @T.prim_func(private=True) - def func(A: T.Buffer((128,), "int32"), B: T.Buffer((128,), "int32")): - C = T.alloc_buffer((128,), "int64") + def func(A: T.Tensor((128,), "int32"), B: T.Tensor((128,), "int32")): + C = T.alloc_tensor((128,), "int64") for i in T.serial(0, 16): for j in T.serial(0, 8): C[i * 8 + j] = T.cast(A[i * 8 + j], "int64") + T.int64(1) @@ -301,7 +301,7 @@ def test_pod_params_and_select(): class Before: @T.prim_func def main( - A: T.Buffer((T.int64(4),), "float32"), B: T.Buffer((T.int64(4),), "float32"), n: T.int64 + A: T.Tensor((T.int64(4),), "float32"), B: T.Tensor((T.int64(4),), "float32"), n: T.int64 ): for i in T.serial(T.int64(4)): B[i] = T.Select(T.int64(1) <= i, A[i + n], T.Cast("float32", i)) @@ -309,7 +309,7 @@ def main( @tvm.script.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((4,), "float32"), B: T.Buffer((4,), "float32"), n: T.int32): + def main(A: T.Tensor((4,), "float32"), B: T.Tensor((4,), "float32"), n: T.int32): for i in range(4): B[i] = T.Select(1 <= i, A[i + n], T.Cast("float32", i)) @@ -321,13 +321,13 @@ def test_if_then_else_index(): @tvm.script.ir_module class Before: @T.prim_func - def main(A: T.Buffer((T.int64(4),), "float32"), B: T.Buffer((1,), "float32"), n: T.int64): + def main(A: T.Tensor((T.int64(4),), "float32"), B: T.Tensor((1,), "float32"), n: T.int64): B[0] = A[T.if_then_else(n < T.int64(0), n + T.int64(1), n)] @tvm.script.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((4,), "float32"), B: T.Buffer((1,), "float32"), n: T.int32): + def main(A: T.Tensor((4,), "float32"), B: T.Tensor((1,), "float32"), n: T.int32): B[0] = A[T.if_then_else(n < 0, n + 1, n)] after = tvm.tirx.transform.ForceNarrowIndexToInt32()(Before) @@ -338,7 +338,7 @@ def test_conditional_index_mixed_width_branches(): @tvm.script.ir_module class Before: @T.prim_func - def main(A: T.Buffer((T.int64(4),), "float32"), B: T.Buffer((4,), "float32"), n: T.int64): + def main(A: T.Tensor((T.int64(4),), "float32"), B: T.Tensor((4,), "float32"), n: T.int64): opaque_index: T.int64 = T.call_extern("opaque_index", n, dtype="int64") B[0] = A[T.if_then_else(n < T.int64(0), opaque_index, n)] B[1] = A[T.if_then_else(n < T.int64(0), n, opaque_index)] @@ -348,7 +348,7 @@ def main(A: T.Buffer((T.int64(4),), "float32"), B: T.Buffer((4,), "float32"), n: @tvm.script.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((4,), "float32"), B: T.Buffer((4,), "float32"), n: T.int32): + def main(A: T.Tensor((4,), "float32"), B: T.Tensor((4,), "float32"), n: T.int32): opaque_index: T.int64 = T.call_extern("opaque_index", n, dtype="int64") B[0] = A[T.if_then_else(n < 0, opaque_index, T.Cast("int64", n))] B[1] = A[T.if_then_else(n < 0, T.Cast("int64", n), opaque_index)] @@ -363,14 +363,14 @@ def test_clz(): @tvm.script.ir_module class Before: @T.prim_func - def main(B: T.Buffer((T.int64(4),), "int32")): + def main(B: T.Tensor((T.int64(4),), "int32")): for i in T.serial(T.int64(4)): B[i] = T.clz(i) @tvm.script.ir_module class Expected: @T.prim_func - def main(B: T.Buffer((4,), "int32")): + def main(B: T.Tensor((4,), "int32")): for i in range(4): B[i] = T.clz(i) - 32 + 64 @@ -382,13 +382,13 @@ def test_right_shift_preserves_sign_extension_after_narrowing(): @tvm.script.ir_module class Before: @T.prim_func - def main(A: T.Buffer((T.int64(6),), "float32"), B: T.Buffer((1,), "float32"), n: T.int64): + def main(A: T.Tensor((T.int64(6),), "float32"), B: T.Tensor((1,), "float32"), n: T.int64): B[0] = A[T.shift_right(T.truncmod(n - T.int64(8), T.int64(6)), T.int64(63))] @tvm.script.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((6,), "float32"), B: T.Buffer((1,), "float32"), n: T.int32): + def main(A: T.Tensor((6,), "float32"), B: T.Tensor((1,), "float32"), n: T.int32): B[0] = A[T.shift_right(T.truncmod(n - 8, 6), 31)] # ForceNarrowIndexToInt32 assumes that index values fit in int32. Under @@ -403,8 +403,8 @@ def test_right_shift_dynamic_and_vector_amounts(): class Before: @T.prim_func def main( - A: T.Buffer((T.int64(6),), "float32"), - B: T.Buffer((T.int64(5),), "float32"), + A: T.Tensor((T.int64(6),), "float32"), + B: T.Tensor((T.int64(5),), "float32"), n: T.int64, shift: T.int64, ): @@ -420,8 +420,8 @@ def main( class Expected: @T.prim_func def main( - A: T.Buffer((6,), "float32"), - B: T.Buffer((5,), "float32"), + A: T.Tensor((6,), "float32"), + B: T.Tensor((5,), "float32"), n: T.int32, shift: T.int32, ): @@ -442,8 +442,8 @@ def test_left_shift_dynamic_and_vector_amounts_remain_valid(): class Before: @T.prim_func def main( - A: T.Buffer((T.int64(1),), "float32"), - B: T.Buffer((T.int64(5),), "float32"), + A: T.Tensor((T.int64(1),), "float32"), + B: T.Tensor((T.int64(5),), "float32"), n: T.int64, shift: T.int64, ): @@ -459,8 +459,8 @@ def main( class Expected: @T.prim_func def main( - A: T.Buffer((1,), "float32"), - B: T.Buffer((5,), "float32"), + A: T.Tensor((1,), "float32"), + B: T.Tensor((5,), "float32"), n: T.int32, shift: T.int32, ): @@ -482,7 +482,7 @@ def test_let_binding(): @tvm.script.ir_module class Before: @T.prim_func - def main(Buf: T.Buffer([n], "int32")): + def main(Buf: T.Tensor([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) @@ -492,7 +492,7 @@ def main(Buf: T.Buffer([n], "int32")): @tvm.script.ir_module class Expected: @T.prim_func - def main(Buf: T.Buffer([n], "int32")): + def main(Buf: T.Tensor([n], "int32")): # The pass narrows indexing variables (n, the For extent) but leaves # an explicitly-typed `T.Cast("int64", ...)` storage alone; a Cast to # int32 is inserted at the use site (the For iter) instead. diff --git a/tests/python/tirx-transform/test_tir_transform_fp8_legalize.py b/tests/python/tirx-transform/test_tir_transform_fp8_legalize.py index 414bde94f0ea..c58295086cc4 100644 --- a/tests/python/tirx-transform/test_tir_transform_fp8_legalize.py +++ b/tests/python/tirx-transform/test_tir_transform_fp8_legalize.py @@ -30,10 +30,10 @@ class Before: @T.prim_func def main(Aptr: T.handle(dtype), Bptr: T.handle(dtype), Dptr: T.handle(dtype)): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), dtype, data=Aptr) - B = T.decl_buffer((100,), dtype, data=Bptr) - D = T.decl_buffer((100,), dtype, data=Dptr) - C = T.decl_buffer((100,), dtype) + A = T.decl_tensor((100,), dtype, data=Aptr) + B = T.decl_tensor((100,), dtype, data=Bptr) + D = T.decl_tensor((100,), dtype, data=Dptr) + C = T.decl_tensor((100,), dtype) for i in T.grid(100): C[i] = A[i] + B[i] D[i] = T.exp(C[i]) @@ -55,10 +55,10 @@ class After: @T.prim_func def main(Aptr: T.handle(dtype), Bptr: T.handle(dtype), Dptr: T.handle(dtype)): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), dtype, data=Aptr) - B = T.decl_buffer((100,), dtype, data=Bptr) - D = T.decl_buffer((100,), dtype, data=Dptr) - C = T.decl_buffer((100,), promote_dtype) + A = T.decl_tensor((100,), dtype, data=Aptr) + B = T.decl_tensor((100,), dtype, data=Bptr) + D = T.decl_tensor((100,), dtype, data=Dptr) + C = T.decl_tensor((100,), promote_dtype) for i in T.grid(100): C[i] = promote_f8(dtype, promote_dtype, A[i]) + promote_f8( dtype, promote_dtype, B[i] @@ -188,10 +188,10 @@ class After: @T.prim_func def main(Aptr: T.handle("uint8"), Bptr: T.handle("uint8"), Dptr: T.handle("uint8")): T.func_attr({"global_symbol": "main"}) - A = T.decl_buffer((100,), "uint8", data=Aptr) - B = T.decl_buffer((100,), "uint8", data=Bptr) - D = T.decl_buffer((100,), "uint8", data=Dptr) - C = T.decl_buffer((100,), promote_dtype) + A = T.decl_tensor((100,), "uint8", data=Aptr) + B = T.decl_tensor((100,), "uint8", data=Bptr) + D = T.decl_tensor((100,), "uint8", data=Dptr) + C = T.decl_tensor((100,), promote_dtype) for i in T.grid(100): C[i] = promote_uint8(dtype, promote_dtype, A[i]) + promote_uint8( dtype, promote_dtype, B[i] @@ -219,7 +219,7 @@ def test_fp8_compute_legalize(dtype, promote_dtype): def test_fp8_compute_legalize_preserves_opaque_buffer_access(dtype, promote_dtype): @T.prim_func def before(): - buffer = T.alloc_buffer((16,), dtype) + buffer = T.alloc_tensor((16,), dtype) T.evaluate(T.call_extern("void", "consume", buffer.data)) before_mod = tvm.IRModule.from_expr(before) diff --git a/tests/python/tirx-transform/test_tir_transform_helpers.py b/tests/python/tirx-transform/test_tir_transform_helpers.py index f4c30904d5c4..e7cb4acc7a2a 100644 --- a/tests/python/tirx-transform/test_tir_transform_helpers.py +++ b/tests/python/tirx-transform/test_tir_transform_helpers.py @@ -27,7 +27,7 @@ def test_annotate_entry_func_single_primfunc(): @tvm.script.ir_module class MockModule: @T.prim_func(private=True) - def func1(A: T.Buffer((16,), "float32")): + def func1(A: T.Tensor((16,), "float32")): for i in T.serial(16): if i == 5: if i == 5: @@ -48,14 +48,14 @@ def func1(A: T.Buffer((16,), "float32")): @tvm.script.ir_module class MockModule: @T.prim_func(private=True) - def func1(A: T.Buffer((16,), "float32")): + def func1(A: T.Tensor((16,), "float32")): for i in T.serial(16): if i == 5: if i == 5: A[i] = 0.0 @T.prim_func(private=True) - def func2(A: T.Buffer((32,), "float32")): + def func2(A: T.Tensor((32,), "float32")): for i in T.serial(32): if i == 15: if i == 15: @@ -239,9 +239,9 @@ def add(a: T.int32, b: T.int32) -> T.int32: @T.prim_func def main( - A: T.Buffer((128, 128), "int32"), - B: T.Buffer((128, 128), "int32"), - C: T.Buffer((128, 128), "int32"), + A: T.Tensor((128, 128), "int32"), + B: T.Tensor((128, 128), "int32"), + C: T.Tensor((128, 128), "int32"), ): T.func_attr({"global_symbol": "main"}) length: T.let[T.int32] = Before.add(64, 64) # Call from host @@ -263,9 +263,9 @@ def add_host(a: T.int32, b: T.int32) -> T.int32: @T.prim_func def main( - A: T.Buffer((128, 128), "int32"), - B: T.Buffer((128, 128), "int32"), - C: T.Buffer((128, 128), "int32"), + A: T.Tensor((128, 128), "int32"), + B: T.Tensor((128, 128), "int32"), + C: T.Tensor((128, 128), "int32"), ): T.func_attr( { @@ -294,7 +294,7 @@ def add(a: T.int32, b: T.int32) -> T.int32: return a + b @T.prim_func - def main(A: T.Buffer((1,), "int32")): + def main(A: T.Tensor((1,), "int32")): T.func_attr({"global_symbol": "main"}) host_value: T.let[T.int32] = Before.add(1, 2) T.device_entry() @@ -314,7 +314,7 @@ def add_host(a: T.int32, b: T.int32) -> T.int32: return a + b @T.prim_func - def main(A: T.Buffer((1,), "int32")): + def main(A: T.Tensor((1,), "int32")): T.func_attr( { "global_symbol": "main", diff --git a/tests/python/tirx-transform/test_tir_transform_lower_intrin.py b/tests/python/tirx-transform/test_tir_transform_lower_intrin.py index 3e060ee61b1a..5ff9771a70f9 100644 --- a/tests/python/tirx-transform/test_tir_transform_lower_intrin.py +++ b/tests/python/tirx-transform/test_tir_transform_lower_intrin.py @@ -60,9 +60,9 @@ def check_value(expr, variables, data, fref): # Build input and output buffers input_bufs = [ - tvm.tirx.decl_buffer((n,), dtype=variables[i].ty, name=f"v{i}") for i in range(num_vars) + tvm.tirx.decl_tensor((n,), dtype=variables[i].ty, name=f"v{i}") for i in range(num_vars) ] - out_buf = tvm.tirx.decl_buffer((n,), dtype=expr.ty, name="C") + out_buf = tvm.tirx.decl_tensor((n,), dtype=expr.ty, name="C") # Build loop body: for each i, bind variables[j] = input_bufs[j][i], then store expr to out loop_var = tvm.tirx.Var("i", "int32") @@ -138,7 +138,7 @@ def collect(node): def test_lower_vector_access_ptr(): - buffer = tvm.tirx.decl_buffer((8,), "float32x2", name="A") + buffer = tvm.tirx.decl_tensor((8,), "float32x2", name="A") access_ptr = buffer.access_ptr(access_mask=3, offset=2, extent=4) assert access_ptr.op.name == "tirx.tvm_access_ptr" @@ -154,7 +154,7 @@ def test_lower_vector_access_ptr(): lowered_body = tvm.tirx.transform.LowerIntrin()(mod)["main"].body assert isinstance(lowered_body, tvm.tirx.SeqStmt) alias = lowered_body.seq[0] - assert _is_buffer_binding(alias, "tirx.decl_buffer") + assert _is_buffer_binding(alias, "tirx.decl_tensor") assert alias.value.args[0].op.name == "tirx.buffer_data" assert alias.value.args[0].args[0].same_as(buffer) lowered = lowered_body.seq[1].value @@ -176,7 +176,7 @@ def test_lower_vector_access_ptr(): @pytest.mark.skipif(not env.has_llvm(), reason="need llvm") def test_lower_vector_access_ptr_with_padded_vector_dtype(): - buffer = tvm.tirx.decl_buffer((8,), "float32x3", name="A") + buffer = tvm.tirx.decl_tensor((8,), "float32x3", name="A") access_ptr = buffer.access_ptr(access_mask=1, offset=2, extent=4) body = tvm.tirx.Evaluate(tvm.tirx.call_extern("void", "consume", access_ptr)) func = tvm.tirx.PrimFunc([buffer], body).with_attr("global_symbol", "main") @@ -185,7 +185,7 @@ def test_lower_vector_access_ptr_with_padded_vector_dtype(): def test_lower_buffer_data_access_ptr_preserves_buffer_identity(): - buffer = tvm.tirx.decl_buffer((16,), "float32", "buffer") + buffer = tvm.tirx.decl_tensor((16,), "float32", "buffer") access = tvm.tirx.tvm_access_ptr("float32", buffer.data, 3, 8, 1) func = tvm.tirx.PrimFunc([buffer], tvm.tirx.Evaluate(access)).with_attr( @@ -202,7 +202,7 @@ def test_lower_buffer_data_access_ptr_preserves_buffer_identity(): @pytest.mark.parametrize("shape", [(), (2, 4)]) def test_lower_access_ptr_uses_flat_alias_for_non_1d_buffer(shape): - buffer = tvm.tirx.decl_buffer(shape, "float32", "buffer") + buffer = tvm.tirx.decl_tensor(shape, "float32", "buffer") access = buffer.access_ptr(access_mask=1) func = tvm.tirx.PrimFunc([buffer], tvm.tirx.Evaluate(access)).with_attr( "target", tvm.target.Target("llvm") @@ -211,7 +211,7 @@ def test_lower_access_ptr_uses_flat_alias_for_non_1d_buffer(shape): lowered = tvm.tirx.transform.LowerIntrin()(tvm.IRModule.from_expr(func))["main"].body assert isinstance(lowered, tvm.tirx.SeqStmt) alias = lowered.seq[0] - assert _is_buffer_binding(alias, "tirx.decl_buffer") + assert _is_buffer_binding(alias, "tirx.decl_tensor") assert len(alias.var.ty.shape) == 1 load = lowered.seq[1].value.args[0] assert isinstance(load, tvm.ir.TensorLoad) diff --git a/tests/python/tirx-transform/test_tir_transform_lower_tvm_builtin.py b/tests/python/tirx-transform/test_tir_transform_lower_tvm_builtin.py index d5baaf3d793e..7ed2f330beb9 100644 --- a/tests/python/tirx-transform/test_tir_transform_lower_tvm_builtin.py +++ b/tests/python/tirx-transform/test_tir_transform_lower_tvm_builtin.py @@ -35,9 +35,9 @@ def test_lower_call_packed(): class Before: @T.prim_func def main( - A: T.Buffer((64, 64), "float32"), - B: T.Buffer((64, 64), "float32"), - C: T.Buffer((64, 64), "float32"), + A: T.Tensor((64, 64), "float32"), + B: T.Tensor((64, 64), "float32"), + C: T.Tensor((64, 64), "float32"), ): T.func_attr({"target": tvm.target.Target("llvm")}) T.attr("", "device_id", T.int32(0)) @@ -47,14 +47,14 @@ def main( class Expected: @T.prim_func def main( - A: T.Buffer((64, 64), "float32"), - B: T.Buffer((64, 64), "float32"), - C: T.Buffer((64, 64), "float32"), + A: T.Tensor((64, 64), "float32"), + B: T.Tensor((64, 64), "float32"), + C: T.Tensor((64, 64), "float32"), ): T.func_attr({"target": tvm.target.Target("llvm")}) stack_ffi_any: T.let[T.handle] = T.tvm_stack_alloca("tvm_ffi_any", 4) stack_array: T.let[T.handle] = T.tvm_stack_alloca("array", 3) - stack_shape = T.decl_buffer( + stack_shape = T.decl_tensor( (T.int64(6),), "int64", data=T.tvm_stack_alloca("shape", 6), layout=None ) stack_shape[0] = T.int64(64) @@ -158,7 +158,7 @@ def packed_echo(value): ) def build_tir(): - Ab = tvm.tirx.decl_buffer((2,), "float32") + Ab = tvm.tirx.decl_tensor((2,), "float32") # Build statements using direct TIR construction (no ir_builder) # 1. Store packed_echo(const) result into Ab[0] @@ -187,13 +187,13 @@ def build_tir(): def test_lower_overflow_int32(): @T.prim_func(check_well_formed=False) - def variance4(rxplaceholder: T.Buffer((T.int64(1), T.int64(32), T.int64(25690112)), "float32")): + def variance4(rxplaceholder: T.Tensor((T.int64(1), T.int64(32), T.int64(25690112)), "float32")): T.func_attr({"global_symbol": "variance4", "tirx.noalias": True}) - rxplaceholder_red = T.alloc_buffer((32,), "float32") - T_subtract = T.alloc_buffer((822083584,), "float32") - rxplaceholder_red_1 = T.decl_buffer((T.int64(32),), data=rxplaceholder_red.data) - rxplaceholder_1 = T.decl_buffer((T.int64(822083584),), data=rxplaceholder.data) - T_subtract_1 = T.decl_buffer((T.int64(822083584),), data=T_subtract.data) + rxplaceholder_red = T.alloc_tensor((32,), "float32") + T_subtract = T.alloc_tensor((822083584,), "float32") + rxplaceholder_red_1 = T.decl_tensor((T.int64(32),), data=rxplaceholder_red.data) + rxplaceholder_1 = T.decl_tensor((T.int64(822083584),), data=rxplaceholder.data) + T_subtract_1 = T.decl_tensor((T.int64(822083584),), data=T_subtract.data) for ax1, ax2 in T.grid(32, 25690112): cse_v1: T.let[T.int32] = ax1 * 25690112 + ax2 T_subtract_1[cse_v1] = rxplaceholder_1[cse_v1] - rxplaceholder_red_1[ax1] @@ -220,8 +220,8 @@ def main(): T.func_attr({"target": T.target("llvm")}) T.attr("dummy", "device_type", 2) # kDLCuda T.attr("dummy", "device_id", 0) - ptr = T.alloc_buffer((16,), "float32") - buf = T.decl_buffer(16, "float32", data=ptr.data) + ptr = T.alloc_tensor((16,), "float32") + buf = T.decl_tensor(16, "float32", data=ptr.data) buf[0] = 0.0 After = tvm.tirx.transform.LowerTVMBuiltin()(Before) @@ -230,8 +230,8 @@ def main(): # Should contain TVMBackendAllocWorkspace and TVMBackendFreeWorkspace assert "TVMBackendAllocWorkspace" in script_output assert "TVMBackendFreeWorkspace" in script_output - # DeclBuffer should appear as a flat statement - assert "T.decl_buffer" in script_output + # DeclTensor should appear as a flat statement + assert "T.decl_tensor" in script_output def test_lower_cpu_allocation(): @@ -244,8 +244,8 @@ def main(): T.func_attr({"target": T.target("llvm")}) T.attr("dummy", "device_type", 1) # kDLCPU T.attr("dummy", "device_id", 0) - ptr = T.alloc_buffer((16,), "float32") - buf = T.decl_buffer(16, "float32", data=ptr.data) + ptr = T.alloc_tensor((16,), "float32") + buf = T.decl_tensor(16, "float32", data=ptr.data) buf[0] = 0.0 @I.ir_module @@ -253,8 +253,8 @@ class Expected: @T.prim_func def main(): T.func_attr({"target": T.target("llvm")}) - ptr = T.alloc_buffer((16,), "float32") - buf = T.decl_buffer(16, "float32", data=ptr.data) + ptr = T.alloc_tensor((16,), "float32") + buf = T.decl_tensor(16, "float32", data=ptr.data) buf[0] = 0.0 After = tvm.tirx.transform.LowerTVMBuiltin()(Before) @@ -270,8 +270,8 @@ class Before: def main(): T.func_attr({"target": T.target("llvm")}) T.attr("dummy", "device_type", 2) # kDLCuda - ptr = T.alloc_buffer((16,), "float32") - buf = T.decl_buffer(16, "float32", data=ptr.data) + ptr = T.alloc_tensor((16,), "float32") + buf = T.decl_tensor(16, "float32", data=ptr.data) buf[0] = 0.0 with pytest.raises(RuntimeError): @@ -294,8 +294,8 @@ class Before: def main(): T.func_attr({"tirx.is_host_func": True}) T.attr("dummy", "device_id", 0) - ptr = T.alloc_buffer((1024 * 1024,), "float32") - buf = T.decl_buffer(1024 * 1024, "float32", data=ptr.data) + ptr = T.alloc_tensor((1024 * 1024,), "float32") + buf = T.decl_tensor(1024 * 1024, "float32", data=ptr.data) buf[0] = 0.0 with pytest.raises(RuntimeError): @@ -316,8 +316,8 @@ class Before: @T.prim_func def main(): T.func_attr({"target": T.target("llvm")}) - ptr = T.alloc_buffer((16,), "float32") - buf = T.decl_buffer(16, "float32", data=ptr.data) + ptr = T.alloc_tensor((16,), "float32") + buf = T.decl_tensor(16, "float32", data=ptr.data) buf[0] = 0.0 # Expected is same as before for this transform 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 eca55da87527..2ff14c385f5a 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 @@ -84,7 +84,7 @@ def test_target_host_removed(): @I.ir_module class before: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"global_symbol": "main", "target": T.target("cuda", host=host)}) T.evaluate(0) @@ -105,7 +105,7 @@ def test_internal_subroutine_call(): @I.ir_module class before: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("llvm", host="llvm")}) before.subroutine(A.data) @@ -138,7 +138,7 @@ def test_subroutine_call_to_externally_visible_subroutine(): @I.ir_module class before: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"global_symbol": "main", "target": T.target("llvm", host="llvm")}) before.subroutine(A.data) @@ -461,7 +461,7 @@ def test_forward_reference_symbolic_variable(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((batch_size + 1,), "int32"), B: T.Buffer((batch_size,), "int32")): + def main(A: T.Tensor((batch_size + 1,), "int32"), B: T.Tensor((batch_size,), "int32")): T.func_attr({"target": T.target("llvm", host="llvm")}) for i in range(batch_size): @@ -478,7 +478,7 @@ def test_buffer_alignment_attached_to_buffer_var(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((16,), "float32", align=64)): + def main(A: T.Tensor((16,), "float32", align=64)): T.func_attr({"global_symbol": "main", "target": T.target("llvm", host="llvm")}) T.evaluate(A[0]) @@ -489,7 +489,7 @@ def main(A: T.Buffer((16,), "float32", align=64)): def collect(node): if isinstance(node, tirx.AttrStmt) and node.attr_key == "storage_alignment": alignment_nodes.append(node.node) - if _is_buffer_binding(node, "tirx.decl_buffer"): + if _is_buffer_binding(node, "tirx.decl_tensor"): declared_buffers.append(node.var) tvm_ffi.structural_walk(after.body, collect) diff --git a/tests/python/tirx-transform/test_tir_transform_narrow_datatype.py b/tests/python/tirx-transform/test_tir_transform_narrow_datatype.py index d26cd17fe8be..91b3d1087ddf 100644 --- a/tests/python/tirx-transform/test_tir_transform_narrow_datatype.py +++ b/tests/python/tirx-transform/test_tir_transform_narrow_datatype.py @@ -48,7 +48,7 @@ def check_const(m, n, target_bits, target_dtype): """Check with constant values using TVMScript closure.""" @T.prim_func - def func(A: T.Buffer((m * n,), "float32"), B: T.Buffer((m * n,), "float32")): + def func(A: T.Tensor((m * n,), "float32"), B: T.Tensor((m * n,), "float32")): for i in T.serial(m): for j in T.serial(n): B[i * n + j] = A[i * n + j] + T.float32(1) @@ -63,8 +63,8 @@ def check_symbolic(m_dtype, n_dtype, target_bits, target_dtype): @T.prim_func def func(A: T.handle("float32"), B: T.handle("float32"), m: T.int32, n: T.int32): - A_buf = T.decl_buffer((m * n,), "float32", data=A) - B_buf = T.decl_buffer((m * n,), "float32", data=B) + A_buf = T.decl_tensor((m * n,), "float32", data=A) + B_buf = T.decl_tensor((m * n,), "float32", data=B) for i in T.serial(m): for j in T.serial(n): B_buf[i * n + j] = A_buf[i * n + j] + T.float32(1) @@ -73,8 +73,8 @@ def func(A: T.handle("float32"), B: T.handle("float32"), m: T.int32, n: T.int32) @T.prim_func def func(A: T.handle("float32"), B: T.handle("float32"), m: T.int64, n: T.int64): - A_buf = T.decl_buffer((m * n,), "float32", data=A) - B_buf = T.decl_buffer((m * n,), "float32", data=B) + A_buf = T.decl_tensor((m * n,), "float32", data=A) + B_buf = T.decl_tensor((m * n,), "float32", data=B) for i in T.serial(m): for j in T.serial(n): B_buf[i * n + j] = A_buf[i * n + j] + T.float32(1) @@ -103,7 +103,7 @@ def test_thread_axis(): # and checks the dtype of thread axis variables after narrowing. def check_const(m, n, target_bits, target_dtype): @T.prim_func - def func(A: T.Buffer((m * n,), "float32"), B: T.Buffer((m * n,), "float32")): + def func(A: T.Tensor((m * n,), "float32"), B: T.Tensor((m * n,), "float32")): bx = T.launch_thread("blockIdx.x", m) tx = T.launch_thread("threadIdx.x", n) B[bx * n + tx] = A[bx * n + tx] + T.float32(1) @@ -139,8 +139,8 @@ def check(m, lanes, target_bits, target_dtype): @T.prim_func def func( - A: T.Buffer((m,), vec_dtype), - B: T.Buffer((m,), vec_dtype), + A: T.Tensor((m,), vec_dtype), + B: T.Tensor((m,), vec_dtype), ): for i in T.serial(m): B[i] = A[i] + T.Broadcast(T.float32(1), lanes) @@ -168,8 +168,8 @@ def check(m, n, target_bits, target_dtype): # The index may overflow in B, while not in A @T.prim_func def func( - A: T.Buffer((m * n,), "float32"), - B: T.Buffer((m * n * 2,), "float32"), + A: T.Tensor((m * n,), "float32"), + B: T.Tensor((m * n * 2,), "float32"), ): for i in T.serial(m): for j in T.serial(n): @@ -187,7 +187,7 @@ def func( def test_condition(): @T.prim_func - def before(A: T.Buffer((128,), "float32"), B: T.Buffer((130,), "float32")): + def before(A: T.Tensor((128,), "float32"), B: T.Tensor((130,), "float32")): for i, j in T.grid(T.int64(2), T.int64(65)): if i * T.int64(65) + j >= T.int64(0) and i * T.int64(65) + j < T.int64(128): A[i * T.int64(65) + j] = 0.0 @@ -199,7 +199,7 @@ def before(A: T.Buffer((128,), "float32"), B: T.Buffer((130,), "float32")): ) @T.prim_func - def expected_after(A: T.Buffer(128, "float32"), B: T.Buffer(130, "float32")): + def expected_after(A: T.Tensor(128, "float32"), B: T.Tensor(130, "float32")): for i, j in T.grid(2, 65): if i * 65 + j >= 0 and i * 65 + j < 128: A[i * 65 + j] = T.float32(0) @@ -216,13 +216,13 @@ def expected_after(A: T.Buffer(128, "float32"), B: T.Buffer(130, "float32")): def test_block(): @T.prim_func - def before(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def before(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): for i in T.serial(0, T.int64(16)): for j in T.serial(0, T.int64(8)): B[i * T.int64(8) + j] = A[i * T.int64(8) + j] + T.float32(1) @T.prim_func - def expected_after(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")): + def expected_after(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")): for i in T.serial(0, T.int32(16)): for j in T.serial(0, T.int32(8)): B[i * T.int32(8) + j] = A[i * T.int32(8) + j] + T.float32(1) @@ -235,7 +235,7 @@ def expected_after(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32" def test_avg_pool2d(): @T.prim_func - def before(PSUM: T.Buffer((313600,), "int32"), PAVG: T.Buffer((313600,), "int32")): + def before(PSUM: T.Tensor((313600,), "int32"), PAVG: T.Tensor((313600,), "int32")): for j in T.parallel(T.int64(0), T.int64(280)): for i in T.serial(T.int64(0), T.int64(35)): for vi in T.vectorized(T.int64(0), T.int64(32)): @@ -268,7 +268,7 @@ def before(PSUM: T.Buffer((313600,), "int32"), PAVG: T.Buffer((313600,), "int32" ) @T.prim_func - def expected_after(PSUM: T.Buffer((313600,), "int32"), PAVG: T.Buffer((313600,), "int32")): + def expected_after(PSUM: T.Tensor((313600,), "int32"), PAVG: T.Tensor((313600,), "int32")): for j in T.parallel(T.int32(0), T.int32(280)): for i in T.serial(T.int32(0), T.int32(35)): for vi in T.vectorized(T.int32(0), T.int32(32)): @@ -298,12 +298,12 @@ def expected_after(PSUM: T.Buffer((313600,), "int32"), PAVG: T.Buffer((313600,), def test_narrow_i64_valued_bufferload_index_to_i32(): @T.prim_func - def before(A: T.Buffer((16,), "int64")): + def before(A: T.Tensor((16,), "int64")): for i in range(T.int64(15)): A[i + T.int64(1)] = A[i] + T.int64(1) @T.prim_func - def expect(A: T.Buffer((16,), "int64")): + def expect(A: T.Tensor((16,), "int64")): for i in range(15): A[i + 1] = A[i] + T.int64(1) diff --git a/tests/python/tirx-transform/test_tir_transform_pointer_value_type_rewrite.py b/tests/python/tirx-transform/test_tir_transform_pointer_value_type_rewrite.py index c9774cbfc612..38b25c9aea11 100644 --- a/tests/python/tirx-transform/test_tir_transform_pointer_value_type_rewrite.py +++ b/tests/python/tirx-transform/test_tir_transform_pointer_value_type_rewrite.py @@ -39,8 +39,8 @@ def test_rewrite_to_shuffle_0(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((16,), "float32"), B: T.Buffer((4,), "float32")): - A_local = T.alloc_buffer((16,), scope="local") + def main(A: T.Tensor((16,), "float32"), B: T.Tensor((4,), "float32")): + A_local = T.alloc_tensor((16,), scope="local") for i in range(4): A_local[T.ramp(i * 4, 1, 4)] = A[T.ramp(i * 4, 1, 4)] for i in range(4): @@ -49,8 +49,8 @@ def main(A: T.Buffer((16,), "float32"), B: T.Buffer((4,), "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((4,), "float32x4", layout=None), B: T.Buffer((4,), "float32")): - A_local = T.alloc_buffer((4,), "float32x4", scope="local", layout=None) + def main(A: T.Tensor((4,), "float32x4", layout=None), B: T.Tensor((4,), "float32")): + A_local = T.alloc_tensor((4,), "float32x4", scope="local", layout=None) for i in range(4): A_local[T.Div(i * 4, 4)] = A[T.Div(i * 4, 4)] for i in range(4): @@ -71,8 +71,8 @@ def test_rewrite_to_shuffle_1(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((8,), "float32"), B: T.Buffer((1,), "float32")): - A_local = T.alloc_buffer((8,), scope="local") + def main(A: T.Tensor((8,), "float32"), B: T.Tensor((1,), "float32")): + A_local = T.alloc_tensor((8,), scope="local") A_local[T.ramp(0, 1, 4)] = A[T.ramp(0, 1, 4)] A_local[T.ramp(4, 1, 4)] = A[T.ramp(4, 1, 4)] B[0] = ( @@ -89,8 +89,8 @@ def main(A: T.Buffer((8,), "float32"), B: T.Buffer((1,), "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((2,), "float32x4", layout=None), B: T.Buffer((1,), "float32")): - A_local = T.alloc_buffer((2,), "float32x4", scope="local", layout=None) + def main(A: T.Tensor((2,), "float32x4", layout=None), B: T.Tensor((1,), "float32")): + A_local = T.alloc_tensor((2,), "float32x4", scope="local", layout=None) A_local[0] = A[0] A_local[1] = A[1] B[0] = ( @@ -114,7 +114,7 @@ def test_address_of(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def main(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): for i in range(4): T.evaluate(T.address_of(A[i * 4])) B[T.ramp(i * 4, 1, 4)] = A[T.ramp(i * 4, 1, 4)] @@ -122,7 +122,7 @@ def main(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((16,), "float32"), B: T.Buffer((4,), "float32x4", layout=None)): + def main(A: T.Tensor((16,), "float32"), B: T.Tensor((4,), "float32x4", layout=None)): for i in range(4): T.evaluate(T.address_of(A[i * 4])) B[T.Div(i * 4, 4)] = A[T.ramp(i * 4, 1, 4)] @@ -137,7 +137,7 @@ def test_scalar_read_without_write(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((16,), "float32")): + def main(A: T.Tensor((16,), "float32")): for i in range(4): T.evaluate(A[i * 4]) @@ -145,7 +145,7 @@ def main(A: T.Buffer((16,), "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer((16,), "float32")): + def main(A: T.Tensor((16,), "float32")): for i in range(4): T.evaluate(A[i * 4]) @@ -159,9 +159,9 @@ def test_decl_buffer_alias_chain_uses_flat_root_map(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((16,), "float32")): - A_view = T.decl_buffer((16,), "float32", data=A.data) - A_view_2 = T.decl_buffer((16,), "float32", data=A_view.data) + def main(A: T.Tensor((16,), "float32")): + A_view = T.decl_tensor((16,), "float32", data=A.data) + A_view_2 = T.decl_tensor((16,), "float32", data=A_view.data) for i in range(4): A_view_2[T.ramp(i * 4, 1, 4)] = T.broadcast(T.float32(1), 4) @@ -176,7 +176,7 @@ def main(A: T.Buffer((16,), "float32")): func.body, lambda node: ( decl_buffers.append(node) - if _is_buffer_binding(node, "tirx.decl_buffer") + if _is_buffer_binding(node, "tirx.decl_tensor") else buffer_stores.append(node) if isinstance(node, tvm.tirx.BufferStore) else None diff --git a/tests/python/tirx-transform/test_tir_transform_prim_func_pass.py b/tests/python/tirx-transform/test_tir_transform_prim_func_pass.py index db4d7cd3e372..5744810a0e45 100644 --- a/tests/python/tirx-transform/test_tir_transform_prim_func_pass.py +++ b/tests/python/tirx-transform/test_tir_transform_prim_func_pass.py @@ -31,7 +31,7 @@ def transform_function(self, func, mod, ctx): x = tvm.tirx.Var("x", "int32") y = tvm.tirx.Var("y", "int32") - b = tvm.tirx.decl_buffer((x,), "float32") + b = tvm.tirx.decl_tensor((x,), "float32") stmt = tvm.tirx.SeqStmt([tvm.tirx.Bind(x, 10), tvm.tirx.Evaluate(x + 1)]) func = tvm.tirx.PrimFunc([x, y, b], stmt) diff --git a/tests/python/tirx-transform/test_tir_transform_remove_assume.py b/tests/python/tirx-transform/test_tir_transform_remove_assume.py index 3e92b7c5e8b1..122bee4eaa94 100644 --- a/tests/python/tirx-transform/test_tir_transform_remove_assume.py +++ b/tests/python/tirx-transform/test_tir_transform_remove_assume.py @@ -27,14 +27,14 @@ def test_remove_assume(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(1, "int32")): + def main(A: T.Tensor(1, "int32")): T.evaluate(T.assume(A[0] == 5)) A[0] = 10 @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(1, "int32")): + def main(A: T.Tensor(1, "int32")): A[0] = 10 After = tvm.tirx.transform.RemoveAssume()(Before) @@ -47,7 +47,7 @@ def test_remove_assume_loop(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(16, "int32")): + def main(A: T.Tensor(16, "int32")): for i in T.serial(16): T.evaluate(T.assume(A[i] == 0)) @@ -57,7 +57,7 @@ def main(A: T.Buffer(16, "int32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(16, "int32")): + def main(A: T.Tensor(16, "int32")): for i in T.serial(16): A[i] = 10 diff --git a/tests/python/tirx-transform/test_tir_transform_remove_no_op.py b/tests/python/tirx-transform/test_tir_transform_remove_no_op.py index c3b71c1bed79..29443f964a02 100644 --- a/tests/python/tirx-transform/test_tir_transform_remove_no_op.py +++ b/tests/python/tirx-transform/test_tir_transform_remove_no_op.py @@ -33,7 +33,7 @@ def test_remove_no_op(): m = tvm.tirx.Var("m", "int32") n = tvm.tirx.Var("n", "int32") dtype = "int64" - Ab = tvm.tirx.decl_buffer((n,), dtype) + Ab = tvm.tirx.decl_tensor((n,), dtype) stmt = tvm.tirx.For( i, 0, @@ -76,7 +76,7 @@ def test_remove_no_op(): def test_remove_no_op_with_invalid_extent(): @T.prim_func - def main(A: T.Buffer((16), "int32"), B: T.Buffer((16), "int32")) -> None: + def main(A: T.Tensor((16), "int32"), B: T.Tensor((16), "int32")) -> None: for i in T.serial(16): for j in T.serial(i - 20): B[i] = A[i] + j @@ -119,12 +119,12 @@ def test_remove_zero_extent_loop(): """A for-loop with no extent is a no-op.""" @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): + def before(A: T.Tensor(16, "int32")): for i in T.serial(0): A[i] = 42 @T.prim_func(private=True) - def expected(A: T.Buffer(16, "int32")): + def expected(A: T.Tensor(16, "int32")): T.evaluate(0) mod = tvm.IRModule.from_expr(before) @@ -140,13 +140,13 @@ def test_remove_unused_let(): """ @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): + def before(A: T.Tensor(16, "int32")): x = 5 for i in T.serial(16): A[i] = 0 @T.prim_func(private=True) - def expected(A: T.Buffer(16, "int32")): + def expected(A: T.Tensor(16, "int32")): x = 5 for i in T.serial(16): A[i] = 0 @@ -164,13 +164,13 @@ def test_remove_let_used_only_in_no_op(): """ @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): + def before(A: T.Tensor(16, "int32")): x = 5 for i in T.serial(0): A[i] = x @T.prim_func(private=True) - def expected(A: T.Buffer(16, "int32")): + def expected(A: T.Tensor(16, "int32")): x = 5 T.evaluate(0) @@ -201,7 +201,7 @@ def test_remove_empty_then_case(): """A no-op then_case can be removed.""" @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): + def before(A: T.Tensor(16, "int32")): for i in T.serial(16): if i < 8: T.evaluate(0) @@ -209,7 +209,7 @@ def before(A: T.Buffer(16, "int32")): A[i] = 42 @T.prim_func(private=True) - def expected(A: T.Buffer(16, "int32")): + def expected(A: T.Tensor(16, "int32")): for i in T.serial(16): if not (i < 8): A[i] = 42 @@ -223,7 +223,7 @@ def test_remove_empty_else_case(): """A no-op else_case can be removed.""" @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): + def before(A: T.Tensor(16, "int32")): for i in T.serial(16): if i < 8: A[i] = 42 @@ -231,7 +231,7 @@ def before(A: T.Buffer(16, "int32")): T.evaluate(0) @T.prim_func(private=True) - def expected(A: T.Buffer(16, "int32")): + def expected(A: T.Tensor(16, "int32")): for i in T.serial(16): if i < 8: A[i] = 42 @@ -248,7 +248,7 @@ def test_suppress_removal_of_unused_write(): """ @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): + def before(A: T.Tensor(16, "int32")): for i in T.serial(16): A[i] = 100 A[i] = 42 @@ -262,7 +262,7 @@ def test_keep_first_write_when_used(): """For two sequential writes, keep the first if it is used""" @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): + def before(A: T.Tensor(16, "int32")): for i in T.serial(16): A[i] = 100 A[i] = A[i] + 1 @@ -280,7 +280,7 @@ def test_keep_partially_overwritten_loop(): """ @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): + def before(A: T.Tensor(16, "int32")): for i in T.serial(16): A[i] = 100 @@ -297,11 +297,11 @@ def test_remove_read_write(): """Writing a value to the same location as was just read is a no-op.""" @T.prim_func(private=True) - def before(A: T.Buffer(1, "int32")): + def before(A: T.Tensor(1, "int32")): A[0] = A[0] @T.prim_func(private=True) - def expected(A: T.Buffer(1, "int32")): + def expected(A: T.Tensor(1, "int32")): T.evaluate(0) mod = tvm.IRModule.from_expr(before) @@ -313,7 +313,7 @@ def test_keep_read_write_to_different_indices(): """Writing a value to a different index should not be removed""" @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): + def before(A: T.Tensor(16, "int32")): for i in T.serial(15): A[i] = A[i + 1] @@ -332,13 +332,13 @@ def test_remove_read_write_same_index_different_expression(): """ @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): + def before(A: T.Tensor(16, "int32")): for io, ii in T.grid(4, 4): i: T.let[T.int32] = 4 * io + ii A[4 * io + ii] = A[i] @T.prim_func(private=True) - def expected(A: T.Buffer(16, "int32")): + def expected(A: T.Tensor(16, "int32")): for io in range(4): for ii in range(4): i: T.let[T.int32] = 4 * io + ii @@ -357,7 +357,7 @@ def test_remove_read_write_same_index_using_constraint(): """ @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): + def before(A: T.Tensor(16, "int32")): for i in T.serial(16): if i != 0: A[i] = A[i - 1] @@ -365,7 +365,7 @@ def before(A: T.Buffer(16, "int32")): A[i] = A[0] @T.prim_func(private=True) - def expected(A: T.Buffer(16, "int32")): + def expected(A: T.Tensor(16, "int32")): for i in T.serial(16): if i != 0: A[i] = A[i - 1] @@ -375,13 +375,13 @@ def expected(A: T.Buffer(16, "int32")): tvm.ir.assert_structural_equal(mod["main"], expected) -@pytest.mark.xfail(reason="Dead alloc removal not yet implemented for flat AllocBuffer") +@pytest.mark.xfail(reason="Dead alloc removal not yet implemented for flat AllocTensor") def test_remove_empty_temporary(): """An allocation with a no-op body is a no-op.""" @T.prim_func(private=True) def before(): - A = T.alloc_buffer((16,), "int32", scope="local") + A = T.alloc_tensor((16,), "int32", scope="local") T.evaluate(0) @T.prim_func(private=True) @@ -393,18 +393,18 @@ def expected(): tvm.ir.assert_structural_equal(mod["main"], expected) -@pytest.mark.xfail(reason="Dead alloc removal not yet implemented for flat AllocBuffer") +@pytest.mark.xfail(reason="Dead alloc removal not yet implemented for flat AllocTensor") def test_remove_empty_temporary_with_decl_buffer(): - """Remove DeclBuffer alongside Allocate + """Remove DeclTensor alongside Allocate - If an unused allocation is removed, any DeclBuffer instances that + If an unused allocation is removed, any DeclTensor instances that refer to it should also be removed. """ @T.prim_func(private=True) def before(): - A = T.decl_buffer([4, 4], "int32", scope="local") - A_flat = T.decl_buffer(16, "int32", scope="local", data=A.data) + A = T.decl_tensor([4, 4], "int32", scope="local") + A_flat = T.decl_tensor(16, "int32", scope="local", data=A.data) T.evaluate(0) @T.prim_func(private=True) @@ -421,13 +421,13 @@ def test_remove_unused_temporary(): """An unused allocation is a no-op.""" @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32")): - B = T.alloc_buffer((16,), "int32", scope="local") + def before(A: T.Tensor(16, "int32")): + B = T.alloc_tensor((16,), "int32", scope="local") for i in T.serial(16): A[i] = 1 @T.prim_func(private=True) - def expected(A: T.Buffer(16, "int32")): + def expected(A: T.Tensor(16, "int32")): for i in T.serial(16): A[i] = 1 @@ -442,7 +442,7 @@ def test_remove_unused_write_into_temporary(): @T.prim_func(private=True) def before(): - A = T.decl_buffer([16], "int32", scope="local") + A = T.decl_tensor([16], "int32", scope="local") for i in T.serial(16): A[i] = 0 @@ -459,8 +459,8 @@ def test_keep_used_write_into_temporary(): """A write into a temporary that is used later must be kept.""" @T.prim_func(private=True) - def before(B: T.Buffer(16, "int32")): - A = T.decl_buffer([16], "int32", scope="local") + def before(B: T.Tensor(16, "int32")): + A = T.decl_tensor([16], "int32", scope="local") for i in T.serial(16): A[i] = 0 @@ -477,8 +477,8 @@ def test_remove_write_into_temporary(): """A write that only impacts a temporary allocation is a no-op.""" @T.prim_func(private=True) - def before(A: T.Buffer(16, "int32"), C: T.Buffer(1, "int32")): - B = T.decl_buffer([16], "int32", scope="local") + def before(A: T.Tensor(16, "int32"), C: T.Tensor(1, "int32")): + B = T.decl_tensor([16], "int32", scope="local") for i in T.serial(16): B[i] = A[i] @@ -490,8 +490,8 @@ def before(A: T.Buffer(16, "int32"), C: T.Buffer(1, "int32")): B[i] = 0 @T.prim_func(private=True) - def expected(A: T.Buffer(16, "int32"), C: T.Buffer(1, "int32")): - B = T.decl_buffer([16], "int32", scope="local") + def expected(A: T.Tensor(16, "int32"), C: T.Tensor(1, "int32")): + B = T.decl_tensor([16], "int32", scope="local") for i in T.serial(16): B[i] = A[i] diff --git a/tests/python/tirx-transform/test_tir_transform_simplify.py b/tests/python/tirx-transform/test_tir_transform_simplify.py index 3c2b7dbcfb9f..21ae84dfcc18 100644 --- a/tests/python/tirx-transform/test_tir_transform_simplify.py +++ b/tests/python/tirx-transform/test_tir_transform_simplify.py @@ -26,8 +26,8 @@ def test_stmt_simplify(): @T.prim_func(private=True) def func(A: T.handle("float32"), C: T.handle("float32"), n: T.int32): - A_ptr = T.decl_buffer((10,), "float32", data=A) - C_ptr = T.decl_buffer((10,), "float32", data=C) + A_ptr = T.decl_tensor((10,), "float32", data=A) + C_ptr = T.decl_tensor((10,), "float32", data=C) n_val: T.let[T.int32] = 10 for i in T.serial(n_val): if i < 12: @@ -48,8 +48,8 @@ def func(A: T.handle("float32"), C: T.handle("float32"), n: T.int32): def test_thread_extent_simplify(): @T.prim_func(private=True) def func(A: T.handle("float32"), C: T.handle("float32"), n: T.int32): - A_ptr = T.decl_buffer((10,), "float32", data=A) - C_ptr = T.decl_buffer((10,), "float32", data=C) + A_ptr = T.decl_tensor((10,), "float32", data=A) + C_ptr = T.decl_tensor((10,), "float32", data=C) n_val: T.let[T.int32] = 10 for tx in T.thread_binding(n_val, thread="threadIdx.x"): for ty in T.thread_binding(1, thread="threadIdx.y"): @@ -73,8 +73,8 @@ def func(A: T.handle("float32"), C: T.handle("float32"), n: T.int32): def test_if_likely(): @T.prim_func(private=True) def func(A: T.handle("float32"), C: T.handle("float32"), n: T.int32): - A_ptr = T.decl_buffer((32,), "float32", data=A) - C_ptr = T.decl_buffer((1024,), "float32", data=C) + A_ptr = T.decl_tensor((32,), "float32", data=A) + C_ptr = T.decl_tensor((1024,), "float32", data=C) for tx in T.thread_binding(32, thread="threadIdx.x"): for ty in T.thread_binding(32, thread="threadIdx.y"): if T.likely(tx * 32 + ty < n): @@ -83,7 +83,7 @@ def func(A: T.handle("float32"), C: T.handle("float32"), n: T.int32): mod = tvm.IRModule.from_expr(func) body = tvm.tirx.transform.StmtSimplify()(mod)["main"].body - # With flat semantics, skip DeclBuffer/AllocBuffer siblings to find the For + # With flat semantics, skip DeclTensor/AllocTensor siblings to find the For if isinstance(body, tvm.tirx.SeqStmt): for_stmts = [s for s in body.seq if isinstance(s, tvm.tirx.For)] body = for_stmts[0] if for_stmts else body @@ -94,13 +94,13 @@ def func(A: T.handle("float32"), C: T.handle("float32"), n: T.int32): def test_loop_body_knows_dynamic_extent_is_positive(): @T.prim_func(private=True) - def before(A: T.Buffer((1,), "float32"), m: T.int32, n: T.int32): + def before(A: T.Tensor((1,), "float32"), m: T.int32, n: T.int32): for i in T.serial(m, n // 4): if n // 4 - m > 0: A[0] = 1.0 @T.prim_func(private=True) - def expected(A: T.Buffer((1,), "float32"), m: T.int32, n: T.int32): + def expected(A: T.Tensor((1,), "float32"), m: T.int32, n: T.int32): for i in T.serial(m, n // 4): A[0] = 1.0 @@ -132,11 +132,11 @@ def test_load_store_noop(): """Store of a value that was just read from the same location is a no-op.""" @T.prim_func(private=True) - def before(A: T.Buffer((1,), "float32")): + def before(A: T.Tensor((1,), "float32")): A[0] = A[0] @T.prim_func(private=True) - def expected(A: T.Buffer((1,), "float32")): + def expected(A: T.Tensor((1,), "float32")): T.evaluate(0) after = _apply_simplify(before) @@ -153,11 +153,11 @@ def test_load_store_noop_after_simplify(): """ @T.prim_func(private=True) - def before(A: T.Buffer((1,), "float32")): + def before(A: T.Tensor((1,), "float32")): A[0] = A[0] + (5.0 - 5.0) @T.prim_func(private=True) - def expected(A: T.Buffer((1,), "float32")): + def expected(A: T.Tensor((1,), "float32")): T.evaluate(0) after = _apply_simplify(before) @@ -173,14 +173,14 @@ def test_nested_condition(): """ @T.prim_func(private=True) - def before(A: T.Buffer((16,), "float32")): + def before(A: T.Tensor((16,), "float32")): for i in T.serial(16): if i == 5: if i == 5: A[i] = 0.0 @T.prim_func(private=True) - def expected(A: T.Buffer((16,), "float32")): + def expected(A: T.Tensor((16,), "float32")): for i in T.serial(16): if i == 5: A[i] = 0.0 @@ -197,14 +197,14 @@ def test_nested_provable_condition(): """ @T.prim_func(private=True) - def before(A: T.Buffer((16,), "float32")): + def before(A: T.Tensor((16,), "float32")): for i in T.serial(16): if i == 5: if i < 7: A[i] = 0.0 @T.prim_func(private=True) - def expected(A: T.Buffer((16,), "float32")): + def expected(A: T.Tensor((16,), "float32")): for i in T.serial(16): if i == 5: A[i] = 0.0 @@ -221,14 +221,14 @@ def test_nested_var_condition(): """ @T.prim_func(private=True) - def before(A: T.Buffer((16,), "float32"), n: T.int32): + def before(A: T.Tensor((16,), "float32"), n: T.int32): for i in T.serial(16): if i == n: if i == n: A[i] = 0.0 @T.prim_func(private=True) - def expected(A: T.Buffer((16,), "float32"), n: T.int32): + def expected(A: T.Tensor((16,), "float32"), n: T.int32): for i in T.serial(16): if i == n: A[i] = 0.0 @@ -247,7 +247,7 @@ def test_altered_buffer_contents(): """ @T.prim_func(private=True) - def before(A: T.Buffer((1,), "int32"), n: T.int32): + def before(A: T.Tensor((1,), "int32"), n: T.int32): if A[0] == n: A[0] = A[0] + 1 if A[0] == n: @@ -267,7 +267,7 @@ def test_negation_of_condition(): """ @T.prim_func(private=True) - def before(A: T.Buffer((16,), "int32")): + def before(A: T.Tensor((16,), "int32")): for i in T.serial(16): if i == 5: if i != 5: @@ -276,7 +276,7 @@ def before(A: T.Buffer((16,), "int32")): A[i] = 1 @T.prim_func(private=True) - def expected(A: T.Buffer((16,), "int32")): + def expected(A: T.Tensor((16,), "int32")): for i in T.serial(16): if i == 5: A[i] = 1 @@ -295,7 +295,7 @@ def test_negation_of_not_equal(): """ @T.prim_func(private=True) - def before(A: T.Buffer((16,), "int32")): + def before(A: T.Tensor((16,), "int32")): for i in T.serial(16): if i != 5: if i == 5: @@ -304,7 +304,7 @@ def before(A: T.Buffer((16,), "int32")): A[i] = 1 @T.prim_func(private=True) - def expected(A: T.Buffer((16,), "int32")): + def expected(A: T.Tensor((16,), "int32")): for i in T.serial(16): if i != 5: A[i] = 1 @@ -321,7 +321,7 @@ def test_negation_of_var_condition(): """ @T.prim_func(private=True) - def before(A: T.Buffer((16,), "int32"), n: T.int32): + def before(A: T.Tensor((16,), "int32"), n: T.int32): for i in T.serial(16): if i == n: if i != n: @@ -330,7 +330,7 @@ def before(A: T.Buffer((16,), "int32"), n: T.int32): A[i] = 1 @T.prim_func(private=True) - def expected(A: T.Buffer((16,), "int32"), n: T.int32): + def expected(A: T.Tensor((16,), "int32"), n: T.int32): for i in T.serial(16): if i == n: A[i] = 1 @@ -349,14 +349,14 @@ def test_literal_constraint_split_boolean_and(): """ @T.prim_func(private=True) - def before(A: T.Buffer((16, 16), "int32"), n: T.int32): + def before(A: T.Tensor((16, 16), "int32"), n: T.int32): for i, j in T.grid(16, 16): if i == n and j == n: if i == n: A[i, j] = 0 @T.prim_func(private=True) - def expected(A: T.Buffer((16, 16), "int32"), n: T.int32): + def expected(A: T.Tensor((16, 16), "int32"), n: T.int32): for i, j in T.grid(16, 16): if i == n and j == n: A[i, j] = 0 @@ -377,7 +377,7 @@ def test_literal_constraint_split_boolean_or(): """ @T.prim_func(private=True) - def before(A: T.Buffer((16, 16), "int32"), n: T.int32): + def before(A: T.Tensor((16, 16), "int32"), n: T.int32): for i, j in T.grid(16, 16): if i == n or j == n: A[i, j] = 0 @@ -388,7 +388,7 @@ def before(A: T.Buffer((16, 16), "int32"), n: T.int32): A[i, j] = 2 @T.prim_func(private=True) - def expected(A: T.Buffer((16, 16), "int32"), n: T.int32): + def expected(A: T.Tensor((16, 16), "int32"), n: T.int32): for i, j in T.grid(16, 16): if i == n or j == n: A[i, j] = 0 @@ -412,14 +412,14 @@ def test_prove_condition_using_let(): """ @T.prim_func(private=True) - def before(A: T.Buffer(4, "bool")): + def before(A: T.Tensor(4, "bool")): for i in T.serial(4): condition: T.let[T.bool] = i < 3 if condition or i >= 3: A[i] = condition @T.prim_func(private=True) - def expected(A: T.Buffer(4, "bool")): + def expected(A: T.Tensor(4, "bool")): for i in T.serial(4): condition: T.let[T.bool] = i < 3 # noqa: F841 A[i] = i < 3 @@ -436,7 +436,7 @@ def test_prove_let_condition(): """ @T.prim_func(private=True) - def before(A: T.Buffer(4, "bool")): + def before(A: T.Tensor(4, "bool")): for i in T.serial(4): condition: T.let[T.bool] = i < 3 if i < 3: @@ -444,7 +444,7 @@ def before(A: T.Buffer(4, "bool")): A[i] = condition @T.prim_func(private=True) - def expected(A: T.Buffer(4, "bool")): + def expected(A: T.Tensor(4, "bool")): for i in T.serial(4): condition: T.let[T.bool] = i < 3 # noqa: F841 if i < 3: @@ -463,7 +463,7 @@ def test_prove_repeated_let_condition(): """ @T.prim_func(private=True) - def before(A: T.Buffer(4, "bool")): + def before(A: T.Tensor(4, "bool")): for i in T.serial(4): condition: T.let[T.bool] = i < 3 if condition: @@ -471,7 +471,7 @@ def before(A: T.Buffer(4, "bool")): A[i] = condition @T.prim_func(private=True) - def expected(A: T.Buffer(4, "bool")): + def expected(A: T.Tensor(4, "bool")): for i in T.serial(4): condition: T.let[T.bool] = i < 3 # noqa: F841 if i < 3: @@ -483,13 +483,13 @@ def expected(A: T.Buffer(4, "bool")): def test_if_then_else_expr(): @T.prim_func(private=True) - def before(A: T.Buffer(16, "float32")): + def before(A: T.Tensor(16, "float32")): for i in T.serial(16): if i < 12: A[i] = T.if_then_else(i < 12, 1.0, 2.0) @T.prim_func(private=True) - def expected(A: T.Buffer(16, "float32")): + def expected(A: T.Tensor(16, "float32")): for i in T.serial(16): if i < 12: A[i] = 1.0 @@ -502,11 +502,11 @@ def test_ceil_log2_int(): """Simplify expressions resulting from topi.math.ceil_log2""" @T.prim_func(private=True) - def before(A: T.Buffer(1, "int32")): + def before(A: T.Tensor(1, "int32")): A[0] = T.cast(T.ceil(T.log2(T.cast(14, "float64"))), dtype="int32") @T.prim_func(private=True) - def expected(A: T.Buffer(1, "int32")): + def expected(A: T.Tensor(1, "int32")): A[0] = 4 after = _apply_simplify(before) @@ -521,7 +521,7 @@ def test_left_ceil_log2_lower_bound(): """ @T.prim_func(private=True) - def before(A: T.Buffer(16, "float32")): + def before(A: T.Tensor(16, "float32")): for i in T.serial(16): x: T.let[T.int32] = T.cast( T.ceil(T.log2(T.cast(i + 1024 + 1, "float64"))), @@ -531,7 +531,7 @@ def before(A: T.Buffer(16, "float32")): A[i] = 0.0 @T.prim_func(private=True) - def expected(A: T.Buffer(16, "float32")): + def expected(A: T.Tensor(16, "float32")): for i in T.serial(16): x: T.let[T.int32] = T.Cast( # noqa: F841 "int32", @@ -552,13 +552,13 @@ def test_left_shift_lower_bound(): """ @T.prim_func(private=True) - def before(A: T.Buffer(16, "float32")): + def before(A: T.Tensor(16, "float32")): for i in T.serial(16): if T.shift_left(1, i) >= 1: A[i] = 0.0 @T.prim_func(private=True) - def expected(A: T.Buffer(16, "float32")): + def expected(A: T.Tensor(16, "float32")): for i in T.serial(16): A[i] = 0.0 @@ -575,13 +575,13 @@ def test_left_shift_upper_bound(): """ @T.prim_func(private=True) - def before(A: T.Buffer(16, "float32")): + def before(A: T.Tensor(16, "float32")): for i in T.serial(16): if T.shift_left(31, i) <= 1015808: A[i] = 0.0 @T.prim_func(private=True) - def expected(A: T.Buffer(16, "float32")): + def expected(A: T.Tensor(16, "float32")): for i in T.serial(16): A[i] = 0.0 @@ -598,7 +598,7 @@ def test_left_shift_of_negative_value(): """ @T.prim_func(private=True) - def before(A: T.Buffer(16, "float32")): + def before(A: T.Tensor(16, "float32")): for i in T.serial(16): if -64 <= T.shift_left(-i, 4): A[i] = 0.0 @@ -618,7 +618,7 @@ def test_left_shift_by_negative_value(): """ @T.prim_func(private=True) - def before(A: T.Buffer(16, "float32")): + def before(A: T.Tensor(16, "float32")): for i in T.serial(16): if T.shift_left(16, -i) <= 16: A[i] = 0.0 @@ -709,7 +709,7 @@ def test_remove_transitively_provable_condition(): for priors, postulate, provable in test_cases: # well formed checker complains of undefined variables in condition @T.prim_func(private=True, check_well_formed=False) - def before_func(A: T.Buffer(1, "bool")): + def before_func(A: T.Tensor(1, "bool")): if priors: A[0] = postulate @@ -718,7 +718,7 @@ def before_func(A: T.Buffer(1, "bool")): if provable: # well formed checker complains of undefined variables in condition @T.prim_func(private=True, check_well_formed=False) - def expected_func(A: T.Buffer(1, "bool")): + def expected_func(A: T.Tensor(1, "bool")): if priors_simplified: A[0] = True @@ -727,7 +727,7 @@ def expected_func(A: T.Buffer(1, "bool")): # well formed checker complains of undefined variables in condition @T.prim_func(private=True, check_well_formed=False) - def expected_func(A: T.Buffer(1, "bool")): + def expected_func(A: T.Tensor(1, "bool")): if priors_simplified: A[0] = postulate_simplified @@ -737,7 +737,7 @@ def expected_func(A: T.Buffer(1, "bool")): def test_suppress_transitively_provable_condition(): @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), i: T.int32, j: T.int32, k: T.int32): + def before(A: T.Tensor(1, "bool"), i: T.int32, j: T.int32, k: T.int32): if i < j and j < k: A[0] = i < k @@ -751,11 +751,11 @@ def test_rewrite_as_and_of_ors(): """If enabled, rewrite boolean expressions into AND of OR""" @T.prim_func(private=True) - def before(A: T.Buffer(3, "bool")): + def before(A: T.Tensor(3, "bool")): T.evaluate(A[0] or (A[1] and A[2])) @T.prim_func(private=True) - def expected(A: T.Buffer(3, "bool")): + def expected(A: T.Tensor(3, "bool")): T.evaluate((A[0] or A[1]) and (A[0] or A[2])) after = _apply_simplify(before, convert_boolean_to_and_of_ors=True) @@ -766,7 +766,7 @@ def test_suppress_rewrite_as_and_of_ors(): """Only rewrite into AND of OR when allowed""" @T.prim_func(private=True) - def before(A: T.Buffer(3, "bool")): + def before(A: T.Tensor(3, "bool")): T.evaluate(A[0] or (A[1] and A[2])) expected = before @@ -786,11 +786,11 @@ def test_rewrite_as_and_of_ors_with_top_level_and(): """ @T.prim_func(private=True) - def before(A: T.Buffer(4, "bool")): + def before(A: T.Tensor(4, "bool")): T.evaluate((A[0] or A[1]) and (A[1] or (A[0] and A[2] and A[3]))) @T.prim_func(private=True) - def expected(A: T.Buffer(4, "bool")): + def expected(A: T.Tensor(4, "bool")): # If the simplification is applied to the OrNode, then a # redundant `(A[1] or A[0])` would't be canceled out. When # applying SimplifyAsAndOfOrs to the top-level AndNode, the @@ -820,11 +820,11 @@ def test_rewrite_as_and_of_ors_with_simplification_between_groups(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), i: T.int32, j: T.int32, k: T.int32): + def before(A: T.Tensor(1, "bool"), i: T.int32, j: T.int32, k: T.int32): A[0] = (i == 0 or j == 10 or k == 20) and (i == 0 or j == 10 or k != 30) @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), i: T.int32, j: T.int32, k: T.int32): + def expected(A: T.Tensor(1, "bool"), i: T.int32, j: T.int32, k: T.int32): A[0] = i == 0 or j == 10 or k == 20 after = _apply_simplify(before, convert_boolean_to_and_of_ors=True) @@ -840,11 +840,11 @@ def test_rewrite_as_and_of_ors_with_simplification_between_reordered_groups(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), i: T.int32, j: T.int32, k: T.int32): + def before(A: T.Tensor(1, "bool"), i: T.int32, j: T.int32, k: T.int32): A[0] = (i == 0 or j == 10 or k == 20) and (j == 10 or k != 30 or i == 0) @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), i: T.int32, j: T.int32, k: T.int32): + def expected(A: T.Tensor(1, "bool"), i: T.int32, j: T.int32, k: T.int32): A[0] = j == 10 or k == 20 or i == 0 after = _apply_simplify(before, convert_boolean_to_and_of_ors=True) @@ -860,11 +860,11 @@ def test_rewrite_as_and_of_or_using_simplification_across_and(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), i: T.int32, j: T.int32, k: T.int32): + def before(A: T.Tensor(1, "bool"), i: T.int32, j: T.int32, k: T.int32): A[0] = (k == 20) and ((i == 0 or j == 10) and (k != 30)) @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), i: T.int32, j: T.int32, k: T.int32): + def expected(A: T.Tensor(1, "bool"), i: T.int32, j: T.int32, k: T.int32): A[0] = (i == 0 or j == 10) and (k == 20) after = _apply_simplify(before, convert_boolean_to_and_of_ors=True) @@ -884,11 +884,11 @@ def test_rewrite_as_and_of_or_using_simplification_within_or(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), i: T.int32, j: T.int32, k: T.int32): + def before(A: T.Tensor(1, "bool"), i: T.int32, j: T.int32, k: T.int32): A[0] = (i == 20) or (j == 0) or (i != 30) @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), i: T.int32, j: T.int32, k: T.int32): + def expected(A: T.Tensor(1, "bool"), i: T.int32, j: T.int32, k: T.int32): A[0] = (j == 0) or (i != 30) after = _apply_simplify(before, convert_boolean_to_and_of_ors=True) @@ -916,12 +916,12 @@ def test_conditional_floor_mod(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), i: T.int32): + def before(A: T.Tensor(1, "bool"), i: T.int32): if T.floormod(0 - i, 2) == 0: A[0] = T.floormod(i, 2) == 0 @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), i: T.int32): + def expected(A: T.Tensor(1, "bool"), i: T.int32): if T.floormod(i, -2) == 0: A[0] = True @@ -938,11 +938,11 @@ def test_simplify_rhs_of_boolean_and_using_lhs(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), n: T.int32): + def before(A: T.Tensor(1, "bool"), n: T.int32): A[0] = n < 5 and n < 10 @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), n: T.int32): + def expected(A: T.Tensor(1, "bool"), n: T.int32): A[0] = n < 5 after = _apply_simplify(before, apply_constraints_to_boolean_branches=True) @@ -957,11 +957,11 @@ def test_simplify_lhs_of_boolean_and_using_rhs(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), n: T.int32): + def before(A: T.Tensor(1, "bool"), n: T.int32): A[0] = n < 10 and n < 5 @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), n: T.int32): + def expected(A: T.Tensor(1, "bool"), n: T.int32): A[0] = n < 5 after = _apply_simplify(before, apply_constraints_to_boolean_branches=True) @@ -977,11 +977,11 @@ def test_simplify_rhs_of_boolean_or_using_lhs(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), n: T.int32): + def before(A: T.Tensor(1, "bool"), n: T.int32): A[0] = n < 10 or n < 5 @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), n: T.int32): + def expected(A: T.Tensor(1, "bool"), n: T.int32): A[0] = n < 10 after = _apply_simplify(before, apply_constraints_to_boolean_branches=True) @@ -996,11 +996,11 @@ def test_simplify_lhs_of_boolean_or_using_rhs(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), n: T.int32): + def before(A: T.Tensor(1, "bool"), n: T.int32): A[0] = n < 5 or n < 10 @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), n: T.int32): + def expected(A: T.Tensor(1, "bool"), n: T.int32): A[0] = n < 10 after = _apply_simplify(before, apply_constraints_to_boolean_branches=True) @@ -1017,11 +1017,11 @@ def test_simplify_rhs_of_boolean_and_using_lhs_without_const(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), n: T.int32, m: T.int32): + def before(A: T.Tensor(1, "bool"), n: T.int32, m: T.int32): A[0] = n < m + 5 and n < m + 10 @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), n: T.int32, m: T.int32): + def expected(A: T.Tensor(1, "bool"), n: T.int32, m: T.int32): A[0] = n < m + 5 after = _apply_simplify( @@ -1040,11 +1040,11 @@ def test_simplify_lhs_of_boolean_and_using_rhs_without_const(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), n: T.int32, m: T.int32): + def before(A: T.Tensor(1, "bool"), n: T.int32, m: T.int32): A[0] = n < m + 10 and n < m + 5 @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), n: T.int32, m: T.int32): + def expected(A: T.Tensor(1, "bool"), n: T.int32, m: T.int32): A[0] = n < m + 5 after = _apply_simplify( @@ -1063,11 +1063,11 @@ def test_simplify_rhs_of_boolean_or_using_lhs_without_const(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), n: T.int32, m: T.int32): + def before(A: T.Tensor(1, "bool"), n: T.int32, m: T.int32): A[0] = n < m + 10 or n < m + 5 @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), n: T.int32, m: T.int32): + def expected(A: T.Tensor(1, "bool"), n: T.int32, m: T.int32): A[0] = n < m + 10 after = _apply_simplify( @@ -1086,11 +1086,11 @@ def test_simplify_lhs_of_boolean_or_using_rhs_without_const(): """ @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), n: T.int32, m: T.int32): + def before(A: T.Tensor(1, "bool"), n: T.int32, m: T.int32): A[0] = n < m + 5 or n < m + 10 @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), n: T.int32, m: T.int32): + def expected(A: T.Tensor(1, "bool"), n: T.int32, m: T.int32): A[0] = n < m + 10 after = _apply_simplify( @@ -1103,12 +1103,12 @@ def test_provable_condition_with_offset(): """Use scoped-constraint to prove inequalities""" @T.prim_func(private=True) - def before(A: T.Buffer(1, "bool"), i: T.int32, j: T.int32): + def before(A: T.Tensor(1, "bool"), i: T.int32, j: T.int32): if i < j: A[0] = i < j + 1 @T.prim_func(private=True) - def expected(A: T.Buffer(1, "bool"), i: T.int32, j: T.int32): + def expected(A: T.Tensor(1, "bool"), i: T.int32, j: T.int32): if i < j: A[0] = True @@ -1141,13 +1141,13 @@ def test_most_restrictive_conditional(): for priors, expr_before, expr_after in test_cases: # well formed checker complains of undefined variables in condition @T.prim_func(private=True, check_well_formed=False) - def before_func(A: T.Buffer(1, "bool")): + def before_func(A: T.Tensor(1, "bool")): if priors: A[0] = expr_before # well formed checker complains of undefined variables in condition @T.prim_func(private=True, check_well_formed=False) - def expected_func(A: T.Buffer(1, "bool")): + def expected_func(A: T.Tensor(1, "bool")): if priors: A[0] = expr_after @@ -1161,7 +1161,7 @@ def test_simplify_trivial_let_buffer_var(): @T.prim_func(private=True) def before(A_ptr: T.handle("float32")): A_ptr_redef: T.let[T.handle("float32")] = A_ptr - A = T.decl_buffer(1, "float32", data=A_ptr_redef) + A = T.decl_tensor(1, "float32", data=A_ptr_redef) A[0] = 42.0 expected = before @@ -1176,13 +1176,13 @@ def test_simplify_trivial_let_elem_offset(): @T.prim_func(private=True) def before(A_ptr: T.handle("float32"), A_offset: T.int32): A_offset_redef = A_offset - A = T.decl_buffer(1, "float32", elem_offset=A_offset_redef, data=A_ptr) + A = T.decl_tensor(1, "float32", elem_offset=A_offset_redef, data=A_ptr) A[0] = 42.0 @T.prim_func(private=True) def expected(A_ptr: T.handle("float32"), A_offset: T.int32): A_offset_redef = A_offset - A = T.decl_buffer(1, "float32", elem_offset=A_offset_redef, data=A_ptr) + A = T.decl_tensor(1, "float32", elem_offset=A_offset_redef, data=A_ptr) A[0] = 42.0 after = _apply_simplify(before) @@ -1195,13 +1195,13 @@ def test_simplify_trivial_let_shape(): @T.prim_func(private=True) def before(A_ptr: T.handle("float32"), A_size: T.int32): A_size_redef = A_size - A = T.decl_buffer([A_size_redef], "float32", data=A_ptr) + A = T.decl_tensor([A_size_redef], "float32", data=A_ptr) A[0] = 42.0 @T.prim_func(private=True) def expected(A_ptr: T.handle("float32"), A_size: T.int32): A_size_redef = A_size - A = T.decl_buffer([A_size_redef], "float32", data=A_ptr) + A = T.decl_tensor([A_size_redef], "float32", data=A_ptr) A[0] = 42.0 after = _apply_simplify(before) @@ -1214,13 +1214,13 @@ def test_simplify_trivial_let_stride(): @T.prim_func(private=True) def before(A_ptr: T.handle("float32"), A_stride: T.int32): A_stride_redef = A_stride - A = T.decl_buffer(1, "float32", strides=[A_stride_redef], data=A_ptr) + A = T.decl_tensor(1, "float32", strides=[A_stride_redef], data=A_ptr) A[0] = 42.0 @T.prim_func(private=True) def expected(A_ptr: T.handle("float32"), A_stride: T.int32): A_stride_redef = A_stride - A = T.decl_buffer(1, "float32", strides=[A_stride_redef], data=A_ptr) + A = T.decl_tensor(1, "float32", strides=[A_stride_redef], data=A_ptr) A[0] = 42.0 after = _apply_simplify(before) @@ -1228,20 +1228,20 @@ def expected(A_ptr: T.handle("float32"), A_stride: T.int32): def test_simplify_buffer_identity_well_formed(): - """Regression: Simplify must not diverge buffer identity between DeclBuffer and BufferLoad. + """Regression: Simplify must not diverge buffer identity between DeclTensor and BufferLoad. The simplifier's Dispatch calls analyzer_->Simplify() directly, bypassing - normal ExprMutator dispatch. If VisitBufferDef remaps a buffer at a DeclBuffer + normal ExprMutator dispatch. If VisitBufferDef remaps a buffer at a DeclTensor site (e.g. inlining n_val -> n in the shape), BufferLoad inside a BufferStore value would NOT pick up the remap because VisitBufferUse is never called. - This causes DeclBuffer/BufferLoad buffer identity divergence. + This causes DeclTensor/BufferLoad buffer identity divergence. """ @T.prim_func(private=True) def before(A_ptr: T.handle("float32"), B_ptr: T.handle("float32"), n: T.int32): n_val = n - A = T.decl_buffer([n_val], "float32", data=A_ptr) - B = T.decl_buffer([n_val], "float32", data=B_ptr) + A = T.decl_tensor([n_val], "float32", data=A_ptr) + B = T.decl_tensor([n_val], "float32", data=B_ptr) B[0] = A[0] after = _apply_simplify(before) @@ -1254,7 +1254,7 @@ def test_buffer_shape_constraint(): @I.ir_module(check_well_formed=False) class Before: @T.prim_func - def main(A: T.Buffer((n * 32,), "float32")): + def main(A: T.Tensor((n * 32,), "float32")): A[T.min(T.int64(0), n)] = T.float32(0) n = T.dynamic("n") @@ -1262,7 +1262,7 @@ def main(A: T.Buffer((n * 32,), "float32")): @I.ir_module(check_well_formed=False) class Expected: @T.prim_func - def main(A: T.Buffer((n * 32,), "float32")): + def main(A: T.Tensor((n * 32,), "float32")): A[T.int64(0)] = T.float32(0) after = tvm.tirx.transform.StmtSimplify()(Before) @@ -1275,7 +1275,7 @@ def test_buffer_shape_constraint_with_offset(): @I.ir_module(check_well_formed=False) class Before: @T.prim_func - def main(A: T.Buffer((n * 32 + 1 - 2,), "float32")): + def main(A: T.Tensor((n * 32 + 1 - 2,), "float32")): A[T.min(T.int64(1), n)] = T.float32(0) n = T.dynamic("n") @@ -1283,7 +1283,7 @@ def main(A: T.Buffer((n * 32 + 1 - 2,), "float32")): @I.ir_module(check_well_formed=False) class Expected: @T.prim_func - def main(A: T.Buffer((n * 32 + 1 - 2,), "float32")): + def main(A: T.Tensor((n * 32 + 1 - 2,), "float32")): A[T.int64(1)] = T.float32(0) after = tvm.tirx.transform.StmtSimplify()(Before) @@ -1292,14 +1292,14 @@ def main(A: T.Buffer((n * 32 + 1 - 2,), "float32")): def test_nested_if_elimination(): @T.prim_func(private=True) - def before(a: T.Buffer((2, 8), "int32"), b: T.Buffer((2, 8), "int32")): + def before(a: T.Tensor((2, 8), "int32"), b: T.Tensor((2, 8), "int32")): for i0, j0 in T.grid(2, 8): b[i0, j0] = T.if_then_else( i0 == 1 and 6 <= j0, 0, T.max(0, T.if_then_else(i0 == 1 and 6 <= j0, 0, a[i0, j0])) ) @T.prim_func(private=True) - def expected(a: T.Buffer((2, 8), "int32"), b: T.Buffer((2, 8), "int32")): + def expected(a: T.Tensor((2, 8), "int32"), b: T.Tensor((2, 8), "int32")): for i0, j0 in T.grid(2, 8): b[i0, j0] = T.if_then_else(i0 == 1 and 6 <= j0, 0, T.max(0, a[i0, j0])) @@ -1314,8 +1314,8 @@ def test_mutable_branch_predicate_preserves_while_bound(else_branch, write_befor # and a store on its back edge both invalidate the entry predicate. from tvm import tirx - x = tirx.decl_buffer((1,), "int32", name="x") - count = tirx.decl_buffer((1,), "int32", name="count") + x = tirx.decl_tensor((1,), "int32", name="x") + count = tirx.decl_tensor((1,), "int32", name="count") loop = tirx.While( T.And(x[0] < 8, count[0] == 0), tirx.BufferStore(x, x[0] + 1, [0]), @@ -1333,7 +1333,7 @@ def test_mutable_branch_predicate_preserves_while_bound(else_branch, write_befor def test_mutable_branch_predicate_preserves_later_load(): @T.prim_func(private=True) - def before(x: T.Buffer((1,), "int32"), out: T.Buffer((1,), "int32")): + def before(x: T.Tensor((1,), "int32"), out: T.Tensor((1,), "int32")): if x[0] < 8: x[0] = x[0] + 1 out[0] = T.Select(x[0] < 8, 1, 0) @@ -1343,7 +1343,7 @@ def before(x: T.Buffer((1,), "int32"), out: T.Buffer((1,), "int32")): def test_mutable_assert_does_not_constrain_later_load(): @T.prim_func(private=True) - def before(x: T.Buffer((1,), "int32"), out: T.Buffer((1,), "int32")): + def before(x: T.Tensor((1,), "int32"), out: T.Tensor((1,), "int32")): assert x[0] < 8, "initial bound" x[0] = x[0] + 1 out[0] = T.Select(x[0] < 8, 1, 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 13b0de42d0d1..1c0d899b9fb7 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 @@ -288,8 +288,8 @@ def test_dynamic_launch_thread(): class before: @T.prim_func def default_function( - A: T.Buffer([seq_len], "int32"), # noqa: F821 - B: T.Buffer([seq_len], "int32"), # noqa: F821 + A: T.Tensor([seq_len], "int32"), # noqa: F821 + B: T.Tensor([seq_len], "int32"), # noqa: F821 seq_len: T.int32, ): T.func_attr({"target": T.target("cuda")}) @@ -305,8 +305,8 @@ def default_function( class expected: @T.prim_func def default_function( - A: T.Buffer((seq_len,), "int32"), # noqa: F821 - B: T.Buffer((seq_len,), "int32"), # noqa: F821 + A: T.Tensor((seq_len,), "int32"), # noqa: F821 + B: T.Tensor((seq_len,), "int32"), # noqa: F821 seq_len: T.int32, ): T.func_attr({"target": T.target("cuda")}) @@ -328,8 +328,8 @@ def default_function_kernel( "tirx.noalias": True, } ) - A = T.decl_buffer(seq_len, "int32", data=A_data) - B = T.decl_buffer(seq_len, "int32", data=B_data) + A = T.decl_tensor(seq_len, "int32", data=A_data) + B = T.decl_tensor(seq_len, "int32", data=B_data) blockIdx_x = T.launch_thread("blockIdx.x", num_blocks) threadIdx_x = T.launch_thread("threadIdx.x", 128) if blockIdx_x * 128 + threadIdx_x < seq_len: @@ -347,13 +347,13 @@ def test_symbolic_var_parameter(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((m,)), B: T.Buffer((m,))): + def main(A: T.Tensor((m,)), B: T.Tensor((m,))): T.func_attr({"target": T.target("cuda")}) T.attr(T.target("cuda"), "target", 0) blockIdx_x = T.launch_thread("blockIdx.x", m) - B_1 = T.decl_buffer((m,), data=B.data) - A_1 = T.decl_buffer((m,), data=A.data) + B_1 = T.decl_tensor((m,), data=B.data) + A_1 = T.decl_tensor((m,), data=A.data) B_1[blockIdx_x] = A_1[blockIdx_x] after = tvm.tirx.transform.SplitHostDevice()(Module) @@ -365,7 +365,7 @@ def test_buffer_used_only_through_data_projection(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((16,), "float32")): + def main(A: T.Tensor((16,), "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) with T.attr(T.target("cuda"), "target", 0): T.evaluate(T.call_extern("consume", A.data, dtype="int32")) @@ -375,7 +375,7 @@ def main(A: T.Buffer((16,), "float32")): declared_buffers = [] def collect(node): - if _is_buffer_binding(node, "tirx.decl_buffer"): + if _is_buffer_binding(node, "tirx.decl_tensor"): declared_buffers.append(node.var) tvm_ffi.structural_walk(kernel.body, collect) @@ -389,7 +389,7 @@ def test_thread_extent_region_extracted_as_device_kernel(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(16, "float32")): + def main(A: T.Tensor(16, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) i = T.launch_thread("threadIdx.x", 16) A[i] = 0.0 @@ -397,7 +397,7 @@ def main(A: T.Buffer(16, "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(16, "float32")): + def main(A: T.Tensor(16, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) T.call_ffi_kernel("main_kernel", A.data, 16, launch_params=["threadIdx.x"]) @@ -413,7 +413,7 @@ def main_kernel(A_data: T.handle("float32")): "tirx.is_global_func": True, } ) - A = T.decl_buffer(16, dtype="float32", data=A_data) + A = T.decl_tensor(16, dtype="float32", data=A_data) i = T.launch_thread("threadIdx.x", 16) A[i] = 0.0 @@ -425,7 +425,7 @@ def test_cuda_launch_preserves_flag_metadata(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(16, "float32")): + def main(A: T.Tensor(16, "float32")): T.func_attr( { "target": T.target("cuda", host="llvm"), @@ -467,7 +467,7 @@ def test_cuda_required_block_size_coexists_with_launch_bounds(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(4, "float32")): + def main(A: T.Tensor(4, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) T.attr(T.target("cuda"), "target", 0) T.attr(0, "tirx.required_block_size", 1) @@ -499,7 +499,7 @@ def test_cuda_launch_preserves_singleton_cluster_dimensions(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) with T.attr(T.target("cuda"), "target", 0): T.launch_thread("blockIdx.x", 4) @@ -530,7 +530,7 @@ def test_device_scope_region_extracted_as_device_kernel(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) T.attr(0, "device_scope", 0) A[0] = 0.0 @@ -538,7 +538,7 @@ def main(A: T.Buffer(1, "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("cuda", host="llvm")}) T.call_ffi_kernel("main_kernel", A.data, launch_params=[]) @@ -554,7 +554,7 @@ def main_kernel(A_data: T.handle("float32")): "tirx.is_global_func": True, } ) - A = T.decl_buffer(1, dtype="float32", data=A_data) + A = T.decl_tensor(1, dtype="float32", data=A_data) T.attr(0, "device_scope", 0) A[0] = 0.0 @@ -568,20 +568,20 @@ def test_lower_device_kernel_launch(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("llvm")}) Before.kernel(A.data) @T.prim_func def kernel(A_data: T.handle("float32")): T.func_attr({"target": T.target("cuda")}) - A = T.decl_buffer(1, dtype="float32", data=A_data) + A = T.decl_tensor(1, dtype="float32", data=A_data) A[0] = 0.0 @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("llvm")}) T.call_ffi_kernel("kernel", A.data, launch_params=[]) @@ -596,7 +596,7 @@ def kernel(A_data: T.handle("float32")): "tirx.is_global_func": True, } ) - A = T.decl_buffer(1, dtype="float32", data=A_data) + A = T.decl_tensor(1, dtype="float32", data=A_data) A[0] = 0.0 After = tvm.tirx.transform.SplitHostDevice()(Before) @@ -609,20 +609,20 @@ def test_externally_visible_kernel_launch(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("llvm")}) Before.kernel(A.data) @T.prim_func def kernel(A_data: T.handle("float32")): T.func_attr({"target": T.target("cuda"), "global_symbol": "kernel_by_another_name"}) - A = T.decl_buffer(1, dtype="float32", data=A_data) + A = T.decl_tensor(1, dtype="float32", data=A_data) A[0] = 0.0 @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("llvm")}) T.call_ffi_kernel("kernel_by_another_name", A.data, launch_params=[]) @@ -637,7 +637,7 @@ def kernel(A_data: T.handle("float32")): "tirx.is_global_func": True, } ) - A = T.decl_buffer(1, dtype="float32", data=A_data) + A = T.decl_tensor(1, dtype="float32", data=A_data) A[0] = 0.0 After = tvm.tirx.transform.SplitHostDevice()(Before) @@ -650,7 +650,7 @@ def test_collect_launch_parameter(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(16, "float32")): + def main(A: T.Tensor(16, "float32")): T.func_attr({"target": T.target("llvm")}) Before.kernel(A.data) @@ -662,14 +662,14 @@ def kernel(A_data: T.handle("float32")): "global_symbol": "kernel", } ) - A = T.decl_buffer(16, dtype="float32", data=A_data) + A = T.decl_tensor(16, dtype="float32", data=A_data) i = T.launch_thread("threadIdx.x", 16) A[i] = 0.0 @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(16, "float32")): + def main(A: T.Tensor(16, "float32")): T.func_attr({"target": T.target("llvm")}) T.call_ffi_kernel("kernel", A.data, 16, launch_params=["threadIdx.x"]) @@ -684,7 +684,7 @@ def kernel(A_data: T.handle("float32")): "tirx.is_global_func": True, } ) - A = T.decl_buffer(16, dtype="float32", data=A_data) + A = T.decl_tensor(16, dtype="float32", data=A_data) i = T.launch_thread("threadIdx.x", 16) A[i] = 0.0 @@ -698,20 +698,20 @@ def test_same_device_different_target(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("llvm")}) Before.kernel(A.data) @T.prim_func def kernel(A_data: T.handle("float32")): T.func_attr({"target": T.target("c")}) - A = T.decl_buffer(16, dtype="float32", data=A_data) + A = T.decl_tensor(16, dtype="float32", data=A_data) A[0] = 0.0 @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(1, "float32")): + def main(A: T.Tensor(1, "float32")): T.func_attr({"target": T.target("llvm")}) T.call_extern("kernel", A.data, dtype="void") @@ -724,7 +724,7 @@ def kernel(A_data: T.handle("float32")): "tirx.is_global_func": True, } ) - A = T.decl_buffer(16, dtype="float32", data=A_data) + A = T.decl_tensor(16, dtype="float32", data=A_data) A[0] = 0.0 After = tvm.tirx.transform.SplitHostDevice()(Before) @@ -737,14 +737,14 @@ def test_bind_before_thread_extent(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(16, "float32"), n: T.int32): + def main(A: T.Tensor(16, "float32"), n: T.int32): T.func_attr({"target": T.target("llvm")}) Before.kernel(A.data, n) @T.prim_func def kernel(A_data: T.handle("float32"), n: T.int32): T.func_attr({"target": T.target("cuda"), "global_symbol": "kernel"}) - A = T.decl_buffer(16, dtype="float32", data=A_data) + A = T.decl_tensor(16, dtype="float32", data=A_data) v: T.let[T.int32] = n + 1 i = T.launch_thread("threadIdx.x", v) A[i] = 0.0 @@ -752,7 +752,7 @@ def kernel(A_data: T.handle("float32"), n: T.int32): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(16, "float32"), n: T.int32): + def main(A: T.Tensor(16, "float32"), n: T.int32): T.func_attr({"target": T.target("llvm")}) T.call_ffi_kernel("kernel", A.data, n, n + 1, launch_params=["threadIdx.x"]) @@ -767,7 +767,7 @@ def kernel(A_data: T.handle("float32"), n: T.int32): "tirx.is_global_func": True, } ) - A = T.decl_buffer(16, dtype="float32", data=A_data) + A = T.decl_tensor(16, dtype="float32", data=A_data) v: T.let[T.int32] = n + 1 i = T.launch_thread("threadIdx.x", v) A[i] = 0.0 diff --git a/tests/python/tirx-transform/test_tir_transform_storage_rewrite.py b/tests/python/tirx-transform/test_tir_transform_storage_rewrite.py index d3ad057317c2..93371ccc4f7d 100644 --- a/tests/python/tirx-transform/test_tir_transform_storage_rewrite.py +++ b/tests/python/tirx-transform/test_tir_transform_storage_rewrite.py @@ -42,10 +42,10 @@ def test_alloc_seq(): def func(n: T.int32): for i in T.serial(n): for j in range(10): - A = T.alloc_buffer((200,), scope=scope_tb) + A = T.alloc_tensor((200,), scope=scope_tb) A[j] = T.float32(1.2) for j in range(10): - B = T.alloc_buffer((200,), scope=scope_tb) + B = T.alloc_tensor((200,), scope=scope_tb) B[j] = T.float32(1.3) mod = tvm.IRModule.from_expr(func) @@ -54,7 +54,7 @@ def func(n: T.int32): num_alloc = [0] def verify(n): - if _is_buffer_binding(n, "tirx.alloc_buffer"): + if _is_buffer_binding(n, "tirx.alloc_tensor"): num_alloc[0] += 1 assert n.var.ty.shape[0].value == 200 @@ -70,11 +70,11 @@ def make_mod(dtype_list, length): @T.prim_func def func(): # Allocate all buffers in parent scope (before any loops) - A = T.alloc_buffer((length,), dtype_list[0], scope="local.L0A") - B = T.alloc_buffer((length,), dtype_list[1], scope="local.L0A") - C = T.alloc_buffer((length,), dtype_list[2], scope="local.L0A") - D = T.alloc_buffer((length,), dtype_list[3], scope="local.L0A") - E = T.alloc_buffer((length,), "int8", scope="local.L0A") + A = T.alloc_tensor((length,), dtype_list[0], scope="local.L0A") + B = T.alloc_tensor((length,), dtype_list[1], scope="local.L0A") + C = T.alloc_tensor((length,), dtype_list[2], scope="local.L0A") + D = T.alloc_tensor((length,), dtype_list[3], scope="local.L0A") + E = T.alloc_tensor((length,), "int8", scope="local.L0A") for j in range(length): A[j] = T.Cast(dtype_list[0], 1) @@ -109,7 +109,7 @@ def offset_generater(dtype_list, length): def dtype_test(dtype_list, length): def verify(n): - if _is_buffer_binding(n, "tirx.alloc_buffer"): + if _is_buffer_binding(n, "tirx.alloc_tensor"): assert n.var.ty.shape[0].value == offset mod = make_mod(dtype_list, length) @@ -136,22 +136,22 @@ def test_address_of(): # In this test, the storage rewrite pass is allowed to # combine buffers B and D, but not C @T.prim_func - def before(A: T.Buffer(8, "float32"), E: T.Buffer(8, "float32")): - B = T.alloc_buffer((8,)) + def before(A: T.Tensor(8, "float32"), E: T.Tensor(8, "float32")): + B = T.alloc_tensor((8,)) for i in range(8): B[i] = ( T.call_extern("deref", T.address_of(A[i]), dtype="float32") + T.call_extern("deref", T.address_of(A[0]), dtype="float32") + T.float32(1) ) - C = T.alloc_buffer((8,)) + C = T.alloc_tensor((8,)) for i in range(8): C[i] = ( T.call_extern("deref", T.address_of(B[i]), dtype="float32") + T.call_extern("deref", T.address_of(B[0]), dtype="float32") + T.float32(2) ) - D = T.alloc_buffer((8,)) + D = T.alloc_tensor((8,)) for i in range(8): D[i] = ( T.call_extern("deref", T.address_of(C[i]), dtype="float32") @@ -166,7 +166,7 @@ def before(A: T.Buffer(8, "float32"), E: T.Buffer(8, "float32")): ) def verify(n): - if _is_buffer_binding(n, "tirx.alloc_buffer"): + if _is_buffer_binding(n, "tirx.alloc_tensor"): total_alloc[0] += n.var.ty.shape[0].value total_alloc = [0] @@ -185,14 +185,14 @@ def test_parallel_alloc(): def func1(n: T.int32): for i in T.parallel(n): for j in range(10): - A = T.alloc_buffer((n,)) + A = T.alloc_tensor((n,)) A[j] = A[j] + T.float32(2) mod = tvm.IRModule.from_expr(func1) body = tvm.tirx.transform.StorageRewrite()(mod)["func1"] - # With flat AllocBuffer, the for body is a SeqStmt; first element is AllocBuffer - assert _is_buffer_binding(body.body.body[0], "tirx.alloc_buffer") + # With flat AllocTensor, the for body is a SeqStmt; first element is AllocTensor + assert _is_buffer_binding(body.body.body[0], "tirx.alloc_tensor") @T.prim_func def func2(n: T.int32): @@ -200,33 +200,33 @@ def func2(n: T.int32): with T.attr(T.int32(1), "pragma_scope", "parallel_launch_point"): for i in T.parallel(n): for j in range(10): - A = T.alloc_buffer((n,)) + A = T.alloc_tensor((n,)) A[j] = A[j] + T.float32(2) mod = tvm.IRModule.from_expr(func2) body = tvm.tirx.transform.StorageRewrite()(mod)["func2"] - assert _is_buffer_binding(body.body.body.body.body[0], "tirx.alloc_buffer") + assert _is_buffer_binding(body.body.body.body.body[0], "tirx.alloc_tensor") def test_while_alloc(): @T.prim_func def func_parallel(n: T.int32): for i in T.parallel(n): - j = T.alloc_buffer((1,), "int32") + j = T.alloc_tensor((1,), "int32") j[0] = 0 while j[0] < 10: - A = T.alloc_buffer((n,)) + A = T.alloc_tensor((n,)) A[j[0]] = A[j[0]] + T.float32(2) j[0] = j[0] + j[0] + 1 @T.prim_func def func_serial(n: T.int32): for i in T.serial(n): - j = T.alloc_buffer((1,), "int32") + j = T.alloc_tensor((1,), "int32") j[0] = 0 while j[0] < 10: - A = T.alloc_buffer((n,)) + A = T.alloc_tensor((n,)) A[j[0]] = A[j[0]] + T.float32(2) j[0] = j[0] + j[0] + 1 @@ -243,15 +243,15 @@ def func_serial(n: T.int32): # } body = tvm.tirx.transform.StorageRewrite()(mod)["func_parallel"] # Navigate to inside the for loop, then check that allocations exist - # The structure with DeclBuffer is: - # parallel (i, 0, n) { DeclBuffer(j, DeclBuffer(A, ...)) } - # or with Allocate+DeclBuffer pairs + # The structure with DeclTensor is: + # parallel (i, 0, n) { DeclTensor(j, DeclTensor(A, ...)) } + # or with Allocate+DeclTensor pairs inner = body.body.body # inside For - # Skip DeclBuffer nodes to find Allocate + # Skip DeclTensor nodes to find Allocate num_alloc = [0] def count_alloc(n): - if _is_buffer_binding(n, "tirx.alloc_buffer"): + if _is_buffer_binding(n, "tirx.alloc_tensor"): num_alloc[0] += 1 tvm_ffi.structural_walk(inner, count_alloc) @@ -269,17 +269,17 @@ def test_alloc_seq_type(): def func(n: T.int32): for i in T.serial(n): for j in range(10): - A = T.alloc_buffer((200,), scope="local.L0A") - A1 = T.alloc_buffer((200,), scope="local.L0A") + A = T.alloc_tensor((200,), scope="local.L0A") + A1 = T.alloc_tensor((200,), scope="local.L0A") A[j] = T.float32(1.2) A1[j] = T.float32(1.3) - B = T.alloc_buffer((200,), "int16", scope="local.L0A") + B = T.alloc_tensor((200,), "int16", scope="local.L0A") B[j] = T.int16(1) - C = T.alloc_buffer((200,), "int16", scope="local.L0A") + C = T.alloc_tensor((200,), "int16", scope="local.L0A") C[j] = T.int16(1) - D = T.alloc_buffer((200,), "int16", scope="local.L0A") + D = T.alloc_tensor((200,), "int16", scope="local.L0A") D[j] = B[j] + C[j] - A2 = T.alloc_buffer((200,), scope="local.L0A") + A2 = T.alloc_tensor((200,), scope="local.L0A") A2[j] = A[j] mod = tvm.IRModule.from_expr(func) @@ -288,7 +288,7 @@ def func(n: T.int32): num_alloc = [0] def verify(n): - if _is_buffer_binding(n, "tirx.alloc_buffer"): + if _is_buffer_binding(n, "tirx.alloc_tensor"): num_alloc[0] += 1 assert n.var.ty.shape[0].value == 500 @@ -303,13 +303,13 @@ def test_alloc_seq_type2(): def func(n: T.int32): for i in T.serial(n): for j in range(10): - A = T.alloc_buffer((200,), scope=scope_tb) + A = T.alloc_tensor((200,), scope=scope_tb) A[j] = T.float32(1.2) for j in range(20): - B = T.alloc_buffer((400,), "int16", scope=scope_tb) + B = T.alloc_tensor((400,), "int16", scope=scope_tb) B[j] = T.int16(1) for j in range(10): - C = T.alloc_buffer((200,), scope=scope_tb) + C = T.alloc_tensor((200,), scope=scope_tb) C[j] = T.float32(1.2) mod = tvm.IRModule.from_expr(func) @@ -318,7 +318,7 @@ def func(n: T.int32): num_alloc = [0] def verify(n): - if _is_buffer_binding(n, "tirx.alloc_buffer"): + if _is_buffer_binding(n, "tirx.alloc_tensor"): num_alloc[0] += 1 assert n.var.ty.shape[0].value == 200 @@ -331,17 +331,17 @@ def test_reuse_small_buffer(): def func(n: T.int32): for i in T.serial(n): for j in range(10): - A = T.alloc_buffer((200,), "int16", scope="local.L0A") + A = T.alloc_tensor((200,), "int16", scope="local.L0A") A[j] = T.int16(1) - B = T.alloc_buffer((200,), "int16", scope="local.L0A") + B = T.alloc_tensor((200,), "int16", scope="local.L0A") B[j] = T.int16(1) - B1 = T.alloc_buffer((200,), "int16", scope="local.L0A") + B1 = T.alloc_tensor((200,), "int16", scope="local.L0A") B1[j] = A[j] + B[j] - C = T.alloc_buffer((400,), "int16", scope="local.L0A") + C = T.alloc_tensor((400,), "int16", scope="local.L0A") C[j] = T.int16(1) - D = T.alloc_buffer((400,), "int16", scope="local.L0A") + D = T.alloc_tensor((400,), "int16", scope="local.L0A") D[j] = T.int16(1) - E = T.alloc_buffer((400,), "int16", scope="local.L0A") + E = T.alloc_tensor((400,), "int16", scope="local.L0A") E[j] = C[j] mod = tvm.IRModule.from_expr(func) @@ -350,7 +350,7 @@ def func(n: T.int32): num_alloc = [0] def verify(n): - if _is_buffer_binding(n, "tirx.alloc_buffer"): + if _is_buffer_binding(n, "tirx.alloc_tensor"): num_alloc[0] += 1 assert n.var.ty.shape[0].value == 800 @@ -360,16 +360,16 @@ def verify(n): def test_access_in_let_value(): @T.prim_func - def func(A: T.Buffer((8,), "float32")): + def func(A: T.Tensor((8,), "float32")): for i in range(8): - B = T.alloc_buffer((1,)) + B = T.alloc_tensor((1,)) B[0] = 3.14 x: T.let[T.float32] = T.exp(B[0]) A[i] = (x + 1.0) / (x - 1.0) @T.prim_func - def func_rewritten(A: T.Buffer((8,), "float32")) -> None: - B = T.alloc_buffer((1,)) + def func_rewritten(A: T.Tensor((8,), "float32")) -> None: + B = T.alloc_tensor((1,)) for i in range(8): B[0] = 3.14 x: T.let[T.float32] = T.exp(B[0]) @@ -382,9 +382,9 @@ def func_rewritten(A: T.Buffer((8,), "float32")) -> None: def test_decl_buffer_is_not_vectorized(): - """StorageRewrite leaves explicit DeclBuffer views unchanged. + """StorageRewrite leaves explicit DeclTensor views unchanged. - Vectorization of DeclBuffer views was dropped because the rewritten result + Vectorization of DeclTensor views was dropped because the rewritten result violates the immutable BufferVar binding invariants. """ @@ -395,7 +395,7 @@ def main() -> None: A_data: T.let[T.handle("int32")] = T.call_extern( "dummy_func", dtype=T.handle("int32").ty ) - A = T.decl_buffer([8], "int32", data=A_data) + A = T.decl_tensor([8], "int32", data=A_data) A[T.ramp(0, 1, 8)] = T.broadcast(42, 8) After = tvm.tirx.transform.StorageRewrite()(Before) @@ -403,14 +403,14 @@ def main() -> None: def test_rewrite_decl_buffer(): - """A DeclBuffer node may appear in StorageRewrite's input""" + """A DeclTensor node may appear in StorageRewrite's input""" @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(16, "float32"), D: T.Buffer(16, "float32")): - B = T.decl_buffer(16, dtype="float32") - C = T.decl_buffer(16, dtype="float32") + def main(A: T.Tensor(16, "float32"), D: T.Tensor(16, "float32")): + B = T.decl_tensor(16, dtype="float32") + C = T.decl_tensor(16, dtype="float32") for i in range(16): B[i] = A[i] @@ -424,9 +424,9 @@ def main(A: T.Buffer(16, "float32"), D: T.Buffer(16, "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(16, "float32"), D: T.Buffer(16, "float32")): - B = T.decl_buffer(16, dtype="float32") - C = T.decl_buffer(16, dtype="float32", data=B.data) + def main(A: T.Tensor(16, "float32"), D: T.Tensor(16, "float32")): + B = T.decl_tensor(16, dtype="float32") + C = T.decl_tensor(16, dtype="float32", data=B.data) for i in range(16): B[i] = A[i] @@ -448,20 +448,20 @@ def test_decl_buffer_alias_chain_uses_flat_root(): @I.ir_module class Before: @T.prim_func - def main(D: T.Buffer(1, "float32")): - A = T.decl_buffer(16, dtype="float32") - B = T.decl_buffer(16, dtype="float32", data=A.data) - C = T.decl_buffer(16, dtype="float32", data=B.data) + def main(D: T.Tensor(1, "float32")): + A = T.decl_tensor(16, dtype="float32") + B = T.decl_tensor(16, dtype="float32", data=A.data) + C = T.decl_tensor(16, dtype="float32", data=B.data) A[0] = 1.0 D[0] = C[0] @I.ir_module class Expected: @T.prim_func - def main(D: T.Buffer(1, "float32")): - A = T.decl_buffer(16, dtype="float32") - B = T.decl_buffer(16, dtype="float32", data=A.data) - C = T.decl_buffer(16, dtype="float32", data=A.data) + def main(D: T.Tensor(1, "float32")): + A = T.decl_tensor(16, dtype="float32") + B = T.decl_tensor(16, dtype="float32", data=A.data) + C = T.decl_tensor(16, dtype="float32", data=A.data) A[0] = 1.0 D[0] = C[0] @@ -474,12 +474,12 @@ def test_decl_buffer_alias_extends_source_lifetime(): """An access through an alias prevents reuse of its source allocation.""" @T.prim_func - def func(D: T.Buffer(1, "float32")): - A = T.decl_buffer(16, dtype="float32") - B = T.decl_buffer(16, dtype="float32", data=A.data) + def func(D: T.Tensor(1, "float32")): + A = T.decl_tensor(16, dtype="float32") + B = T.decl_tensor(16, dtype="float32", data=A.data) A[0] = 1.0 - C = T.decl_buffer(16, dtype="float32") + C = T.decl_tensor(16, dtype="float32") C[0] = 2.0 D[0] = B[0] + C[0] @@ -488,27 +488,27 @@ def func(D: T.Buffer(1, "float32")): tvm_ffi.structural_walk( after.body, lambda node: allocations.append(node) - if _is_buffer_binding(node, "tirx.alloc_buffer") + if _is_buffer_binding(node, "tirx.alloc_tensor") else None, ) assert len(allocations) == 2 def test_no_orphaned_decl_buffer(): - """A DeclBuffer of an unused Allocate should be removed + """A DeclTensor of an unused Allocate should be removed StorageRewrite removes any allocations that are unused. When it - does so, any DeclBuffer that refers to that allocation should also + does so, any DeclTensor that refers to that allocation should also be removed. """ @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer(16, "float32"), D: T.Buffer(16, "float32")): - B = T.decl_buffer(16, dtype="float32") - C = T.decl_buffer(16, dtype="float32") - Unused = T.decl_buffer(16, dtype="float32") + def main(A: T.Tensor(16, "float32"), D: T.Tensor(16, "float32")): + B = T.decl_tensor(16, dtype="float32") + C = T.decl_tensor(16, dtype="float32") + Unused = T.decl_tensor(16, dtype="float32") for i in range(16): B[i] = A[i] @@ -522,9 +522,9 @@ def main(A: T.Buffer(16, "float32"), D: T.Buffer(16, "float32")): @I.ir_module class Expected: @T.prim_func - def main(A: T.Buffer(16, "float32"), D: T.Buffer(16, "float32")): - B = T.decl_buffer(16, dtype="float32") - C = T.decl_buffer(16, dtype="float32", data=B.data) + def main(A: T.Tensor(16, "float32"), D: T.Tensor(16, "float32")): + B = T.decl_tensor(16, dtype="float32") + C = T.decl_tensor(16, dtype="float32", data=B.data) for i in range(16): B[i] = A[i] diff --git a/tests/python/tirx-transform/test_tir_transform_unroll_loop.py b/tests/python/tirx-transform/test_tir_transform_unroll_loop.py index 8986388b655f..8fd535ac6eb1 100644 --- a/tests/python/tirx-transform/test_tir_transform_unroll_loop.py +++ b/tests/python/tirx-transform/test_tir_transform_unroll_loop.py @@ -25,7 +25,7 @@ def test_unroll_loop(): @I.ir_module class Module: @T.prim_func - def main(Ab: T.Buffer((n,), "int64"), n: T.int64): # noqa: F821 + def main(Ab: T.Tensor((n,), "int64"), n: T.int64): # noqa: F821 for i in T.serial(n, n + 2): for j in T.unroll(8): Ab[j + 1] = Ab[i] + T.int64(1) @@ -53,7 +53,7 @@ def main(Ab: T.Buffer((n,), "int64"), n: T.int64): # noqa: F821 @I.ir_module class ModuleWithPragma: @T.prim_func - def main(Ab: T.Buffer((n,), "int64"), n: T.int64): # noqa: F821 + def main(Ab: T.Tensor((n,), "int64"), n: T.int64): # noqa: F821 with T.attr(T.int32(0), "pragma_auto_unroll_max_step", 16): for i in T.serial(n, n + 2): for j in T.unroll(8): @@ -76,7 +76,7 @@ def test_unroll_fake_loop(): @I.ir_module class Module: @T.prim_func - def main(Ab: T.Buffer((n,), "int32"), n: T.int64): # noqa: F821 + def main(Ab: T.Tensor((n,), "int32"), n: T.int64): # noqa: F821 for i in T.serial(1): Ab[i * 2] = 3 for j in T.serial(10): @@ -97,16 +97,16 @@ class Before: @T.prim_func def main(): for i in T.unroll(2): - buf = T.alloc_buffer([16], "float32") + buf = T.alloc_tensor([16], "float32") buf[0] = 0.0 @I.ir_module class Expected: @T.prim_func def main(): - buf1 = T.alloc_buffer([16], "float32") + buf1 = T.alloc_tensor([16], "float32") buf1[0] = 0.0 - buf2 = T.alloc_buffer([16], "float32") + buf2 = T.alloc_tensor([16], "float32") buf2[0] = 0.0 after = tvm.tirx.transform.UnrollLoop()(Before) @@ -118,20 +118,20 @@ def test_unroll_local_access(): @I.ir_module class Before: @T.prim_func - def main(B: T.Buffer((64,), "float32")): + def main(B: T.Tensor((64,), "float32")): for bx in T.thread_binding(4, thread="blockIdx.x"): for tx in T.thread_binding(4, thread="threadIdx.x"): - A_local = T.alloc_buffer((4,), scope="local") + A_local = T.alloc_tensor((4,), scope="local") for i in T.serial(4): A_local[i] = T.float32(i) @I.ir_module class Expected: @T.prim_func - def main(B: T.Buffer((64,), "float32")): + def main(B: T.Tensor((64,), "float32")): for bx in T.thread_binding(4, thread="blockIdx.x"): for tx in T.thread_binding(4, thread="threadIdx.x"): - A_local = T.alloc_buffer((4,), scope="local") + A_local = T.alloc_tensor((4,), scope="local") A_local[0] = T.float32(0) A_local[1] = T.float32(1) A_local[2] = T.float32(2) diff --git a/tests/python/tirx-transform/test_tir_transform_vectorize.py b/tests/python/tirx-transform/test_tir_transform_vectorize.py index 036865d0fc60..e675d4890f72 100644 --- a/tests/python/tirx-transform/test_tir_transform_vectorize.py +++ b/tests/python/tirx-transform/test_tir_transform_vectorize.py @@ -40,14 +40,14 @@ def test_vectorize_loop(extent, target): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((16,), "float32")): + def main(A: T.Tensor((16,), "float32")): for j in T.vectorized(0, extent): A[j] = 1 @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((16,), "float32")): + def main(A: T.Tensor((16,), "float32")): A[T.Ramp(0, 1, extent)] = T.Broadcast(1, extent) with tvm.target.Target(target): @@ -59,7 +59,7 @@ def test_vectorize_vector(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((4,), "float32x4"), n: T.int32): + def main(A: T.Tensor((4,), "float32x4"), n: T.int32): for i in range(n): for j in T.vectorized(4): A[j] = T.Broadcast(T.float32(1), 4) @@ -78,7 +78,7 @@ def test_vectorize_vector_scalable_error(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32")): for j in T.vectorized(T.vscale() * 4): A[T.ramp(j * 4, 1, 4)] = T.Broadcast(T.float32(1), 4) @@ -92,7 +92,7 @@ def test_vectorize_vector_scalable_error2(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((25,), "float32xvscalex4")): + def main(A: T.Tensor((25,), "float32xvscalex4")): for j in T.vectorized(4): A[j] = T.Broadcast(T.float32(1), T.vscale() * 4) @@ -105,7 +105,7 @@ def test_vectorize_vector_scalable_error3(): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32")): for j in T.vectorized(4): A[T.ramp(j * T.vscale() * 4, 1, T.vscale() * 4)] = T.Broadcast( T.float32(1), T.vscale() * 4 @@ -121,7 +121,7 @@ def test_vectorize_vector_scalable_error4(): @I.ir_module class Module: @T.prim_func(private=True) - def main(A: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32")): for j in T.vectorized(T.vscale() * 4): A[T.ramp(j * T.vscale() * 4, 1, T.vscale() * 4)] = T.Broadcast( T.float32(1), T.vscale() * 4 @@ -140,7 +140,7 @@ def test_vectorize_with_if(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((25,), "float32"), n: T.int32, x: T.int32): + def main(A: T.Tensor((25,), "float32"), n: T.int32, x: T.int32): for i in T.vectorized(extent): if x < n: A[i] = A[i] + T.float32(1) @@ -151,7 +151,7 @@ def main(A: T.Buffer((25,), "float32"), n: T.int32, x: T.int32): @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((25,), "float32"), n: T.int32, x: T.int32): + def main(A: T.Tensor((25,), "float32"), n: T.int32, x: T.int32): if x < n: A[T.Ramp(0, 1, extent)] = A[T.Ramp(0, 1, extent)] + T.Broadcast( T.float32(1), extent @@ -173,7 +173,7 @@ def test_vectorize_if_scalable_extent(): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((25,), "float32"), n: T.int32, x: T.int32): + def main(A: T.Tensor((25,), "float32"), n: T.int32, x: T.int32): for i in T.vectorized(extent): if x < n: A[i] = A[i] + T.float32(1) @@ -184,7 +184,7 @@ def main(A: T.Buffer((25,), "float32"), n: T.int32, x: T.int32): @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((25,), "float32"), n: T.int32, x: T.int32): + def main(A: T.Tensor((25,), "float32"), n: T.int32, x: T.int32): if x < n: A[T.Ramp(0, 1, extent)] = A[T.Ramp(0, 1, extent)] + T.Broadcast( T.float32(1), extent @@ -211,7 +211,7 @@ def test_vectorize_let(extent, target): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32")): for i in T.vectorized(extent): v: T.let = A[i] + T.float32(1) A[i] = v + T.float32(2) @@ -219,7 +219,7 @@ def main(A: T.Buffer((25,), "float32")): @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32")): v: T.let = A[T.Ramp(0, 1, extent)] + T.Broadcast(T.float32(1), extent) A[T.Ramp(0, 1, extent)] = v + T.Broadcast(T.float32(2), extent) @@ -233,7 +233,7 @@ def test_vectorize_with_le_cond(extent, target): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((16,), "float32"), n: T.int32): + def main(A: T.Tensor((16,), "float32"), n: T.int32): for i in T.vectorized(extent): if i <= n: A[i] = A[i] + T.float32(1) @@ -250,7 +250,7 @@ def test_vectorize_with_ge_cond(extent, target): @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((16,), "float32"), n: T.int32): + def main(A: T.Tensor((16,), "float32"), n: T.int32): for i in T.vectorized(extent): if i >= n: A[i] = A[i] + T.float32(1) @@ -267,14 +267,14 @@ def test_vectorize_if_then_else_scalarize(extent, target): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32")): for i in T.vectorized(extent): A[i] = T.if_then_else(i > 0, A[i] + T.float32(1), A[i]) @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32")): for i_s in range(extent): A[i_s] = T.if_then_else(i_s > 0, A[i_s] + T.float32(1), A[i_s]) @@ -288,7 +288,7 @@ def test_vectorize_if_then_else_vector(extent, target): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((25,), "float32"), n: T.int32): + def main(A: T.Tensor((25,), "float32"), n: T.int32): for i in range(n): for j in T.vectorized(extent): A[i * extent + j] = T.if_then_else(i > 0, A[i * extent + j], 0) @@ -296,7 +296,7 @@ def main(A: T.Buffer((25,), "float32"), n: T.int32): @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((25,), "float32"), n: T.int32): + def main(A: T.Tensor((25,), "float32"), n: T.int32): for i in range(n): A[T.Ramp(i * extent, 1, extent)] = T.if_then_else( i > 0, A[T.Ramp(i * extent, 1, extent)], T.Broadcast(0, extent) @@ -337,16 +337,16 @@ def test_vectorize_while_fail(): class Module: @T.prim_func def main( - A: T.Buffer((64,), "float32"), - B: T.Buffer((64,), "float32"), - C: T.Buffer((64,), "float32"), + A: T.Tensor((64,), "float32"), + B: T.Tensor((64,), "float32"), + C: T.Tensor((64,), "float32"), ): # Initialize C to 0 for j in range(64): C[j] = T.float32(0) # While loop inside vectorized loop (should fail) - i = T.decl_buffer((1,), "int32", scope="local") + i = T.decl_tensor((1,), "int32", scope="local") i[0] = 0 for j in T.vectorized(64): while i[0] < 10: @@ -370,14 +370,14 @@ def test_vectorize_with_reinterpret(extent, vec_str, target): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((16,), "int32"), B: T.Buffer((16,), "float32")): + def main(A: T.Tensor((16,), "int32"), B: T.Tensor((16,), "float32")): for i in T.vectorized(0, extent): B[i] = T.reinterpret("float32", A[i]) @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((16,), "int32"), B: T.Buffer((16,), "float32")): + def main(A: T.Tensor((16,), "int32"), B: T.Tensor((16,), "float32")): B[T.Ramp(0, 1, extent)] = T.reinterpret(vec_str, A[T.Ramp(0, 1, extent)]) with tvm.target.Target(target): @@ -410,14 +410,14 @@ def test_vectorize_binary(op, extent, target): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((25,), "float32"), B: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32"), B: T.Tensor((25,), "float32")): for j in T.vectorized(extent): A[j] = op(T.float32(3), B[j]) @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((25,), "float32"), B: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32"), B: T.Tensor((25,), "float32")): A[T.Ramp(0, 1, extent)] = op(T.Broadcast(T.float32(3), extent), B[T.Ramp(0, 1, extent)]) with tvm.target.Target(target): @@ -431,14 +431,14 @@ def test_vectorize_logical(op, extent, target): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((25,), "bool"), B: T.Buffer((25,), "bool")): + def main(A: T.Tensor((25,), "bool"), B: T.Tensor((25,), "bool")): for j in T.vectorized(extent): A[j] = op(T.bool(1), B[j]) @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((25,), "bool"), B: T.Buffer((25,), "bool")): + def main(A: T.Tensor((25,), "bool"), B: T.Tensor((25,), "bool")): A[T.Ramp(0, 1, extent)] = op(T.Broadcast(T.bool(1), extent), B[T.Ramp(0, 1, extent)]) with tvm.target.Target(target): @@ -451,14 +451,14 @@ def test_vectorize_select(extent, target): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((25,), "float32"), B: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32"), B: T.Tensor((25,), "float32")): for j in T.vectorized(extent): A[j] = T.Select(T.bool(True), A[j], B[j]) @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((25,), "float32"), B: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32"), B: T.Tensor((25,), "float32")): A[T.Ramp(0, 1, extent)] = T.Select( T.Broadcast(T.bool(True), extent), A[T.Ramp(0, 1, extent)], @@ -478,14 +478,14 @@ def test_vectorize_cast(extent, vec_str, target): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((25,), "int32"), B: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "int32"), B: T.Tensor((25,), "float32")): for j in T.vectorized(extent): A[j] = T.Cast("int32", B[j]) @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((25,), "int32"), B: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "int32"), B: T.Tensor((25,), "float32")): A[T.Ramp(0, 1, extent)] = T.Cast(vec_str, B[T.Ramp(0, 1, extent)]) with tvm.target.Target(target): @@ -499,7 +499,7 @@ def test_illegal_extent(): @I.ir_module(check_well_formed=False) class Mod: @T.prim_func - def main(A: T.Buffer((25,), "int32")): + def main(A: T.Tensor((25,), "int32")): for j in T.vectorized(n): A[j] = 3 @@ -512,7 +512,7 @@ def test_illegal_vscale_in_non_sve_compilation(): @I.ir_module class Mod: @T.prim_func - def main(A: T.Buffer((16,), "float32")): + def main(A: T.Tensor((16,), "float32")): for j in T.vectorized(0, 4 * T.vscale()): A[j] = 13 @@ -524,7 +524,7 @@ def main(A: T.Buffer((16,), "float32")): def test_vectorize_and_predicate_all_buffer_loads_stores(): @T.prim_func - def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def before(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i_0 in T.serial(T.ceildiv(14, 4)): for i_1 in T.vectorized(4): @@ -532,7 +532,7 @@ def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): B[i_0 * 4 + i_1] = A[i_0 * 4 + i_1] + 1.0 @T.prim_func - def expected(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def expected(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i_0 in range(4): T.evaluate( @@ -565,7 +565,7 @@ def test_vectorize_and_predicate_some_buffer_loads_stores(): # Currently revert to scalarizing the block if not all accesses # have been predicated, otherwise incorrect code is generated. @T.prim_func - def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def before(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i_0 in T.serial(T.ceildiv(14, 4)): for i_1 in T.vectorized(4): @@ -573,7 +573,7 @@ def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): B[i_0 * 4 + i_1] = A[i_0] + 1.0 @T.prim_func - def expected(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def expected(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i_0, i_1_s in T.grid(4, 4): if i_0 * 4 + i_1_s < 14: @@ -587,7 +587,7 @@ def expected(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): def test_vectorize_and_predicate_multiple_access_statements(): @T.prim_func - def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def before(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i_0 in T.serial(T.ceildiv(14, 4)): for i_1 in T.vectorized(4): @@ -596,7 +596,7 @@ def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): B[i_0 * 4 + i_1] = 1.0 @T.prim_func - def expected(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def expected(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i_0 in range(4): T.evaluate( @@ -632,7 +632,7 @@ def test_vectorize_nested_predicates_preserve_both_masks(): ) @T.prim_func - def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def before(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): for i_0 in T.serial(4): for i_1 in T.vectorized(4): if i_0 * 4 + i_1 < 15: @@ -656,7 +656,7 @@ def collect_predicates(node): def test_vectorize_and_predicate_invalid_conditions(): @T.prim_func - def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def before(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i_0 in T.serial(T.ceildiv(14, 4)): for i_1 in T.vectorized(4): @@ -668,7 +668,7 @@ def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): A[i_0 * 4 + i_1] = 2.0 @T.prim_func - def expected(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def expected(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i_0 in range(4): for i_1_s in range(4): @@ -692,7 +692,7 @@ def test_vectorize_with_explicitly_disabled_buffer_level_predication(): # by default. However, it has been explicitly disabled by the pass context # option, so no buffer-level predicates should be added. @T.prim_func - def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def before(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i_0 in T.serial(T.ceildiv(14, 4)): for i_1 in T.vectorized(4): @@ -700,7 +700,7 @@ def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): B[i_0 * 4 + i_1] = A[i_0 * 4 + i_1] + 1.0 @T.prim_func - def expected(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def expected(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for i_0, i_1_s in T.grid(4, 4): if i_0 * 4 + i_1_s < 14: @@ -715,7 +715,7 @@ def expected(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): def test_vectorize_and_predicate_buffer_load_stores_with_sve_func_attr_target(): @T.prim_func - def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def before(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True, "target": sve_target}) for i_0 in T.serial(T.ceildiv(14, 4)): for i_1 in T.vectorized(4): @@ -723,7 +723,7 @@ def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): B[i_0 * 4 + i_1] = A[i_0 * 4 + i_1] + 1.0 @T.prim_func - def expected(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def expected(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True, "target": sve_target}) for i_0 in range(4): T.evaluate( @@ -753,7 +753,7 @@ def expected(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): def test_vectorize_and_predicate_buffer_load_stores_with_sve_attr_scope_target(): @T.prim_func - def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def before(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with T.attr(sve_target, "target", 0): for i_0 in T.serial(T.ceildiv(14, 4)): @@ -762,7 +762,7 @@ def before(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): B[i_0 * 4 + i_1] = A[i_0 * 4 + i_1] + 1.0 @T.prim_func - def expected(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def expected(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.func_attr({"global_symbol": "main", "tirx.noalias": True}) with T.attr(sve_target, "target", 0): for i_0 in range(4): @@ -799,14 +799,14 @@ def test_vectorize_llvm_pure_intrin(extent, vec_str, target): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((25,), "float32"), B: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32"), B: T.Tensor((25,), "float32")): for j in T.vectorized(extent): A[j] = T.call_llvm_pure_intrin("float32", "llvm.sqrt", B[j]) @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((25,), "float32"), B: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "float32"), B: T.Tensor((25,), "float32")): A[T.Ramp(0, 1, extent)] = T.call_llvm_pure_intrin( vec_str, "llvm.sqrt", B[T.Ramp(0, 1, extent)] ) @@ -825,14 +825,14 @@ def test_vectorize_llvm_pure_intrin_fail(extent, vec_str, target): @I.ir_module class Before: @T.prim_func - def main(A: T.Buffer((25,), "int32"), B: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "int32"), B: T.Tensor((25,), "float32")): for j in T.vectorized(extent): A[j] = T.call_llvm_pure_intrin("int32", "llvm.lround", B[j]) @I.ir_module class After: @T.prim_func - def main(A: T.Buffer((25,), "int32"), B: T.Buffer((25,), "float32")): + def main(A: T.Tensor((25,), "int32"), B: T.Tensor((25,), "float32")): A[T.Ramp(0, 1, extent)] = T.call_llvm_pure_intrin( vec_str, "llvm.lround", B[T.Ramp(0, 1, extent)] ) diff --git a/tests/python/tirx/codegen/test_codegen_ampere.py b/tests/python/tirx/codegen/test_codegen_ampere.py index 709afd8d9470..646c1ec405cf 100644 --- a/tests/python/tirx/codegen/test_codegen_ampere.py +++ b/tests/python/tirx/codegen/test_codegen_ampere.py @@ -90,10 +90,10 @@ def test_ptx_mma_m16n8k16(a_type, no_c_ptr): # fmt: off @T.prim_func def main( - D: T.Buffer((16, 8), "float32"), - A: T.Buffer((16, 16), a_type), - B: T.Buffer((16, 8), b_type), - C: T.Buffer((16, 8), "float32"), + D: T.Tensor((16, 8), "float32"), + A: T.Tensor((16, 16), a_type), + B: T.Tensor((16, 8), b_type), + C: T.Tensor((16, 8), "float32"), ): T.device_entry() cta_id = T.cta_id([1]) @@ -167,10 +167,10 @@ def test_ptx_mma_m16n8k8(a_type, no_c_ptr): # fmt: off @T.prim_func def main( - D: T.Buffer((16, 8), "float32"), - A: T.Buffer((16, 8), a_type), - B: T.Buffer((8, 8), b_type), - C: T.Buffer((16, 8), "float32"), + D: T.Tensor((16, 8), "float32"), + A: T.Tensor((16, 8), a_type), + B: T.Tensor((8, 8), b_type), + C: T.Tensor((16, 8), "float32"), ): T.device_entry() cta_id = T.cta_id([1]) diff --git a/tests/python/tirx/codegen/test_codegen_blackwell.py b/tests/python/tirx/codegen/test_codegen_blackwell.py index 075cd622f948..6cd2497bb24e 100644 --- a/tests/python/tirx/codegen/test_codegen_blackwell.py +++ b/tests/python/tirx/codegen/test_codegen_blackwell.py @@ -54,7 +54,7 @@ def visit(node): if isinstance(node, tvm.tirx.Bind) and node.var.name == "remote_mbar_ptr": bindings.append(node) if ( - _is_buffer_binding(node, "tirx.decl_buffer") + _is_buffer_binding(node, "tirx.decl_tensor") and getattr(node.value.args[0], "name", None) == "remote_mbar_ptr" ): buffers.append(node) @@ -94,13 +94,13 @@ def test_tmem_alloc_dealloc_relinquish(): # fmt: off @T.prim_func - def test_tmem(A: T.Buffer((16, 16), "float16")): + def test_tmem(A: T.Tensor((16, 16), "float16")): T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([4]) lane_id = T.lane_id([32]) tid = T.thread_id([128]) - # tmem_addr = T.alloc_buffer((1,), "uint32", scope="shared", align=8) + # tmem_addr = T.alloc_tensor((1,), "uint32", scope="shared", align=8) tmem_addr = T.shared_scalar("uint32") # alloc TMEM @@ -131,7 +131,7 @@ def test_tmem(A: T.Buffer((16, 16), "float16")): def test_mbarrier_try_wait_once_codegen(): # fmt: off @T.prim_func - def test_try_wait_once(A: T.Buffer((16, 16), "float16")): + def test_try_wait_once(A: T.Tensor((16, 16), "float16")): T.device_entry() T.cta_id([1]) T.thread_id([128]) @@ -323,7 +323,7 @@ def nested_remote_view(): def test_fence_before_after_thread_sync(): # fmt: off @T.prim_func - def test_fence(A: T.Buffer((16, 16), "float16")): + def test_fence(A: T.Tensor((16, 16), "float16")): T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([4]) @@ -352,14 +352,14 @@ def test_tcgen05_ld_st_roundtrip(): # fmt: off @T.prim_func - def test_ld_st(A: T.Buffer((HEIGHT, WIDTH), "float32"), B: T.Buffer((HEIGHT, WIDTH), "float32")): # noqa: E501 + def test_ld_st(A: T.Tensor((HEIGHT, WIDTH), "float32"), B: T.Tensor((HEIGHT, WIDTH), "float32")): # noqa: E501 T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([4]) lane_id = T.lane_id([32]) tx = T.thread_id([128]) - reg = T.alloc_buffer((WIDTH,), "float32", scope="local") - # tmem_addr = T.alloc_buffer((1,), "uint32", scope="shared", align=8) + reg = T.alloc_tensor((WIDTH,), "float32", scope="local") + # tmem_addr = T.alloc_tensor((1,), "uint32", scope="shared", align=8) tmem_addr = T.shared_scalar("uint32") # alloc TMEM @@ -431,20 +431,20 @@ def test_tcgen05_cp_ld_roundtrip(): # fmt: off @T.prim_func - def test_cp_ld(A: T.Buffer((HEIGHT, WIDTH), dtype, layout=T.TileLayout(T.S[(HEIGHT, WIDTH // 4, 4) : (4, HEIGHT * 4, 1)])), # noqa: E501 - B: T.Buffer((HEIGHT, WIDTH), dtype, layout=T.TileLayout(T.S[(HEIGHT, WIDTH // 4, 4) : (4, HEIGHT * 4, 1)]))): # noqa: E501 + def test_cp_ld(A: T.Tensor((HEIGHT, WIDTH), dtype, layout=T.TileLayout(T.S[(HEIGHT, WIDTH // 4, 4) : (4, HEIGHT * 4, 1)])), # noqa: E501 + B: T.Tensor((HEIGHT, WIDTH), dtype, layout=T.TileLayout(T.S[(HEIGHT, WIDTH // 4, 4) : (4, HEIGHT * 4, 1)]))): # noqa: E501 T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([4]) lane_id = T.lane_id([32]) tx = T.thread_id([128]) - A_smem = T.alloc_buffer((HEIGHT, WIDTH), dtype, scope="shared", layout=A_layout) - reg = T.alloc_buffer((WIDTH,), dtype, scope="local") - # tmem_addr = T.alloc_buffer((1,), "uint32", scope="shared", align=8) + A_smem = T.alloc_tensor((HEIGHT, WIDTH), dtype, scope="shared", layout=A_layout) + reg = T.alloc_tensor((WIDTH,), dtype, scope="local") + # tmem_addr = T.alloc_tensor((1,), "uint32", scope="shared", align=8) tmem_addr = T.shared_scalar("uint32") - descA = T.alloc_buffer((1,), "uint64", scope="local") - bar = T.alloc_buffer((1,), "uint64", scope="shared", align=8) - phase = T.alloc_buffer((1,), "int32", scope="local") + descA = T.alloc_tensor((1,), "uint64", scope="local") + bar = T.alloc_tensor((1,), "uint64", scope="shared", align=8) + phase = T.alloc_tensor((1,), "int32", scope="local") # alloc TMEM if warp_id == 0: @@ -543,25 +543,25 @@ def test_tcgen05_mma_ss_no_tma(swizzle): # fmt: off @T.prim_func - def test_mma_ss_no_tma(A: T.Buffer((M, K), a_type, layout=T.TileLayout(T.S[M, K])), - B: T.Buffer((N, K), b_type, layout=T.TileLayout(T.S[N, K])), - C: T.Buffer((M, N), d_type)): + def test_mma_ss_no_tma(A: T.Tensor((M, K), a_type, layout=T.TileLayout(T.S[M, K])), + B: T.Tensor((N, K), b_type, layout=T.TileLayout(T.S[N, K])), + C: T.Tensor((M, N), d_type)): T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([4]) lane_id = T.lane_id([32]) tx = T.thread_id([128]) - dyn = T.alloc_buffer((dyn_smem_bytes,), "uint8", scope="shared") + dyn = T.alloc_tensor((dyn_smem_bytes,), "uint8", scope="shared") tmem_addr = T.decl_scalar("uint32", dyn.data, scope="shared", elem_offset=0) - A_smem = T.decl_buffer((M, K), a_type, dyn.data, elem_offset=256, layout=A_layout) - B_smem = T.decl_buffer((N, K), b_type, dyn.data, elem_offset=256 + M*K, layout=B_layout) - bar = T.decl_buffer((1,), "uint64", dyn.data, scope="shared", elem_offset=8) + A_smem = T.decl_tensor((M, K), a_type, dyn.data, elem_offset=256, layout=A_layout) + B_smem = T.decl_tensor((N, K), b_type, dyn.data, elem_offset=256 + M*K, layout=B_layout) + bar = T.decl_tensor((1,), "uint64", dyn.data, scope="shared", elem_offset=8) - reg = T.alloc_buffer((N,), d_type, scope="local") - descA = T.alloc_buffer((1,), "uint64", scope="local") - descB = T.alloc_buffer((1,), "uint64", scope="local") - descI = T.alloc_buffer((1,), "uint32", scope="local") - phase = T.alloc_buffer((1,), "int32", scope="local") + reg = T.alloc_tensor((N,), d_type, scope="local") + descA = T.alloc_tensor((1,), "uint64", scope="local") + descB = T.alloc_tensor((1,), "uint64", scope="local") + descI = T.alloc_tensor((1,), "uint32", scope="local") + phase = T.alloc_tensor((1,), "int32", scope="local") # alloc TMEM if warp_id == 0: @@ -644,11 +644,11 @@ def test_tcgen05_mma_pred_codegen(): def test_mma_pred(): T.device_entry() T.thread_id([1]) - tmem_addr = T.alloc_buffer((1,), "uint32", scope="local") - desc_a = T.alloc_buffer((1,), "uint64", scope="local") - desc_b = T.alloc_buffer((1,), "uint64", scope="local") - desc_i = T.alloc_buffer((1,), "uint32", scope="local") - pred = T.alloc_buffer((1,), "uint32", scope="local") + tmem_addr = T.alloc_tensor((1,), "uint32", scope="local") + desc_a = T.alloc_tensor((1,), "uint64", scope="local") + desc_b = T.alloc_tensor((1,), "uint64", scope="local") + desc_i = T.alloc_tensor((1,), "uint32", scope="local") + pred = T.alloc_tensor((1,), "uint32", scope="local") tmem_addr[0] = T.uint32(0) desc_a[0] = T.uint64(0) diff --git a/tests/python/tirx/codegen/test_codegen_cuda.py b/tests/python/tirx/codegen/test_codegen_cuda.py index 2e0123c7bd84..6d6405377760 100644 --- a/tests/python/tirx/codegen/test_codegen_cuda.py +++ b/tests/python/tirx/codegen/test_codegen_cuda.py @@ -119,7 +119,7 @@ def test_cuda_module_destructor_preserves_current_device(): torch = pytest.importorskip("torch") @T.prim_func - def main(A: T.Buffer((1,), "int32")): + def main(A: T.Tensor((1,), "int32")): T.device_entry() tx = T.thread_id([1]) if tx == 0: @@ -145,7 +145,7 @@ def main(A: T.Buffer((1,), "int32")): def test_vector_access_ptr_preserves_packed_offset(monkeypatch): - buffer = tvm.tirx.decl_buffer((8,), "int4x4", name="A") + buffer = tvm.tirx.decl_tensor((8,), "int4x4", name="A") data = tvm.tirx.Var("A_data", tvm.tirx.buffer_data_pointer_type(buffer)) access_ptr = buffer.access_ptr(access_mask=3, offset=2, extent=4) body = tvm.tirx.SeqStmt( @@ -153,7 +153,7 @@ def test_vector_access_ptr_preserves_packed_offset(monkeypatch): tvm.tirx.Bind( buffer, tvm.ir.Call( - "tirx.decl_buffer", + "tirx.decl_tensor", [ data, tvm.ir.Tuple(buffer.shape), @@ -184,7 +184,7 @@ def test_vector_access_ptr_preserves_packed_offset(monkeypatch): def _cuda_ldg_scalar_kernel(dtype: str): @T.prim_func - def main(src: T.Buffer((1,), dtype), out: T.Buffer((1,), dtype)): + def main(src: T.Tensor((1,), dtype), out: T.Tensor((1,), dtype)): T.device_entry() tx = T.thread_id([32]) if tx == 0: @@ -199,7 +199,7 @@ def _cuda_ldg_vector_kernel(dtype: str, vec: str): if vec_len == 2: @T.prim_func - def main(src: T.Buffer((2,), dtype), out: T.Buffer((2,), dtype)): + def main(src: T.Tensor((2,), dtype), out: T.Tensor((2,), dtype)): T.device_entry() tx = T.thread_id([32]) tmp0 = T.alloc_local((1,), dtype) @@ -212,7 +212,7 @@ def main(src: T.Buffer((2,), dtype), out: T.Buffer((2,), dtype)): return main @T.prim_func - def main(src: T.Buffer((4,), dtype), out: T.Buffer((4,), dtype)): + def main(src: T.Tensor((4,), dtype), out: T.Tensor((4,), dtype)): T.device_entry() tx = T.thread_id([32]) tmp0 = T.alloc_local((1,), dtype) @@ -241,7 +241,7 @@ def main(src: T.Buffer((4,), dtype), out: T.Buffer((4,), dtype)): def test_tirx_launch_bounds_omits_min_blocks_without_persistent_schedule(): @T.prim_func - def main(A: T.Buffer((4,), "int32")): + def main(A: T.Tensor((4,), "int32")): T.device_entry() bx = T.cta_id([4]) tx = T.thread_id([128]) @@ -255,7 +255,7 @@ def main(A: T.Buffer((4,), "int32")): def test_tirx_launch_bounds_min_blocks_attr_sets_one_block_per_sm(): @T.prim_func - def main(A: T.Buffer((4,), "int32")): + def main(A: T.Tensor((4,), "int32")): T.device_entry() T.attr({"tirx.launch_bounds_min_blocks_per_sm": 1}) bx = T.cta_id([4]) @@ -270,7 +270,7 @@ def main(A: T.Buffer((4,), "int32")): def test_tirx_launch_bounds_max_blocks_per_cluster_emits_third_operand(): @T.prim_func - def main(A: T.Buffer((4,), "int32")): + def main(A: T.Tensor((4,), "int32")): T.device_entry() T.attr( { @@ -290,7 +290,7 @@ def main(A: T.Buffer((4,), "int32")): def test_tirx_max_registers_attr_emits_cuda_maxnreg(): @T.prim_func - def main(A: T.Buffer((4,), "int32")): + def main(A: T.Tensor((4,), "int32")): T.device_entry() T.attr({"tirx.max_registers": 92}) bx = T.cta_id([4]) @@ -306,7 +306,7 @@ def main(A: T.Buffer((4,), "int32")): def test_tirx_max_registers_rejects_launch_bounds(): @T.prim_func - def main(A: T.Buffer((4,), "int32")): + def main(A: T.Tensor((4,), "int32")): T.device_entry() T.attr( { @@ -325,7 +325,7 @@ def main(A: T.Buffer((4,), "int32")): def test_tirx_required_block_size_emits_cuda_block_size(): @T.prim_func - def main(A: T.Buffer((8,), "int32")): + def main(A: T.Tensor((8,), "int32")): T.device_entry() T.attr({"tirx.required_block_size": 1}) bx, by = T.cta_id([4, 2]) @@ -342,7 +342,7 @@ def main(A: T.Buffer((8,), "int32")): def test_tirx_required_block_size_emits_launch_bounds_when_requested(): @T.prim_func - def main(A: T.Buffer((4,), "int32")): + def main(A: T.Tensor((4,), "int32")): T.device_entry() T.attr( { @@ -366,7 +366,7 @@ def main(A: T.Buffer((4,), "int32")): def test_tirx_cuda_kernel_return_zero_codegen_is_void_early_return(): @T.prim_func - def main(A: T.Buffer((4,), "int32")): + def main(A: T.Tensor((4,), "int32")): T.device_entry() bx = T.cta_id([4]) tx = T.thread_id([32]) @@ -385,7 +385,7 @@ def main(A: T.Buffer((4,), "int32")): def test_serial_pragma_unroll_codegen(): @T.prim_func - def main(A: T.Buffer((4,), "int32")): + def main(A: T.Tensor((4,), "int32")): T.device_entry() tx = T.thread_id([32]) if tx == 0: @@ -402,7 +402,7 @@ def main(A: T.Buffer((4,), "int32")): def test_serial_pragma_unroll_count_codegen(): @T.prim_func - def main(A: T.Buffer((4,), "int32")): + def main(A: T.Tensor((4,), "int32")): T.device_entry() tx = T.thread_id([32]) if tx == 0: @@ -415,7 +415,7 @@ def main(A: T.Buffer((4,), "int32")): def test_serial_disable_unroll_pragma_immediately_precedes_dynamic_for(): @T.prim_func - def main(A: T.Buffer((4,), "int32")): + def main(A: T.Tensor((4,), "int32")): T.device_entry() tx = T.thread_id([32]) if tx == 0: @@ -430,7 +430,7 @@ def main(A: T.Buffer((4,), "int32")): def test_cluster_cta_id_codegen_uses_coordinate_sregs(): @T.prim_func - def main(A: T.Buffer((1,), "int32")): + def main(A: T.Tensor((1,), "int32")): T.device_entry() cbx, cby = T.cta_id_in_cluster([2, 2]) tx = T.thread_id([32]) @@ -447,7 +447,7 @@ def main(A: T.Buffer((1,), "int32")): @pytest.mark.gpu def test_cuda_handle_uint64_reinterpret_codegen(): @T.prim_func - def main(A: T.Buffer((1,), "uint64")): + def main(A: T.Tensor((1,), "uint64")): T.device_entry() tx = T.thread_id([32]) if tx == 0: @@ -464,7 +464,7 @@ def main(A: T.Buffer((1,), "uint64")): @pytest.mark.skipif(not env.has_cuda(), reason="need cuda") def test_cuda_atomic_add(): @T.prim_func - def main(A: T.Buffer((1,), "int32"), B: T.Buffer((1,), "float32")): + def main(A: T.Tensor((1,), "int32"), B: T.Tensor((1,), "float32")): T.device_entry() cta_id = T.cta_id([1]) tx = T.thread_id([32]) @@ -490,7 +490,7 @@ def run_and_check(): def test_ptx_ld_acquire_and_volatile_codegen(): @T.prim_func - def main(A: T.Buffer((1,), "uint64"), B: T.Buffer((1,), "int32"), C: T.Buffer((1,), "uint32")): + def main(A: T.Tensor((1,), "uint64"), B: T.Tensor((1,), "int32"), C: T.Tensor((1,), "uint32")): T.device_entry() tx = T.thread_id([32]) if tx == 0: @@ -510,7 +510,7 @@ def main(A: T.Buffer((1,), "uint64"), B: T.Buffer((1,), "int32"), C: T.Buffer((1 def test_ptx_f32x2_value_codegen(): @T.prim_func - def main(A: T.Buffer((2,), "uint64"), B: T.Buffer((2,), "float32")): + def main(A: T.Tensor((2,), "uint64"), B: T.Tensor((2,), "float32")): T.device_entry() tx = T.thread_id([32]) if tx == 0: @@ -540,7 +540,7 @@ def test_ptx_neg_f32_codegen(): """`neg{.ftz}.f32` (ISA 9.7.3.10) -- the exact form, without fast-math .ftz.""" @T.prim_func - def main(A: T.Buffer((2,), "float32")): + def main(A: T.Tensor((2,), "float32")): T.device_entry() tx = T.thread_id([32]) if tx == 0: @@ -555,7 +555,7 @@ def test_ptx_sub_f16x2_codegen(): """The packed half line `sub{.rnd}{.ftz}{.sat}.f16x2` (ISA 9.7.4.2).""" @T.prim_func - def main(A: T.Buffer((3,), "uint32")): + def main(A: T.Tensor((3,), "uint32")): T.device_entry() tx = T.thread_id([32]) if tx == 0: @@ -574,10 +574,10 @@ def test_sparse_decode_conversion_intrinsics_codegen(monkeypatch): @T.prim_func def main( - U16: T.Buffer((1,), "uint16"), - U32: T.Buffer((2,), "uint32"), - U64: T.Buffer((1,), "uint64"), - F32: T.Buffer((2,), "float32"), + U16: T.Tensor((1,), "uint16"), + U32: T.Tensor((2,), "uint32"), + U64: T.Tensor((1,), "uint64"), + F32: T.Tensor((2,), "float32"), ): T.device_entry() tx = T.thread_id([32]) @@ -601,10 +601,10 @@ def main( def test_megamoe_extracted_intrinsics_codegen(): @T.prim_func def main( - U32: T.Buffer((4,), "uint32"), - I32: T.Buffer((1,), "int32"), - U64: T.Buffer((1,), "uint64"), - F32: T.Buffer((4,), "float32"), + U32: T.Tensor((4,), "uint32"), + I32: T.Tensor((1,), "int32"), + U64: T.Tensor((1,), "uint64"), + F32: T.Tensor((4,), "float32"), ): T.device_entry() tx = T.thread_id([32]) @@ -696,9 +696,9 @@ def test_cuda_ldg_vector_rejects_unsupported_dtype(): def test_ptx_cp_async_bulk_non_tma_form_codegen(): @T.prim_func def main( - A: T.Buffer((128,), "float32"), - B: T.Buffer((128,), "float32"), - C: T.Buffer((1,), "uint64"), + A: T.Tensor((128,), "float32"), + B: T.Tensor((128,), "float32"), + C: T.Tensor((1,), "uint64"), ): T.device_entry() tx = T.thread_id([32]) @@ -728,12 +728,12 @@ def main( def test_ptx_sync_and_clc_codegen(): @T.prim_func - def main(A: T.Buffer((1,), "uint32")): + def main(A: T.Tensor((1,), "uint32")): T.device_entry() tx = T.thread_id([32]) if tx == 0: - bar = T.alloc_buffer((5,), "uint64", scope="shared", align=16) - response = T.alloc_buffer((4,), "uint32", scope="shared", align=16) + bar = T.alloc_tensor((5,), "uint64", scope="shared", align=16) + response = T.alloc_tensor((4,), "uint32", scope="shared", align=16) T.ptx.cp.async_.mbarrier.arrive.shared.b64(bar.ptr_to([0])) T.ptx.cp.async_.mbarrier.arrive.noinc.shared__cta.b64(bar.ptr_to([0])) T.cuda.mbarrier_wait(bar.ptr_to([0]), T.int32(0)) @@ -800,11 +800,11 @@ def main(A: T.Buffer((1,), "uint32")): def test_ptx_mbarrier_arrive_new_forms_codegen(): @T.prim_func - def main(Pred: T.Buffer((1,), "int32")): + def main(Pred: T.Tensor((1,), "int32")): T.device_entry() tx = T.thread_id([32]) if tx == 0: - bar = T.alloc_buffer((6,), "uint64", scope="shared", align=16) + bar = T.alloc_tensor((6,), "uint64", scope="shared", align=16) state = T.local_scalar("uint64") T.ptx.mbarrier.arrive.relaxed.cta.shared__cta.b64(bar.ptr_to([0])) T.ptx.mbarrier.arrive.relaxed.cluster.shared__cluster.b64(bar.ptr_to([1])) @@ -835,7 +835,7 @@ def main(Pred: T.Buffer((1,), "int32")): def test_cuda_ldg_vector_scatter_codegen(): @T.prim_func - def main(src: T.Buffer((4,), "int32"), out: T.Buffer((4,), "int32")): + def main(src: T.Tensor((4,), "int32"), out: T.Tensor((4,), "int32")): T.device_entry() tx = T.thread_id([32]) tmp0 = T.alloc_local((1,), "int32") @@ -895,14 +895,14 @@ def main(A_map: T.TensorMap()): def test_tma_cache_policy_operand_codegen(): @T.prim_func - def main(Cache: T.Buffer((1,), "uint64")): + def main(Cache: T.Tensor((1,), "uint64")): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) B_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) T.device_entry() tx = T.thread_id([32]) if tx == 0: - smem = T.alloc_buffer((128,), "float32", scope="shared", align=128) + smem = T.alloc_tensor((128,), "float32", scope="shared", align=128) bar = T.shared_scalar("uint64") T.ptx[_TMA_G2S_CG2_CACHE]( smem.data, T.address_of(A_map), 0, 0, T.address_of(bar), Cache[0] @@ -940,7 +940,7 @@ def main(Cache: T.Buffer((1,), "uint64")): def test_cuda_thread_fence(): @T.prim_func - def main(A: T.Buffer((16, 16), "int32")): + def main(A: T.Tensor((16, 16), "int32")): T.device_entry() cta_id = T.cta_id([1]) tx = T.thread_id([32]) @@ -953,7 +953,7 @@ def main(A: T.Buffer((16, 16), "int32")): def test_cuda_nano_sleep(): @T.prim_func - def main(A: T.Buffer((16, 16), "int32")): + def main(A: T.Tensor((16, 16), "int32")): T.device_entry() cta_id = T.cta_id([1]) tx = T.thread_id([32]) @@ -966,7 +966,7 @@ def main(A: T.Buffer((16, 16), "int32")): def test_cuda_atomic_cas(): @T.prim_func - def main(A: T.Buffer((16, 16), "int32")): + def main(A: T.Tensor((16, 16), "int32")): T.device_entry() cta_id = T.cta_id([1]) tx = T.thread_id([32]) @@ -988,7 +988,7 @@ def test_add_one(): """ @T.prim_func - def main(a: T.Buffer((16, 16), "int32"), b: T.Buffer((16, 16), "int32")): + def main(a: T.Tensor((16, 16), "int32"), b: T.Tensor((16, 16), "int32")): T.device_entry() cta_id = T.cta_id([1]) tx = T.thread_id([32]) @@ -1022,7 +1022,7 @@ def test_print(): """ @T.prim_func - def main(a: T.Buffer((16, 16), "int32")): + def main(a: T.Tensor((16, 16), "int32")): T.device_entry() cta_id = T.cta_id([1]) tx = T.thread_id([32]) @@ -1050,15 +1050,15 @@ def run_and_check(): def test_warp_shuffle_xor_sync(): # fmt: off @T.prim_func - def func(A: T.Buffer((32,), dtype='float32', align=16)): + def func(A: T.Tensor((32,), dtype='float32', align=16)): T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([1]) lane_id = T.lane_id([32]) - A_local = T.alloc_buffer([1], "float32", scope="local") - i = T.alloc_buffer([1], "int32", scope="local") + A_local = T.alloc_tensor([1], "float32", scope="local") + i = T.alloc_tensor([1], "int32", scope="local") A_local[0] = T.float32(31 - lane_id) i[0] = 16 @@ -1112,7 +1112,7 @@ def test_ptx_cp_async(cp_size, cache_hint, prefetch_size, predicate, fill_mode): # fmt: off @T.prim_func - def main(A: T.Buffer((N), "float16")): + def main(A: T.Tensor((N), "float16")): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([32]) @@ -1158,7 +1158,7 @@ def run_and_check(): def test_ptx_ldmatrix(trans, num): # fmt: off @T.prim_func - def main(A: T.Buffer((16, 16), "float16"), B: T.Buffer((16, 16), "float16")): + def main(A: T.Tensor((16, 16), "float16"), B: T.Tensor((16, 16), "float16")): T.device_entry() cta_id = T.cta_id([1]) tx = T.thread_id([32]) @@ -1225,7 +1225,7 @@ def run_and_check(): def test_uint32_loop_var_and_scope_id_emit_unsigned(): @T.prim_func - def main(A: T.Buffer((128,), "int32")): + def main(A: T.Tensor((128,), "int32")): T.device_entry() _ = T.cta_id([1]) tx = T.thread_id([128], dtype="uint32") @@ -1243,11 +1243,11 @@ def main(A: T.Buffer((128,), "int32")): @pytest.mark.skipif(not env.has_cuda(), reason="need cuda") def test_uint32_loop_var_runs_correctly(): @T.prim_func - def main(A: T.Buffer((128,), "int32"), B: T.Buffer((128,), "int32")): + def main(A: T.Tensor((128,), "int32"), B: T.Tensor((128,), "int32")): T.device_entry() _ = T.cta_id([1]) tx = T.thread_id([128], dtype="uint32") - acc = T.alloc_buffer((1,), "int32", scope="local") + acc = T.alloc_tensor((1,), "int32", scope="local") acc[0] = 0 for k in T.serial(4, dtype="uint32"): acc[0] = acc[0] + A[tx] + T.int32(k) diff --git a/tests/python/tirx/codegen/test_codegen_dsmem.py b/tests/python/tirx/codegen/test_codegen_dsmem.py index f93043cb3d60..be359b62fcd1 100644 --- a/tests/python/tirx/codegen/test_codegen_dsmem.py +++ b/tests/python/tirx/codegen/test_codegen_dsmem.py @@ -48,7 +48,7 @@ def test_ptx_cp_async_bulk_s2c_codegen(): # fmt: off @T.prim_func - def main(A: T.Buffer((128,), "float16")): + def main(A: T.Tensor((128,), "float16")): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([1]) @@ -79,7 +79,7 @@ def test_ptx_cp_async_bulk_s2c_codegen_address_conversion(): # fmt: off @T.prim_func - def main(A: T.Buffer((64,), "float32")): + def main(A: T.Tensor((64,), "float32")): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([1]) @@ -110,7 +110,7 @@ def test_mapa_pointer_bind_codegen(): # fmt: off @T.prim_func - def main(A: T.Buffer((1,), "uint64")): + def main(A: T.Tensor((1,), "uint64")): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([1]) @@ -118,7 +118,7 @@ def main(A: T.Buffer((1,), "uint64")): mapped = T.alloc_local([1], "uint64") T.ptx.mapa.u64(mapped[0], mbar.ptr_to([0]), T.uint32(0)) remote_ptr = T.reinterpret(ptr_ty, mapped[0]) - remote_mbar = T.decl_buffer([1], "uint64", data=remote_ptr, scope="shared") + remote_mbar = T.decl_tensor([1], "uint64", data=remote_ptr, scope="shared") A[0] = remote_mbar[0] # fmt: on @@ -127,7 +127,7 @@ def main(A: T.Buffer((1,), "uint64")): loads = [] def collect(node): - if _is_buffer_binding(node, "tirx.decl_buffer"): + if _is_buffer_binding(node, "tirx.decl_tensor"): decl_buffers.append(node) elif isinstance(node, tvm.tirx.Bind) and isinstance(node.var.ty, PointerType): binds.append(node) diff --git a/tests/python/tirx/codegen/test_codegen_hopper.py b/tests/python/tirx/codegen/test_codegen_hopper.py index 4f8af60a2f13..9cdc4afa5417 100644 --- a/tests/python/tirx/codegen/test_codegen_hopper.py +++ b/tests/python/tirx/codegen/test_codegen_hopper.py @@ -38,7 +38,7 @@ def _get_source(func: tvm.tirx.PrimFunc) -> tuple[str, tvm.IRModule]: def _run_tensormap_encode(shape, dtype, encode_args): # fmt: off @T.prim_func - def main(A: T.Buffer(shape, dtype=dtype, align=32)): + def main(A: T.Tensor(shape, dtype=dtype, align=32)): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) T.call_packed("runtime.cuTensorMapEncodeTiled", A_map, dtype, len(shape), A.data, *encode_args) # noqa: E501 @@ -65,7 +65,7 @@ def run_and_check(): def test_ptx_setmaxnreg(inc): # fmt: off @T.prim_func - def func(A: T.Buffer(1)): + def func(A: T.Tensor(1)): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([128]) @@ -86,12 +86,12 @@ def func(A: T.Buffer(1)): def test_stmatrix_sync_aligned(trans): # fmt: off @T.prim_func - def func(A: T.Buffer((16, 16), "float16")): + def func(A: T.Tensor((16, 16), "float16")): T.device_entry() cta_id = T.cta_id([1]) tx = T.thread_id([32]) - A_smem = T.alloc_buffer((16, 16), "float16", scope="shared", align=16) - reg = T.alloc_buffer((8,), "float16", scope="local") + A_smem = T.alloc_tensor((16, 16), "float16", scope="shared", align=16) + reg = T.alloc_tensor((8,), "float16", scope="local") for i in range(8): reg[i] = tx * 8 + i # stmatrix stores 4 b32 registers; reg is fp16, so they ride a uint32 @@ -153,7 +153,7 @@ def run_and_check(): def test_ptx_stmatrix(trans, num): # fmt: off @T.prim_func - def main(A: T.Buffer((16, 16), "float16")): + def main(A: T.Tensor((16, 16), "float16")): T.device_entry() cta_id = T.cta_id([1]) tx = T.thread_id([32]) @@ -225,7 +225,7 @@ def test_ptx_stmatrix_noncontiguous(trans, num): # fmt: off @T.prim_func - def main(A: T.Buffer((16, 16), "float16")): + def main(A: T.Tensor((16, 16), "float16")): T.device_entry() cta_id = T.cta_id([1]) tx = T.thread_id([32]) @@ -289,7 +289,7 @@ def run_and_check(): def test_bar_arrive(): # fmt: off @T.prim_func - def func(A: T.Buffer(1)): + def func(A: T.Tensor(1)): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([128]) @@ -306,7 +306,7 @@ def func(A: T.Buffer(1)): def test_bar_sync(): # fmt: off @T.prim_func - def func(A: T.Buffer(1)): + def func(A: T.Tensor(1)): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([128]) @@ -323,7 +323,7 @@ def func(A: T.Buffer(1)): def test_barrier_sync_unaligned(): # fmt: off @T.prim_func - def func(A: T.Buffer(1)): + def func(A: T.Tensor(1)): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([128]) @@ -340,7 +340,7 @@ def func(A: T.Buffer(1)): def test_fence_mbarrier_init_release_clsuter(): # fmt: off @T.prim_func - def func(A: T.Buffer(1)): + def func(A: T.Tensor(1)): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([128]) @@ -356,7 +356,7 @@ def func(A: T.Buffer(1)): def test_ptx_elect_sync(): # fmt: off @T.prim_func - def func(A: T.Buffer(1)): + def func(A: T.Tensor(1)): T.device_entry() cta_id = T.cta_id([1]) tx = T.thread_id([128]) @@ -377,7 +377,7 @@ def func(A: T.Buffer(1)): def test_ptx_fence(sem, scope): # fmt: off @T.prim_func - def func(A: T.Buffer(1)): + def func(A: T.Tensor(1)): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([128]) @@ -393,7 +393,7 @@ def func(A: T.Buffer(1)): def test_fence_proxy_async(): # fmt: off @T.prim_func - def func(A: T.Buffer(1)): + def func(A: T.Tensor(1)): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([128]) @@ -432,7 +432,7 @@ def get_ir(shape, tma_args): # fmt: off @T.prim_func def main( - A: T.Buffer(shape, dtype=dtype, align=16), B: T.Buffer(shape, dtype=dtype, align=16) + A: T.Tensor(shape, dtype=dtype, align=16), B: T.Tensor(shape, dtype=dtype, align=16) ): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) @@ -449,7 +449,7 @@ def main( for threadIdx in T.thread_binding(128, thread="threadIdx.x"): bar = T.shared_scalar("uint64") phase: T.int32 - A_smem = T.alloc_buffer(shape, dtype, scope="shared", align=128) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", align=128) phase = 0 if threadIdx == 0: @@ -633,8 +633,8 @@ def get_ir(swizzle, dtype): # fmt: off @T.prim_func def main( - A: T.Buffer(total_elems, dtype=dtype, align=16), - B: T.Buffer(total_elems, dtype=dtype, align=16), + A: T.Tensor(total_elems, dtype=dtype, align=16), + B: T.Tensor(total_elems, dtype=dtype, align=16), ): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) @@ -649,7 +649,7 @@ def main( T.device_entry() for blockIdx in T.thread_binding(1, thread="blockIdx.x"): for threadIdx in T.thread_binding(128, thread="threadIdx.x"): - A_smem = T.alloc_buffer((total_elems,), dtype, scope="shared", align=128) + A_smem = T.alloc_tensor((total_elems,), dtype, scope="shared", align=128) bar = T.shared_scalar("uint64") phase: T.int32 @@ -735,8 +735,8 @@ def get_ir(shape, tma_args): # fmt: off @T.prim_func def main( - A: T.Buffer(shape, dtype="float32", align=16), - B: T.Buffer(shape, dtype="float32", align=16), + A: T.Tensor(shape, dtype="float32", align=16), + B: T.Tensor(shape, dtype="float32", align=16), ): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) @@ -754,7 +754,7 @@ def main( for tx in T.thread_binding(128, thread="threadIdx.x"): bar = T.shared_scalar("uint64") phase: T.int32 - A_smem = T.alloc_buffer(shape[::-1], "float32", scope="shared", align=128) + A_smem = T.alloc_tensor(shape[::-1], "float32", scope="shared", align=128) phase = 0 if tx == 0: @@ -838,8 +838,8 @@ def get_ir(shape, tma_args): # fmt: off @T.prim_func def main( - A: T.Buffer(shape, dtype="float32", align=16), - B: T.Buffer(shape, dtype="float32", align=16), + A: T.Tensor(shape, dtype="float32", align=16), + B: T.Tensor(shape, dtype="float32", align=16), ): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) @@ -862,7 +862,7 @@ def main( for tx in T.thread_binding(128, thread="threadIdx.x"): bar = T.shared_scalar("uint64") phase: T.int32 - A_smem = T.alloc_buffer(shape[::-1], "float32", scope="shared", align=128) + A_smem = T.alloc_tensor(shape[::-1], "float32", scope="shared", align=128) phase = 0 if tx == 0: @@ -979,7 +979,7 @@ def get_ir(shape, tma_args): # fmt: off @T.prim_func - def main(A: T.Buffer(shape, dtype='float32', align=16)): + def main(A: T.Tensor(shape, dtype='float32', align=16)): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) T.call_packed("runtime.cuTensorMapEncodeTiled", A_map, "float32", len(shape), A.data, *tma_args) # noqa: E501 @@ -988,7 +988,7 @@ def main(A: T.Buffer(shape, dtype='float32', align=16)): cta_id = T.cta_id([1]) tx = T.thread_id([128]) - A_smem = T.alloc_buffer(elems, "float32", scope="shared", align=128) + A_smem = T.alloc_tensor(elems, "float32", scope="shared", align=128) if tx == 0: for i in T.serial(0, elems): @@ -1062,9 +1062,9 @@ def get_accum_list(C, C_elems): # fmt: off @T.prim_func def main( - A: T.Buffer(shapeA, dtype=in_dtype, align=16), - B: T.Buffer(shapeB, dtype=in_dtype, align=16), - C: T.Buffer(shapeC, dtype=out_dtype, align=16), + A: T.Tensor(shapeA, dtype=in_dtype, align=16), + B: T.Tensor(shapeB, dtype=in_dtype, align=16), + C: T.Tensor(shapeC, dtype=out_dtype, align=16), ): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) @@ -1080,14 +1080,14 @@ def main( cta_id = T.cta_id([1]) tx = T.thread_id([128]) # A warpgroup is 128 threads - A_smem = T.alloc_buffer(shapeA, in_dtype, scope="shared", align=1024) - B_smem = T.alloc_buffer(shapeB, in_dtype, scope="shared", align=1024) + A_smem = T.alloc_tensor(shapeA, in_dtype, scope="shared", align=1024) + B_smem = T.alloc_tensor(shapeB, in_dtype, scope="shared", align=1024) bar = T.shared_scalar("uint64") phase: T.int32 descA: T.uint64 descB: T.uint64 - C_local = T.alloc_buffer((C_elems,), out_dtype, scope="local") + C_local = T.alloc_tensor((C_elems,), out_dtype, scope="local") # init phase and bar phase = 0 @@ -1237,9 +1237,9 @@ def get_accum_list(C, C_elems): # fmt: off @T.prim_func def main( - A: T.Buffer(shapeA, dtype=in_dtype, align=16), - B: T.Buffer(shapeB, dtype=in_dtype, align=16), - C: T.Buffer(shapeC, dtype=out_dtype, align=16), + A: T.Tensor(shapeA, dtype=in_dtype, align=16), + B: T.Tensor(shapeB, dtype=in_dtype, align=16), + C: T.Tensor(shapeC, dtype=out_dtype, align=16), ): B_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) @@ -1251,17 +1251,17 @@ def main( cta_id = T.cta_id([1]) tx = T.thread_id([128]) # A warpgroup is 128 threads - B_smem = T.alloc_buffer(shapeB, in_dtype, scope="shared", align=1024) - # bar = T.alloc_buffer((1,), "uint64", scope="shared", align=8) + B_smem = T.alloc_tensor(shapeB, in_dtype, scope="shared", align=1024) + # bar = T.alloc_tensor((1,), "uint64", scope="shared", align=8) bar = T.shared_scalar("uint64") - # descB = T.alloc_buffer((1,), "uint64", scope="local") + # descB = T.alloc_tensor((1,), "uint64", scope="local") descB: T.uint64 - A_local = T.alloc_buffer((A_elems,), in_dtype, scope="local") - C_local = T.alloc_buffer((C_elems,), out_dtype, scope="local") + A_local = T.alloc_tensor((A_elems,), in_dtype, scope="local") + C_local = T.alloc_tensor((C_elems,), out_dtype, scope="local") A_elems_b32 = T.meta_var(A_elems // (32 // in_dtype_bits)) - A_local_b32 = T.decl_buffer((A_elems_b32,), "uint32", data=A_local.data) + A_local_b32 = T.decl_tensor((A_elems_b32,), "uint32", data=A_local.data) # load A to regs for i in T.serial(0, A_elems // 4): @@ -1384,12 +1384,12 @@ def run_and_check(): @pytest.mark.skipif(not env.has_cuda_compute(9), reason="need cuda compute >= 9.0") def test_mapa(): @T.prim_func - def func(A: T.Buffer(1)): + def func(A: T.Tensor(1)): T.device_entry() cbx = T.cta_id_in_cluster([2]) cta_id = T.cta_id([2]) tx = T.thread_id([128]) - A_smem = T.alloc_buffer([1], "uint32", scope="shared") + A_smem = T.alloc_tensor([1], "uint32", scope="shared") mapped = T.alloc_local([1], "uint64") if cbx == 0 and tx == 0: T.ptx.mapa.u64(mapped[0], A_smem.data, T.uint32(cbx)) diff --git a/tests/python/tirx/codegen/test_codegen_nki.py b/tests/python/tirx/codegen/test_codegen_nki.py index ca8965e7d361..83186308cef5 100644 --- a/tests/python/tirx/codegen/test_codegen_nki.py +++ b/tests/python/tirx/codegen/test_codegen_nki.py @@ -39,11 +39,11 @@ def compare_strings_ignore_whitespace(s1, s2): def test_nki_add_1(): # fmt: off @T.prim_func - def func(A: T.Buffer((128, 512)), B: T.Buffer((128, 512))): + def func(A: T.Tensor((128, 512)), B: T.Tensor((128, 512))): T.func_attr({"num_inputs": 1}) T.device_entry() - A_sbuf = T.alloc_buffer((128, 512), "float32", scope="trn.sbuf",) - B_sbuf = T.alloc_buffer((128, 512), "float32", scope="trn.sbuf",) + A_sbuf = T.alloc_tensor((128, 512), "float32", scope="trn.sbuf",) + B_sbuf = T.alloc_tensor((128, 512), "float32", scope="trn.sbuf",) with T.attr(0, "tensorized_nki_instruction", 1): for i in range(0, 128): for j in range(0, 512): @@ -93,11 +93,11 @@ def func_kernel(A_ptr, B_ptr: nt.mutable_tensor, ): def test_nki_add_2(): # fmt: off @T.prim_func - def func(A: T.Buffer((128, 2048)), B: T.Buffer((128, 2048))): + def func(A: T.Tensor((128, 2048)), B: T.Tensor((128, 2048))): T.func_attr({"num_inputs": 1}) T.device_entry() - A_sbuf = T.alloc_buffer((128, 512), "float32", scope="trn.sbuf",) - B_sbuf = T.alloc_buffer((128, 512), "float32", scope="trn.sbuf",) + A_sbuf = T.alloc_tensor((128, 512), "float32", scope="trn.sbuf",) + B_sbuf = T.alloc_tensor((128, 512), "float32", scope="trn.sbuf",) for k in range(0, 4): with T.attr(0, "tensorized_nki_instruction", 1): for i in range(0, 128): @@ -170,22 +170,22 @@ def test_nki_matmul_1(): @T.prim_func def func( - lhsT: T.Buffer((K, M), "float16"), - rhs: T.Buffer((K, N), "float16"), + lhsT: T.Tensor((K, M), "float16"), + rhs: T.Tensor((K, N), "float16"), result: T.buffer((M, N), "float16"), ): T.func_attr({"num_inputs": 2}) - result_tiles = T.alloc_buffer( + result_tiles = T.alloc_tensor( (TILE_M, NUM_BLOCK_M, TILES_IN_BLOCK_M, TILES_IN_BLOCK_N, TILE_N), "float32", scope="trn.sbuf", ) - rhs_tiles = T.alloc_buffer((TILE_K, TILES_IN_BLOCK_K, BLOCK_N), "float16", scope="trn.sbuf") - lhsT_tiles = T.alloc_buffer( + rhs_tiles = T.alloc_tensor((TILE_K, TILES_IN_BLOCK_K, BLOCK_N), "float16", scope="trn.sbuf") + lhsT_tiles = T.alloc_tensor( (TILE_K, TILES_IN_BLOCK_K, BLOCK_M), "float16", scope="trn.sbuf" ) - res_tile = T.alloc_buffer((1, TILE_M, TILE_N), "float32", scope="trn.psum") - result_packed = T.alloc_buffer((TILE_K, BLOCK_N), "float32", scope="trn.sbuf") + res_tile = T.alloc_tensor((1, TILE_M, TILE_N), "float32", scope="trn.psum") + result_packed = T.alloc_tensor((TILE_K, BLOCK_N), "float32", scope="trn.sbuf") for n in range(NUM_BLOCK_N): with T.attr(0, "tensorized_nki_instruction", 1): for i0 in range(TILE_M): diff --git a/tests/python/tirx/codegen/test_codegen_nvshmem.py b/tests/python/tirx/codegen/test_codegen_nvshmem.py index fc1c6ce577c3..f58c0ca03c2e 100644 --- a/tests/python/tirx/codegen/test_codegen_nvshmem.py +++ b/tests/python/tirx/codegen/test_codegen_nvshmem.py @@ -77,7 +77,7 @@ def _test_func(): def test_thread_info(sess): @T.prim_func - def main(res: T.Buffer((2,), "int32")): + def main(res: T.Tensor((2,), "int32")): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([nwarps * 32]) @@ -97,7 +97,7 @@ def test_transfer(sess, scope, shape, nwarps, nelems, op_name): # fmt: off @T.prim_func - def main(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)): + def main(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)): T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([nwarps]) @@ -136,7 +136,7 @@ def test_signal_op(sess, sig_op): # fmt: off @T.prim_func - def main(res: T.Buffer((1,), "uint64")): + def main(res: T.Tensor((1,), "uint64")): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([nwarps * 32]) @@ -170,9 +170,9 @@ def test_put_signal(sess, scope, shape, nwarps, nelems, cmp_value): @T.prim_func def main( - A: T.Buffer(shape, dtype), - B: T.Buffer(shape, dtype), - signal_array: T.Buffer((1,), "uint64"), + A: T.Tensor(shape, dtype), + B: T.Tensor(shape, dtype), + signal_array: T.Tensor((1,), "uint64"), ): T.device_entry() cta_id = T.cta_id([1]) @@ -225,7 +225,7 @@ def test_fence_barrier(sess): # fmt: off @T.prim_func - def main(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype), res: T.Buffer((1,), "uint64")): # noqa: E501 + def main(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype), res: T.Tensor((1,), "uint64")): # noqa: E501 T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([nwarps]) diff --git a/tests/python/tirx/codegen/test_cuda_cta_reduce.py b/tests/python/tirx/codegen/test_cuda_cta_reduce.py index b16334eb4cbf..e6f109e0ccd9 100644 --- a/tests/python/tirx/codegen/test_cuda_cta_reduce.py +++ b/tests/python/tirx/codegen/test_cuda_cta_reduce.py @@ -49,14 +49,14 @@ def test_cta_sum_4_warps(): # fmt: off @T.prim_func - def func(out: T.Buffer((N,), 'float32')): + def func(out: T.Tensor((N,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([NUM_WARPS]) lane_id = T.lane_id([32]) tid = T.thread_id([N]) - scratch = T.alloc_buffer((NUM_WARPS,), "float32", scope="shared") + scratch = T.alloc_tensor((NUM_WARPS,), "float32", scope="shared") val: T.f32 = T.float32(tid + 1) val = T.cuda.cta_sum(val, NUM_WARPS, scratch.ptr_to([0])) out[tid] = val @@ -77,14 +77,14 @@ def test_cta_sum_8_warps(): # fmt: off @T.prim_func - def func(out: T.Buffer((N,), 'float32')): + def func(out: T.Tensor((N,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([NUM_WARPS]) lane_id = T.lane_id([32]) tid = T.thread_id([N]) - scratch = T.alloc_buffer((NUM_WARPS,), "float32", scope="shared") + scratch = T.alloc_tensor((NUM_WARPS,), "float32", scope="shared") val: T.f32 = T.float32(tid + 1) val = T.cuda.cta_sum(val, NUM_WARPS, scratch.ptr_to([0])) out[tid] = val @@ -104,14 +104,14 @@ def test_cta_max_4_warps(): # fmt: off @T.prim_func - def func(out: T.Buffer((N,), 'float32')): + def func(out: T.Tensor((N,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([NUM_WARPS]) lane_id = T.lane_id([32]) tid = T.thread_id([N]) - scratch = T.alloc_buffer((NUM_WARPS,), "float32", scope="shared") + scratch = T.alloc_tensor((NUM_WARPS,), "float32", scope="shared") val: T.f32 = T.float32(tid + 1) val = T.cuda.cta_max(val, NUM_WARPS, scratch.ptr_to([0])) out[tid] = val @@ -130,14 +130,14 @@ def test_cta_min_4_warps(): # fmt: off @T.prim_func - def func(out: T.Buffer((N,), 'float32')): + def func(out: T.Tensor((N,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([NUM_WARPS]) lane_id = T.lane_id([32]) tid = T.thread_id([N]) - scratch = T.alloc_buffer((NUM_WARPS,), "float32", scope="shared") + scratch = T.alloc_tensor((NUM_WARPS,), "float32", scope="shared") val: T.f32 = T.float32(tid + 1) val = T.cuda.cta_min(val, NUM_WARPS, scratch.ptr_to([0])) out[tid] = val @@ -156,14 +156,14 @@ def test_cta_sum_1_warp(): # fmt: off @T.prim_func - def func(out: T.Buffer((N,), 'float32')): + def func(out: T.Tensor((N,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([NUM_WARPS]) lane_id = T.lane_id([32]) tid = T.thread_id([N]) - scratch = T.alloc_buffer((NUM_WARPS,), "float32", scope="shared") + scratch = T.alloc_tensor((NUM_WARPS,), "float32", scope="shared") val: T.f32 = T.float32(tid + 1) val = T.cuda.cta_sum(val, NUM_WARPS, scratch.ptr_to([0])) out[tid] = val @@ -183,14 +183,14 @@ def test_cta_sum_all_warp_counts(num_warps): # fmt: off @T.prim_func - def func(out: T.Buffer((N,), 'float32')): + def func(out: T.Tensor((N,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([num_warps]) lane_id = T.lane_id([32]) tid = T.thread_id([N]) - scratch = T.alloc_buffer((num_warps,), "float32", scope="shared") + scratch = T.alloc_tensor((num_warps,), "float32", scope="shared") val: T.f32 = T.float32(tid + 1) val = T.cuda.cta_sum(val, num_warps, scratch.ptr_to([0])) out[tid] = val diff --git a/tests/python/tirx/codegen/test_cuda_wait_until.py b/tests/python/tirx/codegen/test_cuda_wait_until.py index 20be1e8da927..f0b6496d8574 100644 --- a/tests/python/tirx/codegen/test_cuda_wait_until.py +++ b/tests/python/tirx/codegen/test_cuda_wait_until.py @@ -45,7 +45,7 @@ def rendezvous(backoff_ns=None, ptx_type=None): """An N-way barrier on a monotone counter, as radix_topk_multi_cta writes it.""" @T.prim_func - def kernel(state: T.Buffer((1,), "int32"), participants: T.int32): + def kernel(state: T.Tensor((1,), "int32"), participants: T.int32): T.device_entry() T.cta_id([2]) lane = T.thread_id([32]) @@ -71,7 +71,7 @@ def packed_contribution(): """Counter in the high half, payload in the low half: DeepEP's notify slot.""" @T.prim_func - def kernel(slot: T.Buffer((1,), "uint64"), out: T.Buffer((1,), "uint64"), n: T.int32): + def kernel(slot: T.Tensor((1,), "uint64"), out: T.Tensor((1,), "uint64"), n: T.int32): T.device_entry() T.cta_id([2]) lane = T.thread_id([32]) @@ -299,7 +299,7 @@ def test_a_wide_word_cannot_be_waited_on(): with pytest.raises(Exception, match="does not take a 128-bit word"): @T.prim_func - def kernel(response: T.Buffer((2,), "uint64")): + def kernel(response: T.Tensor((2,), "uint64")): T.device_entry() T.cta_id([1]) lane = T.thread_id([32]) diff --git a/tests/python/tirx/codegen/test_cuda_warp_reduce.py b/tests/python/tirx/codegen/test_cuda_warp_reduce.py index 2713c0f76d2a..488017a7922e 100644 --- a/tests/python/tirx/codegen/test_cuda_warp_reduce.py +++ b/tests/python/tirx/codegen/test_cuda_warp_reduce.py @@ -47,7 +47,7 @@ def test_warp_sum_full(): # fmt: off @T.prim_func - def func(out: T.Buffer((32,), 'float32')): + def func(out: T.Tensor((32,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) @@ -71,7 +71,7 @@ def test_warp_sum_partial_8(): # fmt: off @T.prim_func - def func(out: T.Buffer((32,), 'float32')): + def func(out: T.Tensor((32,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) @@ -101,7 +101,7 @@ def test_warp_max_partial_4(): # fmt: off @T.prim_func - def func(out: T.Buffer((32,), 'float32')): + def func(out: T.Tensor((32,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) @@ -127,7 +127,7 @@ def test_warp_min_full(): # fmt: off @T.prim_func - def func(out: T.Buffer((32,), 'float32')): + def func(out: T.Tensor((32,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) @@ -149,7 +149,7 @@ def test_warp_sum_partial_2(): # fmt: off @T.prim_func - def func(out: T.Buffer((32,), 'float32')): + def func(out: T.Tensor((32,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) @@ -178,7 +178,7 @@ def test_warp_sum_all_widths(width): # fmt: off @T.prim_func - def func(out: T.Buffer((32,), 'float32')): + def func(out: T.Tensor((32,), 'float32')): T.device_entry() cta_id = T.cta_id([1]) diff --git a/tests/python/tirx/codegen/test_ptx_addr.py b/tests/python/tirx/codegen/test_ptx_addr.py index a3b8449c4f46..d9c5bdff1fd0 100644 --- a/tests/python/tirx/codegen/test_ptx_addr.py +++ b/tests/python/tirx/codegen/test_ptx_addr.py @@ -86,10 +86,10 @@ def test_ptx_addr_table_validation_rejects_wrong_operand_classes(): def test_ptx_addr_coercion_ir_order_and_shared_codegen(): @T.prim_func - def kernel(global_buf: T.Buffer((8,), "uint64"), raw_shared: T.uint32, raw_global: T.uint64): + def kernel(global_buf: T.Tensor((8,), "uint64"), raw_shared: T.uint32, raw_global: T.uint64): T.device_entry() tx = T.thread_id([32]) - shared_buf = T.alloc_buffer((8,), "uint64", scope="shared") + shared_buf = T.alloc_tensor((8,), "uint64", scope="shared") value = T.local_scalar("uint64") if tx == 0: T.ptx.ld.shared.b64(value, T.ptx.addr(shared_buf.data, 4)) @@ -116,14 +116,14 @@ def kernel(global_buf: T.Buffer((8,), "uint64"), raw_shared: T.uint32, raw_globa def test_ptx_addr_scalar_vector_cache_predicate_and_multi_address_codegen(): @T.prim_func def kernel( - src: T.Buffer((64,), "uint32"), - dst: T.Buffer((64,), "uint32"), - policy: T.Buffer((1,), "uint64"), + src: T.Tensor((64,), "uint32"), + dst: T.Tensor((64,), "uint32"), + policy: T.Tensor((1,), "uint64"), ): T.device_entry() tx = T.thread_id([32]) - shared_buf = T.alloc_buffer((64,), "uint32", scope="shared") - barrier = T.alloc_buffer((1,), "uint64", scope="shared") + shared_buf = T.alloc_tensor((64,), "uint32", scope="shared") + barrier = T.alloc_tensor((1,), "uint64", scope="shared") values = T.alloc_local((2,), "uint32") T.ptx.ld.global_.b32(values[0], T.ptx.addr(src.data, 16)) T.ptx.ld.global_.L2__cache_hint.b32(values[1], T.ptx.addr(src.data, -16), policy[0]) @@ -194,7 +194,7 @@ def test_ptx_addr_zero_sign_boundaries_and_helper_names(): def test_ptx_addr_unrolled_expression_and_dynamic_rejection(): @T.prim_func - def unrolled(src: T.Buffer((16,), "uint32")): + def unrolled(src: T.Tensor((16,), "uint32")): T.device_entry() tx = T.thread_id([32]) value = T.local_scalar("uint32") @@ -208,14 +208,14 @@ def unrolled(src: T.Buffer((16,), "uint32")): assert "ld.global.b32 %0, [%1+32];" in source @T.prim_func - def thread_dynamic(src: T.Buffer((16,), "uint32")): + def thread_dynamic(src: T.Tensor((16,), "uint32")): T.device_entry() tx = T.thread_id([32]) value = T.local_scalar("uint32") T.ptx.ld.global_.b32(value, T.ptx.addr(src.data, tx * 4)) @T.prim_func - def loop_dynamic(src: T.Buffer((16,), "uint32")): + def loop_dynamic(src: T.Tensor((16,), "uint32")): T.device_entry() tx = T.thread_id([32]) value = T.local_scalar("uint32") @@ -253,7 +253,7 @@ def global_u32(raw: T.uint32): with pytest.raises(ValueError, match="does not support T.ptx.addr"): @T.prim_func - def ptr_operand(src: T.Buffer((8,), "uint32")): + def ptr_operand(src: T.Tensor((8,), "uint32")): T.device_entry() result = T.local_scalar("uint32") T.ptx.isspacep.global_(result, T.ptx.addr(src.data, 4)) @@ -263,10 +263,10 @@ def test_ptx_addr_tma_tmem_and_independent_immediate_rejections(): with pytest.raises(ValueError, match="does not support T.ptx.addr"): @T.prim_func - def tma(tmap: T.Buffer((8,), "uint64")): + def tma(tmap: T.Tensor((8,), "uint64")): T.device_entry() - shared_buf = T.alloc_buffer((16,), "uint32", scope="shared") - barrier = T.alloc_buffer((1,), "uint64", scope="shared") + shared_buf = T.alloc_tensor((16,), "uint32", scope="shared") + barrier = T.alloc_tensor((1,), "uint64", scope="shared") T.ptx["cp.async.bulk.tensor.1d.shared::cta.global.mbarrier::complete_tx::bytes"]( shared_buf.data, T.ptx.addr(tmap.data, 16), T.int32(0), barrier.data ) @@ -293,7 +293,7 @@ def tmem(raw: T.uint32): def test_ptx_addr_printer_script_and_json_roundtrip(): @T.prim_func - def kernel(src: T.Buffer((8,), "uint32"), dst: T.Buffer((8,), "uint32")): + def kernel(src: T.Tensor((8,), "uint32"), dst: T.Tensor((8,), "uint32")): T.device_entry() value = T.local_scalar("uint32") T.ptx.ld.global_.b32(value, T.ptx.addr(src.data, -16)) @@ -312,7 +312,7 @@ def test_ptx_addr_legacy_positional_offsets_rejected(): with pytest.raises(ValueError): @T.prim_func - def scalar_load(src: T.Buffer((8,), "uint32")): + def scalar_load(src: T.Tensor((8,), "uint32")): T.device_entry() value = T.local_scalar("uint32") T.ptx.ld.global_.b32(value, src.data, 16) @@ -320,7 +320,7 @@ def scalar_load(src: T.Buffer((8,), "uint32")): with pytest.raises(ValueError): @T.prim_func - def vector_load(src: T.Buffer((8,), "uint32")): + def vector_load(src: T.Tensor((8,), "uint32")): T.device_entry() values = T.alloc_local((2,), "uint32") T.ptx.ld.global_.v2.b32(values[0], values[1], src.data, 16) @@ -328,14 +328,14 @@ def vector_load(src: T.Buffer((8,), "uint32")): with pytest.raises(ValueError): @T.prim_func - def scalar_store(dst: T.Buffer((8,), "uint32")): + def scalar_store(dst: T.Tensor((8,), "uint32")): T.device_entry() T.ptx.st.global_.b32(dst.data, 16, T.uint32(0)) with pytest.raises(ValueError): @T.prim_func - def vector_store(dst: T.Buffer((8,), "uint32")): + def vector_store(dst: T.Tensor((8,), "uint32")): T.device_entry() T.ptx.st.global_.v2.b32(dst.data, 16, T.uint32(0), T.uint32(0)) diff --git a/tests/python/tirx/codegen/test_ptx_dialect.py b/tests/python/tirx/codegen/test_ptx_dialect.py index 20c2638204cd..c23a8cc9832b 100644 --- a/tests/python/tirx/codegen/test_ptx_dialect.py +++ b/tests/python/tirx/codegen/test_ptx_dialect.py @@ -73,7 +73,7 @@ def test_ptx_registration(): def test_ptx_prefetch_codegen(): @T.prim_func - def kernel(A: T.Buffer((32,), "float32")): + def kernel(A: T.Tensor((32,), "float32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -87,7 +87,7 @@ def kernel(A: T.Buffer((32,), "float32")): def test_ptx_ld_st_codegen(): @T.prim_func - def kernel(A: T.Buffer((32,), "uint32"), B: T.Buffer((32,), "uint32")): + def kernel(A: T.Tensor((32,), "uint32"), B: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -112,7 +112,7 @@ def test_ptx_ld_s32_wide_destination_codegen(): """A signed scalar load may sign-extend into a wider destination register.""" @T.prim_func - def kernel(A: T.Buffer((32,), "int32"), Out: T.Buffer((32,), "int64")): + def kernel(A: T.Tensor((32,), "int32"), Out: T.Tensor((32,), "int64")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -129,11 +129,11 @@ def kernel(A: T.Buffer((32,), "int32"), Out: T.Buffer((32,), "int64")): def test_ptx_st_shared_coercion(): @T.prim_func - def kernel(out: T.Buffer((1,), "uint32")): + def kernel(out: T.Tensor((1,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) - smem = T.alloc_buffer((4,), "uint32", scope="shared") + smem = T.alloc_tensor((4,), "uint32", scope="shared") if tx == 0: # Shared-space slot fed a shared-scope pointer: engine must # auto-wrap with cvta_generic_to_shared. @@ -148,11 +148,11 @@ def kernel(out: T.Buffer((1,), "uint32")): def test_ptx_explicit_cvta(): @T.prim_func - def kernel(out: T.Buffer((1,), "uint64")): + def kernel(out: T.Tensor((1,), "uint64")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) - smem = T.alloc_buffer((4,), "uint32", scope="shared") + smem = T.alloc_tensor((4,), "uint32", scope="shared") smem[tx % 4] = T.uint32(0) if tx == 0: T.ptx.cvta.to.shared.u64(out[0], smem.data) @@ -164,7 +164,7 @@ def kernel(out: T.Buffer((1,), "uint64")): def test_ptx_red_codegen(): @T.prim_func - def kernel(A: T.Buffer((1,), "uint32")): + def kernel(A: T.Tensor((1,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -178,7 +178,7 @@ def kernel(A: T.Buffer((1,), "uint32")): def test_ptx_predication_codegen(): @T.prim_func - def kernel(A: T.Buffer((32,), "uint32")): + def kernel(A: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -195,7 +195,7 @@ def kernel(A: T.Buffer((32,), "uint32")): def test_ptx_red_vector_codegen_and_roundtrip(): @T.prim_func - def kernel(A: T.Buffer((16,), "uint32")): + def kernel(A: T.Tensor((16,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -229,7 +229,7 @@ def kernel(A: T.Buffer((16,), "uint32")): with pytest.raises(ValueError, match="already a 32-bit pair"): @T.prim_func - def packed_v8(A: T.Buffer((1,), "uint32")): + def packed_v8(A: T.Tensor((1,), "uint32")): T.device_entry() v = T.local_scalar("uint32") T.ptx.red.global_.add.noftz.v8.f16x2(A.ptr_to([0]), v, v, v, v, v, v, v, v) @@ -237,7 +237,7 @@ def packed_v8(A: T.Buffer((1,), "uint32")): with pytest.raises(AttributeError, match="not a valid modifier"): @T.prim_func - def f32_max(A: T.Buffer((1,), "float32")): + def f32_max(A: T.Tensor((1,), "float32")): T.device_entry() v = T.local_scalar("float32") T.ptx.red.global_.max.v2.f32(A.ptr_to([0]), v, v) @@ -245,7 +245,7 @@ def f32_max(A: T.Buffer((1,), "float32")): def test_ptx_atom_bitbucket_codegen_and_roundtrip(): @T.prim_func - def kernel(A: T.Buffer((16,), "uint32")): + def kernel(A: T.Tensor((16,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -281,7 +281,7 @@ def kernel(A: T.Buffer((16,), "uint32")): @requires_nvcc def test_ptx_predicated_destination_preserves_old_value(): @T.prim_func - def kernel(A: T.Buffer((1,), "float32"), Out: T.Buffer((32,), "float32")): + def kernel(A: T.Tensor((1,), "float32"), Out: T.Tensor((32,), "float32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -301,7 +301,7 @@ def kernel(A: T.Buffer((1,), "float32"), Out: T.Buffer((32,), "float32")): @requires_nvcc def test_ptx_predicated_destination_is_undefined_by_default(): @T.prim_func - def kernel(A: T.Buffer((1,), "float32"), Out: T.Buffer((32,), "float32")): + def kernel(A: T.Tensor((1,), "float32"), Out: T.Tensor((32,), "float32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -323,7 +323,7 @@ def test_ptx_string_form_matches_chain(): def make(fn): @T.prim_func - def kernel(A: T.Buffer((32,), "uint32"), B: T.Buffer((32,), "uint32")): + def kernel(A: T.Tensor((32,), "uint32"), B: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -628,10 +628,10 @@ def test_ptx_92_cp_reduce_negative_grids(): def test_ptx_92_cp_bulk_roundtrip(): @T.prim_func - def kernel(src: T.Buffer((64,), "uint32"), dst: T.Buffer((64,), "uint32")): + def kernel(src: T.Tensor((64,), "uint32"), dst: T.Tensor((64,), "uint32")): T.device_entry() - smem = T.alloc_buffer((64,), "uint32", scope="shared") - mbar = T.alloc_buffer((2,), "uint64", scope="shared") + smem = T.alloc_tensor((64,), "uint32", scope="shared") + mbar = T.alloc_tensor((2,), "uint64", scope="shared") T.ptx[ "cp.async.bulk.shared::cta.global." "mbarrier::complete_tx::bytes.L2::cache_hint.ignore_oob" @@ -662,7 +662,7 @@ def test_ptx_trace_time_errors(): with pytest.raises(ValueError, match="shared state space"): @T.prim_func - def bad_global_addr(out: T.Buffer((1,), "uint32")): + def bad_global_addr(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.ld.global_.b32(out[0], T.uint32(0)) @@ -670,7 +670,7 @@ def bad_global_addr(out: T.Buffer((1,), "uint32")): with pytest.raises(AttributeError, match="not a valid modifier"): @T.prim_func - def bad_modifier(out: T.Buffer((1,), "uint32")): + def bad_modifier(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.ld.global_.bogus.b32(out[0], T.uint32(0)) @@ -680,7 +680,7 @@ def bad_modifier(out: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="must have dtype"): @T.prim_func - def bad_value_dtype(A: T.Buffer((1,), "uint32")): + def bad_value_dtype(A: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.st.global_.u32(A.ptr_to([0]), T.float64(1.0)) @@ -688,7 +688,7 @@ def bad_value_dtype(A: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="missing required modifier"): @T.prim_func - def missing_type(A: T.Buffer((1,), "uint32")): + def missing_type(A: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.st.global_(A.ptr_to([0]), T.uint32(0)) @@ -697,7 +697,7 @@ def missing_type(A: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="requires a scope"): @T.prim_func - def acquire_without_scope(out: T.Buffer((1,), "uint32"), A: T.Buffer((1,), "uint32")): + def acquire_without_scope(out: T.Tensor((1,), "uint32"), A: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.ld.global_.acquire.b32(out[0], A.ptr_to([0])) @@ -705,7 +705,7 @@ def acquire_without_scope(out: T.Buffer((1,), "uint32"), A: T.Buffer((1,), "uint with pytest.raises(ValueError, match="pointer or uint64 handle"): @T.prim_func - def bad_addr_dtype(out: T.Buffer((1,), "uint32")): + def bad_addr_dtype(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.ld.global_.b32(out[0], T.float32(0)) @@ -717,7 +717,7 @@ def test_ptx_destination_errors(): with pytest.raises(ValueError, match="must have dtype"): @T.prim_func - def wrong_dst_dtype(out: T.Buffer((1,), "float64"), A: T.Buffer((1,), "uint32")): + def wrong_dst_dtype(out: T.Tensor((1,), "float64"), A: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.ld.global_.u32(out[0], A.ptr_to([0])) @@ -726,7 +726,7 @@ def wrong_dst_dtype(out: T.Buffer((1,), "float64"), A: T.Buffer((1,), "uint32")) with pytest.raises(ValueError, match="writable scalar"): @T.prim_func - def let_destination(out: T.Buffer((1,), "uint32"), A: T.Buffer((1,), "uint32")): + def let_destination(out: T.Tensor((1,), "uint32"), A: T.Tensor((1,), "uint32")): T.device_entry() bound: T.let = out[0] + T.uint32(1) T.ptx.ld.global_.b32(bound, A.ptr_to([0])) @@ -735,7 +735,7 @@ def let_destination(out: T.Buffer((1,), "uint32"), A: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="writable scalar"): @T.prim_func - def rvalue_destination(A: T.Buffer((1,), "uint32")): + def rvalue_destination(A: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.ld.global_.b32(T.uint32(0), A.ptr_to([0])) @@ -749,7 +749,7 @@ def test_ptx_register_group_codegen(): """ @T.prim_func - def kernel(A: T.Buffer((4,), "float32")): + def kernel(A: T.Tensor((4,), "float32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -775,7 +775,7 @@ def test_ptx_register_group_errors(): with pytest.raises(ValueError, match=r"expects \d+ operand"): @T.prim_func - def wrong_arity(A: T.Buffer((4,), "uint32")): + def wrong_arity(A: T.Tensor((4,), "uint32")): T.device_entry() packed = T.local_scalar("uint64") T.ptx.mov.b64(packed, A[0]) @@ -786,7 +786,7 @@ def wrong_arity(A: T.Buffer((4,), "uint32")): with pytest.raises(ValueError, match="writable scalar"): @T.prim_func - def non_lvalue_lane(A: T.Buffer((4,), "uint32")): + def non_lvalue_lane(A: T.Tensor((4,), "uint32")): T.device_entry() packed = T.local_scalar("uint64") lo = T.local_scalar("uint32") @@ -799,7 +799,7 @@ def non_lvalue_lane(A: T.Buffer((4,), "uint32")): with pytest.raises(ValueError, match="must have one dtype"): @T.prim_func - def mixed_lane_dtypes(A: T.Buffer((4,), "uint32")): + def mixed_lane_dtypes(A: T.Tensor((4,), "uint32")): T.device_entry() packed = T.local_scalar("uint64") f = T.local_scalar("float32") @@ -812,7 +812,7 @@ def mixed_lane_dtypes(A: T.Buffer((4,), "uint32")): with pytest.raises(ValueError, match="is ambiguous"): @T.prim_func - def bare_float_literal(A: T.Buffer((4,), "uint32")): + def bare_float_literal(A: T.Tensor((4,), "uint32")): T.device_entry() packed = T.local_scalar("uint64") T.ptx.mov.b64(packed, 1.5, 2.5) @@ -820,7 +820,7 @@ def bare_float_literal(A: T.Buffer((4,), "uint32")): # An explicit constant is accepted and picks the float32 helper. @T.prim_func - def typed_literal(A: T.Buffer((4,), "float32")): + def typed_literal(A: T.Tensor((4,), "float32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -842,11 +842,11 @@ def test_ptx_optional_operand_arity_dispatch(): """ @T.prim_func - def kernel(A: T.Buffer((4,), "uint32")): + def kernel(A: T.Tensor((4,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) - bar = T.alloc_buffer((2,), "uint64", scope="shared") + bar = T.alloc_tensor((2,), "uint64", scope="shared") T.ptx.bar.sync(T.uint32(0)) T.ptx.bar.sync(T.uint32(0), T.uint32(64)) T.ptx.mbarrier.arrive.shared.b64(bar.ptr_to([0])) @@ -874,11 +874,11 @@ def test_ptx_mbarrier_92_shapes_render_and_roundtrip(): """PTX 9.2 noComplete sink/state and wait shapes render exactly.""" @T.prim_func - def kernel(out: T.Buffer((4,), "uint32")): + def kernel(out: T.Tensor((4,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) - bar = T.alloc_buffer((2,), "uint64", scope="shared") + bar = T.alloc_tensor((2,), "uint64", scope="shared") state = T.local_scalar("uint64") pending = T.local_scalar("uint32") wait_complete = T.local_scalar("uint32") @@ -961,7 +961,7 @@ def test_ptx_bit_width_axis(): """ @T.prim_func - def kernel(A: T.Buffer((4,), "float32"), Out: T.Buffer((4,), "float32")): + def kernel(A: T.Tensor((4,), "float32"), Out: T.Tensor((4,), "float32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -995,7 +995,7 @@ def test_ptx_relaxed_load_store_typing(): """Scalar/vector ld, st and ldu accept ISA section 9.4.1's wider register carriers.""" @T.prim_func - def kernel(A: T.Buffer((8,), "uint64")): + def kernel(A: T.Tensor((8,), "uint64")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -1029,7 +1029,7 @@ def kernel(A: T.Buffer((8,), "uint64")): with pytest.raises(ValueError, match="must have dtype"): @T.prim_func - def floating_source_for_integer_type(A: T.Buffer((1,), "uint16")): + def floating_source_for_integer_type(A: T.Tensor((1,), "uint16")): T.device_entry() T.ptx.st.global_.u16(A.ptr_to([0]), T.float32(1)) @@ -1078,7 +1078,7 @@ def test_ptx_vec256_cache_policy(): ) @T.prim_func - def vec256_calls(A: T.Buffer((8,), "uint32")): + def vec256_calls(A: T.Tensor((8,), "uint32")): T.device_entry() policy = T.local_scalar("uint64") x0 = T.local_scalar("uint32") @@ -1144,7 +1144,7 @@ def test_ptx_st_bulk_size_carriers_and_st_async_byte_bridge(): assert (signed_value & 0xFF) == bits @T.prim_func - def carrier_calls(A: T.Buffer((8,), "uint8")): + def carrier_calls(A: T.Tensor((8,), "uint8")): T.device_entry() T.cta_id([1]) T.thread_id([1]) @@ -1188,7 +1188,7 @@ def test_ptx_integer_arithmetic_dispatch(): """ @T.prim_func - def kernel(A: T.Buffer((4,), "int32")): + def kernel(A: T.Tensor((4,), "int32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -1228,7 +1228,7 @@ def kernel(A: T.Buffer((4,), "int32")): with pytest.raises(ValueError, match="must have dtype int64"): @T.prim_func - def narrow_wide_dst(out: T.Buffer((1,), "int32")): + def narrow_wide_dst(out: T.Tensor((1,), "int32")): T.device_entry() T.ptx.mul.wide.s32(out[0], T.int32(2), T.int32(3)) @@ -1237,7 +1237,7 @@ def narrow_wide_dst(out: T.Buffer((1,), "int32")): with pytest.raises(ValueError, match="hi.s32"): @T.prim_func - def sat_on_lo(out: T.Buffer((1,), "int32")): + def sat_on_lo(out: T.Tensor((1,), "int32")): T.device_entry() T.ptx.mad.lo.sat.s32(out[0], T.int32(2), T.int32(3), T.int32(4)) @@ -1248,21 +1248,21 @@ def sat_on_lo(out: T.Buffer((1,), "int32")): with pytest.raises(ValueError, match="sm_120f-only"): @T.prim_func - def add_sat_sm120(out: T.Buffer((1,), "uint32")): + def add_sat_sm120(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.add.sat.u32(out[0], T.uint32(2), T.uint32(3)) with pytest.raises(ValueError, match=r"not on the add\.u64"): @T.prim_func - def add_sat_no_syntax_line(out: T.Buffer((1,), "uint64")): + def add_sat_no_syntax_line(out: T.Tensor((1,), "uint64")): T.device_entry() T.ptx.add.sat.u64(out[0], T.uint64(2), T.uint64(3)) with pytest.raises(ValueError, match=r"not on the sub\.u32"): @T.prim_func - def sub_sat_no_syntax_line(out: T.Buffer((1,), "uint32")): + def sub_sat_no_syntax_line(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.sub.sat.u32(out[0], T.uint32(2), T.uint32(3)) @@ -1278,7 +1278,7 @@ def test_ptx_floating_point_dispatch(): """ @T.prim_func - def kernel(A: T.Buffer((4,), "float32")): + def kernel(A: T.Tensor((4,), "float32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -1334,7 +1334,7 @@ def kernel(A: T.Buffer((4,), "float32")): with pytest.raises(ValueError, match="only on the .f32 line"): @T.prim_func - def sqrt_approx_f64(out: T.Buffer((1,), "float64")): + def sqrt_approx_f64(out: T.Tensor((1,), "float64")): T.device_entry() T.ptx.sqrt.approx.f64(out[0], T.float64(2.0)) @@ -1342,7 +1342,7 @@ def sqrt_approx_f64(out: T.Buffer((1,), "float64")): with pytest.raises(ValueError, match="rcp.approx.ftz.f64"): @T.prim_func - def rcp_approx_f64(out: T.Buffer((1,), "float64")): + def rcp_approx_f64(out: T.Tensor((1,), "float64")): T.device_entry() T.ptx.rcp.approx.f64(out[0], T.float64(2.0)) @@ -1350,7 +1350,7 @@ def rcp_approx_f64(out: T.Buffer((1,), "float64")): with pytest.raises(ValueError, match="missing required modifier"): @T.prim_func - def div_without_mode(out: T.Buffer((1,), "float32")): + def div_without_mode(out: T.Tensor((1,), "float32")): T.device_entry() T.ptx.div.f32(out[0], T.float32(1.0), T.float32(2.0)) @@ -1365,7 +1365,7 @@ def test_ptx_half_precision_dispatch(): """ @T.prim_func - def kernel(A: T.Buffer((4,), "uint32")): + def kernel(A: T.Tensor((4,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -1411,7 +1411,7 @@ def kernel(A: T.Buffer((4,), "uint32")): with pytest.raises(ValueError, match="oob line spells no"): @T.prim_func - def oob_with_ftz(out: T.Buffer((1,), "uint16")): + def oob_with_ftz(out: T.Tensor((1,), "uint16")): T.device_entry() T.ptx.fma.rn.oob.ftz.f16(out[0], T.uint16(0), T.uint16(0), T.uint16(0)) @@ -1419,7 +1419,7 @@ def oob_with_ftz(out: T.Buffer((1,), "uint16")): with pytest.raises(ValueError, match="separate syntax lines"): @T.prim_func - def sat_and_relu(out: T.Buffer((1,), "uint16")): + def sat_and_relu(out: T.Tensor((1,), "uint16")): T.device_entry() T.ptx.fma.rn.sat.relu.f16(out[0], T.uint16(0), T.uint16(0), T.uint16(0)) @@ -1427,7 +1427,7 @@ def sat_and_relu(out: T.Buffer((1,), "uint16")): with pytest.raises(ValueError, match="mandatorily"): @T.prim_func - def ex2_bf16_no_ftz(out: T.Buffer((1,), "uint16")): + def ex2_bf16_no_ftz(out: T.Tensor((1,), "uint16")): T.device_entry() T.ptx.ex2.approx.bf16(out[0], T.uint16(0)) @@ -1435,7 +1435,7 @@ def ex2_bf16_no_ftz(out: T.Buffer((1,), "uint16")): with pytest.raises(ValueError, match="takes no .ftz"): @T.prim_func - def abs_bf16_ftz(out: T.Buffer((1,), "uint16")): + def abs_bf16_ftz(out: T.Tensor((1,), "uint16")): T.device_entry() T.ptx.abs.ftz.bf16(out[0], T.uint16(0)) @@ -1455,7 +1455,7 @@ def test_ptx_mixed_precision_dispatch(): """ @T.prim_func - def kernel(A: T.Buffer((4,), "float32")): + def kernel(A: T.Tensor((4,), "float32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -1487,7 +1487,7 @@ def kernel(A: T.Buffer((4,), "float32")): with pytest.raises(ValueError, match="takes no .ftz"): @T.prim_func - def mixed_with_ftz(out: T.Buffer((1,), "float32")): + def mixed_with_ftz(out: T.Tensor((1,), "float32")): T.device_entry() T.ptx.add.rn.ftz.f32.f16(out[0], T.uint16(0), T.float32(0)) @@ -1495,7 +1495,7 @@ def mixed_with_ftz(out: T.Buffer((1,), "float32")): with pytest.raises(ValueError, match="only exists on the .f32"): @T.prim_func - def mixed_f64(out: T.Buffer((1,), "float64")): + def mixed_f64(out: T.Tensor((1,), "float64")): T.device_entry() T.ptx.add.rn.f64.f16(out[0], T.uint16(0), T.float64(0)) @@ -1504,7 +1504,7 @@ def mixed_f64(out: T.Buffer((1,), "float64")): with pytest.raises(AttributeError, match="not a valid modifier"): @T.prim_func - def mul_mixed(out: T.Buffer((1,), "float32")): + def mul_mixed(out: T.Tensor((1,), "float32")): T.device_entry() T.ptx.mul.rn.f32.f16(out[0], T.uint16(0), T.float32(0)) @@ -1520,7 +1520,7 @@ def test_ptx_comparison_selection_dispatch(): """ @T.prim_func - def kernel(A: T.Buffer((4,), "int32")): + def kernel(A: T.Tensor((4,), "int32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -1565,7 +1565,7 @@ def kernel(A: T.Buffer((4,), "int32")): with pytest.raises(ValueError, match="signed"): @T.prim_func - def unsigned_op_on_signed(out: T.Buffer((1,), "uint32")): + def unsigned_op_on_signed(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.setp.lo.s32(out[0], T.int32(1), T.int32(2)) @@ -1573,7 +1573,7 @@ def unsigned_op_on_signed(out: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="floating-point comparison"): @T.prim_func - def float_op_on_integer(out: T.Buffer((1,), "uint32")): + def float_op_on_integer(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.setp.nan.u32(out[0], T.uint32(1), T.uint32(2)) @@ -1581,7 +1581,7 @@ def float_op_on_integer(out: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="only with eq/ne"): @T.prim_func - def ordered_on_bitsize(out: T.Buffer((1,), "uint32")): + def ordered_on_bitsize(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.setp.lt.b32(out[0], T.uint32(1), T.uint32(2)) @@ -1589,7 +1589,7 @@ def ordered_on_bitsize(out: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="only to .f32 comparisons"): @T.prim_func - def ftz_off_f32(out: T.Buffer((1,), "uint32")): + def ftz_off_f32(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.setp.eq.ftz.s32(out[0], T.int32(1), T.int32(2)) @@ -1597,7 +1597,7 @@ def ftz_off_f32(out: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="f32 selector line"): @T.prim_func - def ftz_on_s32_selector(out: T.Buffer((1,), "uint32")): + def ftz_on_s32_selector(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.slct.ftz.b32.s32(out[0], T.uint32(1), T.uint32(2), T.int32(3)) @@ -1607,7 +1607,7 @@ def ftz_on_s32_selector(out: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match=r"operand 'c'.*int32.*uint32"): @T.prim_func - def relaxed_selector(out: T.Buffer((1,), "uint32")): + def relaxed_selector(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.slct.f32.s32(out[0], T.float32(1), T.int32(2), T.uint32(3)) @@ -1662,7 +1662,7 @@ def test_ptx_half_comparison_dispatch(): """ @T.prim_func - def kernel(A: T.Buffer((4,), "uint32")): + def kernel(A: T.Tensor((4,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -1706,7 +1706,7 @@ def kernel(A: T.Buffer((4,), "uint32")): with pytest.raises(ValueError, match="no half-precision syntax"): @T.prim_func - def bad_pair(out: T.Buffer((1,), "uint16")): + def bad_pair(out: T.Tensor((1,), "uint16")): T.device_entry() T.ptx.set.eq.u16.f16x2(out[0], T.uint32(0), T.uint32(0)) @@ -1714,7 +1714,7 @@ def bad_pair(out: T.Buffer((1,), "uint16")): with pytest.raises(ValueError, match="does not take it"): @T.prim_func - def ftz_bf16_source(out: T.Buffer((1,), "uint32")): + def ftz_bf16_source(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.set.eq.ftz.u32.bf16(out[0], T.uint16(0), T.uint16(0)) @@ -1722,14 +1722,14 @@ def ftz_bf16_source(out: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="destination takes no"): @T.prim_func - def ftz_bf16_dest(out: T.Buffer((1,), "uint16")): + def ftz_bf16_dest(out: T.Tensor((1,), "uint16")): T.device_entry() T.ptx.set.eq.ftz.bf16.f32(out[0], T.float32(0), T.float32(0)) with pytest.raises(ValueError, match="spells no .ftz"): @T.prim_func - def setp_ftz_bf16(out: T.Buffer((1,), "uint32")): + def setp_ftz_bf16(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.setp.eq.ftz.bf16(out[0], T.uint16(0), T.uint16(0)) @@ -1737,7 +1737,7 @@ def setp_ftz_bf16(out: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="integer source"): @T.prim_func - def unordered_on_integer(out: T.Buffer((1,), "uint16")): + def unordered_on_integer(out: T.Tensor((1,), "uint16")): T.device_entry() T.ptx.set.equ.f16.s32(out[0], T.int32(0), T.int32(0)) @@ -1746,7 +1746,7 @@ def unordered_on_integer(out: T.Buffer((1,), "uint16")): with pytest.raises(ValueError, match="bit-size source"): @T.prim_func - def ordered_on_bitsize(out: T.Buffer((1,), "uint16")): + def ordered_on_bitsize(out: T.Tensor((1,), "uint16")): T.device_entry() T.ptx.set.lt.f16.b16(out[0], T.uint16(0), T.uint16(0)) @@ -1754,7 +1754,7 @@ def ordered_on_bitsize(out: T.Buffer((1,), "uint16")): with pytest.raises(AttributeError, match="not a valid modifier"): @T.prim_func - def unsigned_alternate(out: T.Buffer((1,), "uint32")): + def unsigned_alternate(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.setp.lo.f16(out[0], T.uint16(0), T.uint16(0)) @@ -1773,7 +1773,7 @@ def test_ptx_logic_shift_dispatch(): """ @T.prim_func - def kernel(A: T.Buffer((4,), "uint32")): + def kernel(A: T.Tensor((4,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -1825,7 +1825,7 @@ def kernel(A: T.Buffer((4,), "uint32")): with pytest.raises(AttributeError, match="not a valid modifier"): @T.prim_func - def cnot_pred(out: T.Buffer((1,), "uint32")): + def cnot_pred(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.cnot.pred(out[0], T.uint32(0)) @@ -1833,7 +1833,7 @@ def cnot_pred(out: T.Buffer((1,), "uint32")): with pytest.raises(AttributeError, match="not a valid modifier"): @T.prim_func - def lop3_xor(out: T.Buffer((1,), "uint32")): + def lop3_xor(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.lop3.xor.b32(out[0], T.uint32(0), T.uint32(0), T.uint32(0), 0x80) @@ -1842,7 +1842,7 @@ def lop3_xor(out: T.Buffer((1,), "uint32")): with pytest.raises(AttributeError, match="not a valid modifier"): @T.prim_func - def shl_signed(out: T.Buffer((1,), "int32")): + def shl_signed(out: T.Tensor((1,), "int32")): T.device_entry() T.ptx.shl.s32(out[0], T.int32(1), T.uint32(2)) @@ -1850,7 +1850,7 @@ def shl_signed(out: T.Buffer((1,), "int32")): # can specialize, but a runtime LUT byte still has no register form and is # rejected at CUDA codegen. @T.prim_func - def lut_runtime(A: T.Buffer((1,), "uint32")): + def lut_runtime(A: T.Tensor((1,), "uint32")): T.device_entry() tx = T.thread_id([32]) if tx == 0: @@ -1861,7 +1861,7 @@ def lut_runtime(A: T.Buffer((1,), "uint32")): _cuda_source(lut_runtime) @T.prim_func - def lut_unrolled(A: T.Buffer((1,), "uint32")): + def lut_unrolled(A: T.Tensor((1,), "uint32")): T.device_entry() tx = T.thread_id([32]) if tx == 0: @@ -1874,7 +1874,7 @@ def lut_unrolled(A: T.Buffer((1,), "uint32")): assert "lop3.b32 %0, %1, %2, %3, 128;" in unrolled_src @T.prim_func - def lut_boundaries(A: T.Buffer((1,), "uint32")): + def lut_boundaries(A: T.Tensor((1,), "uint32")): T.device_entry() T.cta_id([1]) T.thread_id([1]) @@ -1889,7 +1889,7 @@ def lut_boundaries(A: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match=r"inclusive range 0\.\.255, got -1"): @T.prim_func - def lut_below_range(A: T.Buffer((1,), "uint32")): + def lut_below_range(A: T.Tensor((1,), "uint32")): T.device_entry() d = T.local_scalar("uint32") T.ptx.lop3.b32(d, A[0], A[0], A[0], -1) @@ -1897,13 +1897,13 @@ def lut_below_range(A: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match=r"inclusive range 0\.\.255, got 256"): @T.prim_func - def lut_above_range(A: T.Buffer((1,), "uint32")): + def lut_above_range(A: T.Tensor((1,), "uint32")): T.device_entry() d = T.local_scalar("uint32") T.ptx.lop3.b32(d, A[0], A[0], A[0], 256) @T.prim_func - def lut_unrolled_out_of_range(A: T.Buffer((1,), "uint32")): + def lut_unrolled_out_of_range(A: T.Tensor((1,), "uint32")): T.device_entry() T.cta_id([1]) T.thread_id([1]) @@ -1926,11 +1926,11 @@ def test_ptx_data_movement_dispatch(): """ @T.prim_func - def kernel(A: T.Buffer((8,), "uint32")): + def kernel(A: T.Tensor((8,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) - smem = T.alloc_buffer((4,), "uint32", scope="shared") + smem = T.alloc_tensor((4,), "uint32", scope="shared") d = T.local_scalar("uint32") p = T.local_scalar("uint32") v = T.local_scalar("uint32") @@ -1998,7 +1998,7 @@ def kernel(A: T.Buffer((8,), "uint32")): with pytest.raises(ValueError, match="go together"): @T.prim_func - def sem_without_scope(A: T.Buffer((1,), "uint32")): + def sem_without_scope(A: T.Tensor((1,), "uint32")): T.device_entry() v = T.local_scalar("uint32") T.ptx.multimem_ld_reduce.relaxed.add.u32(v, A.ptr_to([0])) @@ -2007,7 +2007,7 @@ def sem_without_scope(A: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="takes no scope"): @T.prim_func - def weak_with_scope(A: T.Buffer((1,), "uint32")): + def weak_with_scope(A: T.Tensor((1,), "uint32")): T.device_entry() v = T.local_scalar("uint32") T.ptx.multimem_ld_reduce.weak.gpu.add.u32(v, A.ptr_to([0])) @@ -2016,7 +2016,7 @@ def weak_with_scope(A: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match=r"\.add takes"): @T.prim_func - def add_s64(A: T.Buffer((1,), "int64")): + def add_s64(A: T.Tensor((1,), "int64")): T.device_entry() v = T.local_scalar("int64") T.ptx.multimem_ld_reduce.add.s64(v, A.ptr_to([0])) @@ -2025,7 +2025,7 @@ def add_s64(A: T.Buffer((1,), "int64")): with pytest.raises(ValueError, match="the scalar line takes"): @T.prim_func - def scalar_f16(A: T.Buffer((1,), "uint16")): + def scalar_f16(A: T.Tensor((1,), "uint16")): T.device_entry() v = T.local_scalar("uint16") T.ptx.multimem_ld_reduce.add.f16(v, A.ptr_to([0])) @@ -2034,7 +2034,7 @@ def scalar_f16(A: T.Buffer((1,), "uint16")): with pytest.raises(ValueError, match="applies to .add"): @T.prim_func - def acc_on_min(A: T.Buffer((1,), "uint32")): + def acc_on_min(A: T.Tensor((1,), "uint32")): T.device_entry() v = T.local_scalar("uint32") T.ptx.multimem_ld_reduce.min.acc__f32.f16x2(v, A.ptr_to([0])) @@ -2043,7 +2043,7 @@ def acc_on_min(A: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="requires .sys"): @T.prim_func - def mmio_gpu(A: T.Buffer((1,), "uint32")): + def mmio_gpu(A: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.st_async.mmio.release.gpu.global_.u32(A.ptr_to([0]), T.uint32(1)) @@ -2059,7 +2059,7 @@ def test_ptx_parallel_sync_dispatch(): """ @T.prim_func - def kernel(A: T.Buffer((8,), "uint32")): + def kernel(A: T.Tensor((8,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -2114,7 +2114,7 @@ def kernel(A: T.Buffer((8,), "uint32")): with pytest.raises(AttributeError, match="not a valid modifier"): @T.prim_func - def redux_add_b32(out: T.Buffer((1,), "uint32")): + def redux_add_b32(out: T.Tensor((1,), "uint32")): T.device_entry() T.ptx.redux_sync.add.b32(out[0], T.uint32(1), T.uint32(0xFFFFFFFF)) @@ -2123,7 +2123,7 @@ def redux_add_b32(out: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match="already a 32-bit pair"): @T.prim_func - def packed_v8(A: T.Buffer((1,), "uint32")): + def packed_v8(A: T.Tensor((1,), "uint32")): T.device_entry() v = T.local_scalar("uint32") T.ptx.atom.global_.add.noftz.v8.f16x2( @@ -2134,7 +2134,7 @@ def packed_v8(A: T.Buffer((1,), "uint32")): with pytest.raises(ValueError, match=r"\.and takes"): @T.prim_func - def red_async_and_u64(A: T.Buffer((1,), "uint64")): + def red_async_and_u64(A: T.Tensor((1,), "uint64")): T.device_entry() v = T.local_scalar("uint64") T.ptx.red_async.relaxed.cluster.shared__cluster.mbarrier__complete_tx__bytes.and_.u64( @@ -2146,11 +2146,11 @@ def test_ptx_lazy_subscript_operands_realize(): """Raw buffer elements realize before PTX predicate operand validation.""" @T.prim_func - def kernel(A: T.Buffer((1,), "uint32")): + def kernel(A: T.Tensor((1,), "uint32")): T.device_entry() T.cta_id([1]) T.thread_id([32]) - regs = T.alloc_buffer((2,), "uint32", scope="local") + regs = T.alloc_tensor((2,), "uint32", scope="local") full = T.uint32(0xFFFFFFFF) T.ptx.elect_sync(regs[0], regs[1], full) T.ptx.vote_sync.all.pred(regs[0], T.ptx.pred(regs[1]), full) @@ -2167,11 +2167,11 @@ def test_ptx_parser_roundtrip(): """script() output re-parses to a structurally equal PrimFunc.""" @T.prim_func - def kernel(A: T.Buffer((32,), "uint32"), B: T.Buffer((32,), "uint32")): + def kernel(A: T.Tensor((32,), "uint32"), B: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) - smem = T.alloc_buffer((4,), "uint32", scope="shared") + smem = T.alloc_tensor((4,), "uint32", scope="shared") if tx == 0: val = T.local_scalar("uint32") smem_addr = T.local_scalar("uint64") @@ -2207,7 +2207,7 @@ def test_ptx_pred_operand_roundtrip(): """ @T.prim_func - def kernel(A: T.Buffer((32,), "uint32")): + def kernel(A: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -2231,7 +2231,7 @@ def test_ptx_wgmma_scale_d_runtime_predicate_roundtrip(): """WGMMA scale-d is a runtime predicate, not a 0/1 text immediate.""" @T.prim_func - def kernel(Out: T.Buffer((128,), "uint32")): + def kernel(Out: T.Tensor((128,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([128]) @@ -2397,7 +2397,7 @@ def test_ptx_codegen_rejects_stale_table_layout(mismatch): ) def test_ptx_tcgen05_mma_block_size_form(): @T.prim_func - def kernel(A: T.Buffer((32,), "uint32")): + def kernel(A: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -2438,7 +2438,7 @@ def test_ptx_tcgen05_mma_block_size_collector_form(): """PTX 9.4 collector qualifiers on block-scaled MMA certify at their sm_107f floor.""" @T.prim_func - def sm107_collector_kernel(A: T.Buffer((32,), "uint32")): + def sm107_collector_kernel(A: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -2482,7 +2482,7 @@ def test_ptx_tcgen05_mma_block_scale_collector_a_without_block_size(): """SM107 activation-stationary FP8 accepts collector A without `.block*`.""" @T.prim_func - def kernel(A: T.Buffer((32,), "uint32")): + def kernel(A: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -2594,7 +2594,7 @@ def untagged_integer(): # A bool expression carries the class in its own dtype, so it needs no tag. @T.prim_func - def bool_needs_no_tag(A: T.Buffer((32,), "uint32")): + def bool_needs_no_tag(A: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -2620,7 +2620,7 @@ def test_ptx_sink_lane_codegen_and_roundtrip(): """ @T.prim_func - def kernel(A: T.Buffer((32,), "uint32")): + def kernel(A: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -2670,7 +2670,7 @@ def sink_every_lane(): def test_ptx_printer_form(): @T.prim_func - def kernel(A: T.Buffer((32,), "uint32"), B: T.Buffer((32,), "uint32")): + def kernel(A: T.Tensor((32,), "uint32"), B: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) @@ -2975,10 +2975,10 @@ def test_ptx_94_cp_bulk_semantics_roundtrip(): """The string namespace dispatches every widened PTX 9.4 sibling.""" @T.prim_func - def kernel(src: T.Buffer((64,), "uint32")): + def kernel(src: T.Tensor((64,), "uint32")): T.device_entry() - smem = T.alloc_buffer((64,), "uint32", scope="shared") - mbar = T.alloc_buffer((2,), "uint64", scope="shared") + smem = T.alloc_tensor((64,), "uint32", scope="shared") + mbar = T.alloc_tensor((2,), "uint64", scope="shared") T.ptx[ "cp.async.bulk.relaxed.sys.shared::cta.global." "mbarrier::complete_tx::bytes." @@ -3188,8 +3188,8 @@ def test_ptx_tcgen05_mapa_address_coercion(): raw_u64 = tvm.tirx.Var("a64", "uint64") ncols = tvm.tirx.Var("n", "uint32") mask = tvm.tirx.Var("mask", "uint16") - out32 = tvm.tirx.decl_buffer((1,), "uint32", name="out32", scope="local") - out64 = tvm.tirx.decl_buffer((1,), "uint64", name="out64", scope="local") + out32 = tvm.tirx.decl_tensor((1,), "uint32", name="out32", scope="local") + out64 = tvm.tirx.decl_tensor((1,), "uint64", name="out64", scope="local") bare_alloc = T.ptx.tcgen05.alloc.cta_group__1.sync.aligned.b32(generic_ptr, ncols) assert bare_alloc.args[0].same_as(generic_ptr) @@ -3221,7 +3221,7 @@ def test_ptx_tcgen05_mapa_address_roundtrip(): """The generic/shared split remains exact through TVMScript print and parse.""" @T.prim_func - def kernel(generic: T.Buffer((4,), "uint64")): + def kernel(generic: T.Tensor((4,), "uint64")): T.device_entry() mapped32 = T.local_scalar("uint32") mapped64 = T.local_scalar("uint64") @@ -3846,7 +3846,7 @@ def test_ptx_coercion_ir_forms(): call = T.ptx.st.release.gpu.global_.b32(global_ptr, val, pred=flag) assert call.args[2].same_as(flag) assert len(call.args) == 2 + 1 + 8 + 1 # operands + pred + slot tokens + marker - out = tvm.tirx.decl_buffer((1,), "uint32", name="out", scope="local") + out = tvm.tirx.decl_tensor((1,), "uint32", name="out", scope="local") call = T.ptx.ld.global_.b32(out[0], global_ptr, pred=flag) assert str(call.args[-1]).strip('"') == "pred" call = T.ptx.ld.global_.b32(out[0], global_ptr, pred=flag, preserve_dst=True) @@ -4551,11 +4551,11 @@ def test_ptx_all_helpers_certify(shard): @requires_nvcc def test_ptx_nvcc_smoke(): @T.prim_func - def kernel(A: T.Buffer((32,), "uint32"), B: T.Buffer((32,), "uint32")): + def kernel(A: T.Tensor((32,), "uint32"), B: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) - smem = T.alloc_buffer((4,), "uint32", scope="shared") + smem = T.alloc_tensor((4,), "uint32", scope="shared") if tx == 0: val = T.local_scalar("uint32") T.ptx.ld.global_.acquire.gpu.b32(val, A.ptr_to([0])) @@ -4572,7 +4572,7 @@ def kernel(A: T.Buffer((32,), "uint32"), B: T.Buffer((32,), "uint32")): @pytest.mark.skipif(not env.has_cuda(), reason="CUDA GPU not available") def test_ptx_ld_st_gpu_roundtrip(): @T.prim_func - def kernel(A: T.Buffer((32,), "uint32"), B: T.Buffer((32,), "uint32")): + def kernel(A: T.Tensor((32,), "uint32"), B: T.Tensor((32,), "uint32")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) diff --git a/tests/python/tirx/codegen/test_ptx_ld_st_ops.py b/tests/python/tirx/codegen/test_ptx_ld_st_ops.py index f6bcec1783ff..d22183e8074c 100644 --- a/tests/python/tirx/codegen/test_ptx_ld_st_ops.py +++ b/tests/python/tirx/codegen/test_ptx_ld_st_ops.py @@ -75,13 +75,13 @@ def _shared_scratch_copy_kernel(num_bytes: int): ld_chain, st_chain = f"ld.shared.{tail}", f"st.shared.{tail}" @T.prim_func - def func(out: T.Buffer((nelems,), smem_dtype)): + def func(out: T.Tensor((nelems,), smem_dtype)): T.device_entry() T.cta_id([1]) T.warp_id([1]) lane = T.lane_id([32]) - src_buf = T.alloc_buffer((nelems,), smem_dtype, scope="shared") - dst_buf = T.alloc_buffer((nelems,), smem_dtype, scope="shared") + src_buf = T.alloc_tensor((nelems,), smem_dtype, scope="shared") + dst_buf = T.alloc_tensor((nelems,), smem_dtype, scope="shared") tmp = T.alloc_local((lanes,), reg_dtype) if fill_offset is not None: if lane < nelems: @@ -110,14 +110,14 @@ def test_ptx_ld_st_codegen_emits_shared_asm(): # fmt: off @T.prim_func - def copy_kernel(D: T.Buffer((4,), 'uint32')) -> None: + def copy_kernel(D: T.Tensor((4,), 'uint32')) -> None: T.device_entry() T.warp_id([4]) T.cta_id([1]) T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - smem = T.alloc_buffer((4,), "uint32", scope="shared") + smem = T.alloc_tensor((4,), "uint32", scope="shared") reg = T.alloc_local((4,), "uint32") if tid_in_wg == 0: T.ptx.st.shared.v4.u32(smem.ptr_to([0]), reg[0], reg[1], reg[2], reg[3]) @@ -139,10 +139,10 @@ def copy_kernel(D: T.Buffer((4,), 'uint32')) -> None: def test_ptx_ld_st_raw_shared_address_codegen(): @T.prim_func - def main(out: T.Buffer((2,), "uint64")): + def main(out: T.Tensor((2,), "uint64")): T.device_entry() tx = T.thread_id([32]) - smem = T.alloc_buffer((2,), "uint64", scope="shared") + smem = T.alloc_tensor((2,), "uint64", scope="shared") values = T.alloc_local((4,), "uint32") if tx == 0: raw_addr: T.uint32 = T.cuda.cvta_generic_to_shared(smem.data) @@ -165,7 +165,7 @@ def test_ptx_ld_st_immediate_offset_codegen(): """An immediate displacement must stay inside the PTX memory operand.""" @T.prim_func - def main(src: T.Buffer((4,), "uint64"), out: T.Buffer((4,), "uint64")): + def main(src: T.Tensor((4,), "uint64"), out: T.Tensor((4,), "uint64")): T.device_entry() tx = T.thread_id([32]) values = T.alloc_local((2,), "uint64") @@ -188,7 +188,7 @@ def test_ptx_ld_global_nc_v8_codegen(): """FlashMLA index loads need ``ld.global.nc`` with a 256B prefetch.""" @T.prim_func - def copy_kernel(src: T.Buffer((8,), "int32"), out: T.Buffer((8,), "int32")) -> None: + def copy_kernel(src: T.Tensor((8,), "int32"), out: T.Tensor((8,), "int32")) -> None: T.device_entry() tx = T.thread_id([32]) tmp = T.alloc_local((8,), "int32") @@ -215,7 +215,7 @@ def test_ptx_ld_global_nc_v4_u64_256b_codegen(): """FlashMLA 32-byte index loads may use four 64-bit PTX outputs.""" @T.prim_func - def copy_kernel(src: T.Buffer((4,), "uint64"), out: T.Buffer((4,), "uint64")) -> None: + def copy_kernel(src: T.Tensor((4,), "uint64"), out: T.Tensor((4,), "uint64")) -> None: T.device_entry() tx = T.thread_id([32]) tmp = T.alloc_local((4,), "uint64") @@ -238,7 +238,7 @@ def test_ptx_ld_vector_scatter_dst_codegen(): """Vector loads may write independent destination pointers.""" @T.prim_func - def copy_kernel(src: T.Buffer((4,), "int32"), out: T.Buffer((4,), "int32")) -> None: + def copy_kernel(src: T.Tensor((4,), "int32"), out: T.Tensor((4,), "int32")) -> None: T.device_entry() tx = T.thread_id([32]) tmp0 = T.alloc_local((1,), "int32") diff --git a/tests/python/tirx/codegen/test_tcgen05_descriptor_encoders.py b/tests/python/tirx/codegen/test_tcgen05_descriptor_encoders.py index 92e2adfe0447..31fb5e295b3d 100644 --- a/tests/python/tirx/codegen/test_tcgen05_descriptor_encoders.py +++ b/tests/python/tirx/codegen/test_tcgen05_descriptor_encoders.py @@ -69,11 +69,11 @@ def test_smem_descriptor_matches_runtime_encoder(ldo, sdo, swizzle): """`base | (addr >> 4)` must reproduce the C bitfield fill exactly.""" @T.prim_func - def kernel(out: T.Buffer((2,), "uint64")): + def kernel(out: T.Tensor((2,), "uint64")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) - smem = T.alloc_buffer((64,), "uint32", scope="shared") + smem = T.alloc_tensor((64,), "uint32", scope="shared") smem[tx] = T.uint32(0) smem[tx + 32] = T.uint32(0) if tx == 0: @@ -112,7 +112,7 @@ def test_instr_descriptor_block_scaled_matches_runtime_encoder( m, n, k, a_dtype, b_dtype, trans_a, trans_b, cta_group ): @T.prim_func - def kernel(out: T.Buffer((1,), "uint64")): + def kernel(out: T.Tensor((1,), "uint64")): T.device_entry() T.cta_id([1]) tx = T.thread_id([32]) diff --git a/tests/python/tirx/iket/iket_profile_workload.py b/tests/python/tirx/iket/iket_profile_workload.py index 5cdc6a9a3865..ab178f352877 100644 --- a/tests/python/tirx/iket/iket_profile_workload.py +++ b/tests/python/tirx/iket/iket_profile_workload.py @@ -35,7 +35,7 @@ @T.prim_func -def canonical_iket_workload(out: T.Buffer((32,), "int32")): +def canonical_iket_workload(out: T.Tensor((32,), "int32")): T.device_entry() profiler = iket.IketProfiler() tx = T.thread_id([32]) @@ -51,7 +51,7 @@ def canonical_iket_workload(out: T.Buffer((32,), "int32")): @T.prim_func -def native_payload_workload(out: T.Buffer((32,), "int32")): +def native_payload_workload(out: T.Tensor((32,), "int32")): T.device_entry() profiler = iket.IketProfiler() tx = T.thread_id([32]) @@ -72,7 +72,7 @@ def native_payload_workload(out: T.Buffer((32,), "int32")): @T.prim_func -def extended_payload_workload(out: T.Buffer((32,), "int32")): +def extended_payload_workload(out: T.Tensor((32,), "int32")): T.device_entry() profiler = iket.IketProfiler() tx = T.thread_id([32]) diff --git a/tests/python/tirx/iket/test_iket_profiler.py b/tests/python/tirx/iket/test_iket_profiler.py index f7aa5e719fc2..ce7b8c9b4465 100644 --- a/tests/python/tirx/iket/test_iket_profiler.py +++ b/tests/python/tirx/iket/test_iket_profiler.py @@ -42,7 +42,7 @@ @T.prim_func -def serial_a(out: T.Buffer((32,), "int32")): +def serial_a(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -51,7 +51,7 @@ def serial_a(out: T.Buffer((32,), "int32")): @T.prim_func -def serial_b(out: T.Buffer((32,), "int32")): +def serial_b(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -60,14 +60,14 @@ def serial_b(out: T.Buffer((32,), "int32")): @T.prim_func -def plain_entry(out: T.Buffer((32,), "int32")): +def plain_entry(out: T.Tensor((32,), "int32")): T.device_entry() tx = T.thread_id([32]) out[tx] = tx + 7 @T.prim_func -def push_pop_kernel(out: T.Buffer((32,), "int32")): +def push_pop_kernel(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -80,7 +80,7 @@ def push_pop_kernel(out: T.Buffer((32,), "int32")): @T.prim_func -def token_loop(n: T.int32, out: T.Buffer((32,), "int32")): +def token_loop(n: T.int32, out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -96,7 +96,7 @@ def token_loop(n: T.int32, out: T.Buffer((32,), "int32")): @T.prim_func -def payload_kernel(out: T.Buffer((32,), "int32")): +def payload_kernel(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -105,7 +105,7 @@ def payload_kernel(out: T.Buffer((32,), "int32")): @T.prim_func -def payload_types(n: T.int64, out: T.Buffer((32,), "int32")): +def payload_types(n: T.int64, out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -128,7 +128,7 @@ def payload_types(n: T.int64, out: T.Buffer((32,), "int32")): @T.prim_func -def payload_presence_mismatch(out: T.Buffer((32,), "int32")): +def payload_presence_mismatch(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -138,7 +138,7 @@ def payload_presence_mismatch(out: T.Buffer((32,), "int32")): @T.prim_func -def payload_type_mismatch(out: T.Buffer((32,), "int32")): +def payload_type_mismatch(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -148,7 +148,7 @@ def payload_type_mismatch(out: T.Buffer((32,), "int32")): @T.prim_func -def sentinel_only_payload(out: T.Buffer((32,), "int32")): +def sentinel_only_payload(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -158,7 +158,7 @@ def sentinel_only_payload(out: T.Buffer((32,), "int32")): @T.prim_func -def payload_float16(out: T.Buffer((32,), "int32")): +def payload_float16(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -167,7 +167,7 @@ def payload_float16(out: T.Buffer((32,), "int32")): @T.prim_func -def payload_bfloat16(out: T.Buffer((32,), "int32")): +def payload_bfloat16(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -176,7 +176,7 @@ def payload_bfloat16(out: T.Buffer((32,), "int32")): @T.prim_func -def payload_pointer(out: T.Buffer((32,), "int32")): +def payload_pointer(out: T.Tensor((32,), "int32")): T.device_entry() tx = T.thread_id([32]) T.evaluate(tvm.tirx.call_intrin("", "tirx.cuda.iket_mark", "bad", out.data)) @@ -184,13 +184,13 @@ def payload_pointer(out: T.Buffer((32,), "int32")): @T.prim_func -def payload_vector(out: T.Buffer((1,), "int32x4")): +def payload_vector(out: T.Tensor((1,), "int32x4")): T.device_entry() T.evaluate(tvm.tirx.call_intrin("", "tirx.cuda.iket_mark", "bad", out[0])) @T.prim_func -def schema_i32(out: T.Buffer((32,), "int32")): +def schema_i32(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -199,7 +199,7 @@ def schema_i32(out: T.Buffer((32,), "int32")): @T.prim_func -def schema_u32(out: T.Buffer((32,), "int32")): +def schema_u32(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -208,7 +208,7 @@ def schema_u32(out: T.Buffer((32,), "int32")): @T.prim_func -def schema_no_payload(out: T.Buffer((32,), "int32")): +def schema_no_payload(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -217,7 +217,7 @@ def schema_no_payload(out: T.Buffer((32,), "int32")): @T.prim_func -def marks_30(out: T.Buffer((32,), "int32")): +def marks_30(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -255,7 +255,7 @@ def marks_30(out: T.Buffer((32,), "int32")): @T.prim_func -def marks_31(out: T.Buffer((32,), "int32")): +def marks_31(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) @@ -324,7 +324,7 @@ def _event_bytes(source, name): def _many_marks(count): marks = "\n".join(f' iket.mark("e{index:04d}")' for index in range(count)) source = f"""@T.prim_func -def main(out: T.Buffer((1,), "int32")): +def main(out: T.Tensor((1,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([1]) @@ -411,7 +411,7 @@ def test_regular_lowering_strips_annotations_and_tokens(): def make_kernel(with_annotation): @T.prim_func - def main(out: T.Buffer((32,), "int32")): + def main(out: T.Tensor((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_fallback.py b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_fallback.py index 9cad5a7c3b71..305a0da0e459 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_fallback.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_fallback.py @@ -77,12 +77,12 @@ def _build_round_trip_kernel(scope, n_threads, shape, dtype): if scope == "warp": @T.prim_func - def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def kernel(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.lane_id([32]) T.thread_id([n_threads]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) Tx.warp.copy(A_smem[full], A[full]) T.cuda.cta_sync() Tx.warp.copy(B[full], A_smem[full]) @@ -90,7 +90,7 @@ def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: elif scope == "warpgroup": @T.prim_func - def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def kernel(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.warpgroup_id([n_threads // 128]) @@ -98,7 +98,7 @@ def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: T.lane_id([32]) T.thread_id_in_wg([128]) T.thread_id([n_threads]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) Tx.wg.copy(A_smem[full], A[full]) T.cuda.cta_sync() Tx.wg.copy(B[full], A_smem[full]) @@ -106,13 +106,13 @@ def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: elif scope == "cta": @T.prim_func - def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def kernel(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.warp_id([n_threads // 32]) T.lane_id([32]) T.thread_id([n_threads]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) Tx.cta.copy(A_smem[full], A[full]) T.cuda.cta_sync() Tx.cta.copy(B[full], A_smem[full]) @@ -172,11 +172,11 @@ def test_fallback_thread_scope(): full = tuple(slice(0, d) for d in shape) @T.prim_func - def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def kernel(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.thread_id([1]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) Tx.copy(A_smem[full], A[full]) T.cuda.cta_sync() Tx.copy(B[full], A_smem[full]) @@ -210,13 +210,13 @@ def test_fallback_emits_gate(): full = tuple(slice(0, d) for d in shape) @T.prim_func - def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def kernel(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.warp_id([8]) # 256 threads => 8 warps T.lane_id([32]) T.thread_id([256]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) Tx.cta.copy(A_smem[full], A[full]) Tx.cta.copy(B[full], A_smem[full]) diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_gmem_smem.py b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_gmem_smem.py index adcfc3ea326d..e7f8d23c8aaa 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_gmem_smem.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_gmem_smem.py @@ -40,7 +40,7 @@ def _build_kernel(scope, n_threads, shape, dtype): if scope == "warpgroup": @T.prim_func - def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def kernel(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.warpgroup_id([n_threads // 128]) @@ -48,7 +48,7 @@ def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: T.lane_id([32]) T.thread_id_in_wg([128]) T.thread_id([n_threads]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) Tx.wg.copy(A_smem[full_slices], A[full_slices]) T.cuda.cta_sync() Tx.wg.copy(B[full_slices], A_smem[full_slices]) @@ -56,12 +56,12 @@ def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: elif scope == "warp": @T.prim_func - def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def kernel(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.lane_id([32]) T.thread_id([n_threads]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) Tx.warp.copy(A_smem[full_slices], A[full_slices]) T.cuda.cta_sync() Tx.warp.copy(B[full_slices], A_smem[full_slices]) @@ -69,13 +69,13 @@ def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: elif scope == "cta": @T.prim_func - def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def kernel(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.warp_id([n_threads // 32]) T.lane_id([32]) T.thread_id([n_threads]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) Tx.cta.copy(A_smem[full_slices], A[full_slices]) T.cuda.cta_sync() Tx.cta.copy(B[full_slices], A_smem[full_slices]) @@ -207,13 +207,13 @@ def test_copy_g2s_s2g(task, dtype, scope): @T.prim_func def copy_sync( - A: T.Buffer(g_shape, dtype, layout=layoutA), B: T.Buffer(g_shape, dtype, layout=layoutB) + A: T.Tensor(g_shape, dtype, layout=layoutA), B: T.Tensor(g_shape, dtype, layout=layoutB) ) -> None: T.device_entry() T.cta_id([2]) T.thread_id([thread_cnt]) - A_smem = T.alloc_buffer(s_shape, dtype, scope="shared", layout=layoutS) + A_smem = T.alloc_tensor(s_shape, dtype, scope="shared", layout=layoutS) # `scope` is parametrized at runtime; select the scope namespace # dynamically (T.cta / T.thread) instead of a literal prefix. getattr(Tx, scope).copy(A_smem[r_smem], A[r_gmem]) @@ -354,7 +354,7 @@ def test_swizzled_smem_emit_must_be_swizzle_aware(): s_layout = ComposeLayout(3, 3, 3, TileLayout(S[shape])) @T.prim_func - def kernel(A: T.Buffer(shape, "float16")) -> None: + def kernel(A: T.Tensor(shape, "float16")) -> None: T.device_entry() T.cta_id([1]) T.warpgroup_id([1]) @@ -362,7 +362,7 @@ def kernel(A: T.Buffer(shape, "float16")) -> None: T.lane_id([32]) T.thread_id_in_wg([128]) T.thread_id([128]) - A_smem = T.alloc_buffer(shape, "float16", scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, "float16", scope="shared", layout=s_layout) Tx.wg.copy(A_smem[0:128, 0:32], A[0:128, 0:32]) # NB: pin sm_90 explicitly — the default cuda target falls back to sm_50 @@ -519,14 +519,14 @@ def test_gmem_smem_swizzle_uses_structured_compose_apply(): @T.prim_func def kernel( - A: T.Buffer(shape, "float16", layout=g_layout), - B: T.Buffer(shape, "float16", layout=g_layout), + A: T.Tensor(shape, "float16", layout=g_layout), + B: T.Tensor(shape, "float16", layout=g_layout), ) -> None: T.device_entry() T.cta_id([1]) T.lane_id([32]) T.thread_id([32]) - smem = T.alloc_buffer(shape, "float16", scope="shared", layout=s_layout) + smem = T.alloc_tensor(shape, "float16", scope="shared", layout=s_layout) Tx.warp.copy(smem, A[:, :]) T.cuda.cta_sync() Tx.warp.copy(B[:, :], smem) diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_ld_stmatrix.py b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_ld_stmatrix.py index ec7c3dbe9832..cb7d5321db75 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_ld_stmatrix.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_ld_stmatrix.py @@ -113,13 +113,13 @@ def _coord(row, cp, t, w): # fmt: off if direction == "ld": @T.prim_func - def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> None: + def kernel(A: T.Tensor((M, N), 'float16'), B: T.Tensor((M, N), 'float16')) -> None: T.device_entry() T.cta_id([1]) T.lane_id([32]) tid = T.thread_id([32]) - A_smem = T.alloc_buffer(s_shape, "float16", scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(s_shape, "float16", scope="shared", layout=s_layout) row = tid // 4 cp = tid % 4 for t in range(num): @@ -127,7 +127,7 @@ def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> No gr, gc = _coord(row, cp, t, w) A_smem[row, cp, t, w] = A[gr, gc] T.cuda.cta_sync() - R_local = T.alloc_buffer(s_shape, "float16", scope="local", layout=r_layout) + R_local = T.alloc_tensor(s_shape, "float16", scope="local", layout=r_layout) Tx.warp.copy(R_local[full], A_smem[full]) r_view = R_local.local() for t in range(num): @@ -136,16 +136,16 @@ def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> No B[gr, gc] = r_view[t * 2 + w] else: # direction == "st" @T.prim_func - def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> None: + def kernel(A: T.Tensor((M, N), 'float16'), B: T.Tensor((M, N), 'float16')) -> None: T.device_entry() T.cta_id([1]) T.lane_id([32]) tid = T.thread_id([32]) - A_smem = T.alloc_buffer(s_shape, "float16", scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(s_shape, "float16", scope="shared", layout=s_layout) row = tid // 4 cp = tid % 4 - R_local = T.alloc_buffer(s_shape, "float16", scope="local", layout=r_layout) + R_local = T.alloc_tensor(s_shape, "float16", scope="local", layout=r_layout) r_view = R_local.local() for t in range(num): for w in range(2): @@ -179,7 +179,7 @@ def _coord(wid, row, cp, t, w): # fmt: off if direction == "ld": @T.prim_func - def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> None: + def kernel(A: T.Tensor((M, N), 'float16'), B: T.Tensor((M, N), 'float16')) -> None: T.device_entry() T.cta_id([1]) @@ -188,7 +188,7 @@ def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> No T.lane_id([32]) T.thread_id_in_wg([128]) tid = T.thread_id([128]) - A_smem = T.alloc_buffer(s_shape, "float16", scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(s_shape, "float16", scope="shared", layout=s_layout) wid = tid // 32 lid = tid % 32 row = lid // 4 @@ -198,7 +198,7 @@ def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> No gr, gc = _coord(wid, row, cp, t, w) A_smem[wid, row, cp, t, w] = A[gr, gc] T.cuda.cta_sync() - R_local = T.alloc_buffer(s_shape, "float16", scope="local", layout=r_layout) + R_local = T.alloc_tensor(s_shape, "float16", scope="local", layout=r_layout) Tx.wg.copy(R_local[full], A_smem[full]) r_view = R_local.local() for t in range(num): @@ -207,7 +207,7 @@ def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> No B[gr, gc] = r_view[t * 2 + w] else: @T.prim_func - def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> None: + def kernel(A: T.Tensor((M, N), 'float16'), B: T.Tensor((M, N), 'float16')) -> None: T.device_entry() T.cta_id([1]) @@ -216,12 +216,12 @@ def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> No T.lane_id([32]) T.thread_id_in_wg([128]) tid = T.thread_id([128]) - A_smem = T.alloc_buffer(s_shape, "float16", scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(s_shape, "float16", scope="shared", layout=s_layout) wid = tid // 32 lid = tid % 32 row = lid // 4 cp = lid % 4 - R_local = T.alloc_buffer(s_shape, "float16", scope="local", layout=r_layout) + R_local = T.alloc_tensor(s_shape, "float16", scope="local", layout=r_layout) r_view = R_local.local() for t in range(num): for w in range(2): @@ -255,14 +255,14 @@ def _coord(wid, row, cp, t, w): # fmt: off if direction == "ld": @T.prim_func - def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> None: + def kernel(A: T.Tensor((M, N), 'float16'), B: T.Tensor((M, N), 'float16')) -> None: T.device_entry() T.cta_id([1]) T.warp_id([4]) T.lane_id([32]) tid = T.thread_id([128]) - A_smem = T.alloc_buffer(s_shape, "float16", scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(s_shape, "float16", scope="shared", layout=s_layout) wid = tid // 32 lid = tid % 32 row = lid // 4 @@ -272,7 +272,7 @@ def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> No gr, gc = _coord(wid, row, cp, t, w) A_smem[wid, row, cp, t, w] = A[gr, gc] T.cuda.cta_sync() - R_local = T.alloc_buffer(s_shape, "float16", scope="local", layout=r_layout) + R_local = T.alloc_tensor(s_shape, "float16", scope="local", layout=r_layout) Tx.cta.copy(R_local[full], A_smem[full]) r_view = R_local.local() for t in range(num): @@ -281,19 +281,19 @@ def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> No B[gr, gc] = r_view[t * 2 + w] else: @T.prim_func - def kernel(A: T.Buffer((M, N), 'float16'), B: T.Buffer((M, N), 'float16')) -> None: + def kernel(A: T.Tensor((M, N), 'float16'), B: T.Tensor((M, N), 'float16')) -> None: T.device_entry() T.cta_id([1]) T.warp_id([4]) T.lane_id([32]) tid = T.thread_id([128]) - A_smem = T.alloc_buffer(s_shape, "float16", scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(s_shape, "float16", scope="shared", layout=s_layout) wid = tid // 32 lid = tid % 32 row = lid // 4 cp = lid % 4 - R_local = T.alloc_buffer(s_shape, "float16", scope="local", layout=r_layout) + R_local = T.alloc_tensor(s_shape, "float16", scope="local", layout=r_layout) r_view = R_local.local() for t in range(num): for w in range(2): @@ -394,19 +394,19 @@ def _build_multi_iter_kernel(outer_ext: int): full = tuple(slice(0, e) for e in shape) @T.prim_func - def kernel(A: T.Buffer(shape, "float16"), B: T.Buffer(shape, "float16")) -> None: + def kernel(A: T.Tensor(shape, "float16"), B: T.Tensor(shape, "float16")) -> None: T.device_entry() T.cta_id([1]) T.lane_id([32]) tid = T.thread_id([32]) - A_smem = T.alloc_buffer(shape, "float16", scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, "float16", scope="shared", layout=s_layout) for a in range(outer_ext): for c in range(2): for d in range(4): for e in range(2): A_smem[a, tid // 4, c, d, tid % 4, e] = A[a, tid // 4, c, d, tid % 4, e] T.cuda.cta_sync() - R_local = T.alloc_buffer(shape, "float16", scope="local", layout=r_layout) + R_local = T.alloc_tensor(shape, "float16", scope="local", layout=r_layout) Tx.warp.copy(R_local[full], A_smem[full]) r_view = R_local.local() for a in range(outer_ext): @@ -473,14 +473,14 @@ def test_ldstmatrix_tcgen05_warpgroup_atom_emits_ldmatrix(): smem_layout = mma_shared_layout("bfloat16", 3, (m, k)) @T.prim_func - def kernel(smem: T.Buffer((m, k), "bfloat16", scope="shared", layout=smem_layout)) -> None: + def kernel(smem: T.Tensor((m, k), "bfloat16", scope="shared", layout=smem_layout)) -> None: T.device_entry() T.cta_id([1]) T.warpgroup_id([1]) T.warp_id_in_wg([4]) T.lane_id([32]) T.thread_id_in_wg([128]) - a_reg = T.alloc_buffer((m, k), "bfloat16", scope="local", layout=reg_layout) + a_reg = T.alloc_tensor((m, k), "bfloat16", scope="local", layout=reg_layout) Tx.wg.copy(a_reg, smem) _, src = _compile_src(kernel) diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_reg.py b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_reg.py index 6384c17709e5..b22030d34ef4 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_reg.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_reg.py @@ -69,7 +69,7 @@ def _build_roundtrip_kernel(scope, n_threads, k, dtype, non_r_scope): if scope == "warpgroup": @T.prim_func - def kernel(B: T.Buffer(shape, dtype)) -> None: + def kernel(B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.warpgroup_id([n_threads // 128]) @@ -77,11 +77,11 @@ def kernel(B: T.Buffer(shape, dtype)) -> None: T.lane_id([32]) T.thread_id_in_wg([128]) tid = T.thread_id([n_threads]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) for kk in range(k): A_smem[tid, kk] = T.cast(tid * 100 + kk + 1, dtype) T.cuda.cta_sync() - R_local = T.alloc_buffer(shape, dtype, scope="local", layout=r_layout) + R_local = T.alloc_tensor(shape, dtype, scope="local", layout=r_layout) Tx.wg.copy(R_local[full_slices], A_smem[full_slices]) for kk in range(k): A_smem[tid, kk] = T.cast(0, dtype) @@ -94,16 +94,16 @@ def kernel(B: T.Buffer(shape, dtype)) -> None: elif scope == "warp": @T.prim_func - def kernel(B: T.Buffer(shape, dtype)) -> None: + def kernel(B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.lane_id([32]) tid = T.thread_id([n_threads]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) for kk in range(k): A_smem[tid, kk] = T.cast(tid * 100 + kk + 1, dtype) T.cuda.cta_sync() - R_local = T.alloc_buffer(shape, dtype, scope="local", layout=r_layout) + R_local = T.alloc_tensor(shape, dtype, scope="local", layout=r_layout) Tx.warp.copy(R_local[full_slices], A_smem[full_slices]) for kk in range(k): A_smem[tid, kk] = T.cast(0, dtype) @@ -116,17 +116,17 @@ def kernel(B: T.Buffer(shape, dtype)) -> None: elif scope == "cta": @T.prim_func - def kernel(B: T.Buffer(shape, dtype)) -> None: + def kernel(B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.warp_id([n_threads // 32]) T.lane_id([32]) tid = T.thread_id([n_threads]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=s_layout) for kk in range(k): A_smem[tid, kk] = T.cast(tid * 100 + kk + 1, dtype) T.cuda.cta_sync() - R_local = T.alloc_buffer(shape, dtype, scope="local", layout=r_layout) + R_local = T.alloc_tensor(shape, dtype, scope="local", layout=r_layout) Tx.cta.copy(R_local[full_slices], A_smem[full_slices]) for kk in range(k): A_smem[tid, kk] = T.cast(0, dtype) @@ -142,7 +142,7 @@ def kernel(B: T.Buffer(shape, dtype)) -> None: if scope == "warpgroup": @T.prim_func - def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def kernel(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.warpgroup_id([n_threads // 128]) @@ -153,7 +153,7 @@ def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: for kk in range(k): A[tid, kk] = T.cast(tid * 100 + kk + 1, dtype) T.cuda.cta_sync() - R_local = T.alloc_buffer(shape, dtype, scope="local", layout=r_layout) + R_local = T.alloc_tensor(shape, dtype, scope="local", layout=r_layout) Tx.wg.copy(R_local[full_slices], A[full_slices]) for kk in range(k): A[tid, kk] = T.cast(0, dtype) @@ -166,7 +166,7 @@ def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: elif scope == "warp": @T.prim_func - def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def kernel(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.lane_id([32]) @@ -174,7 +174,7 @@ def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: for kk in range(k): A[tid, kk] = T.cast(tid * 100 + kk + 1, dtype) T.cuda.cta_sync() - R_local = T.alloc_buffer(shape, dtype, scope="local", layout=r_layout) + R_local = T.alloc_tensor(shape, dtype, scope="local", layout=r_layout) Tx.warp.copy(R_local[full_slices], A[full_slices]) for kk in range(k): A[tid, kk] = T.cast(0, dtype) @@ -187,7 +187,7 @@ def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: elif scope == "cta": @T.prim_func - def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def kernel(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() T.cta_id([1]) T.warp_id([n_threads // 32]) @@ -196,7 +196,7 @@ def kernel(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: for kk in range(k): A[tid, kk] = T.cast(tid * 100 + kk + 1, dtype) T.cuda.cta_sync() - R_local = T.alloc_buffer(shape, dtype, scope="local", layout=r_layout) + R_local = T.alloc_tensor(shape, dtype, scope="local", layout=r_layout) Tx.cta.copy(R_local[full_slices], A[full_slices]) for kk in range(k): A[tid, kk] = T.cast(0, dtype) @@ -277,13 +277,13 @@ def test_reg_roundtrip_gapped_permuted_storage(): # fmt: off @T.prim_func - def kernel(A: T.Buffer(shape, 'float32'), B: T.Buffer(shape, 'float32')) -> None: + def kernel(A: T.Tensor(shape, 'float32'), B: T.Tensor(shape, 'float32')) -> None: T.device_entry() T.cta_id([1]) T.lane_id([32]) T.thread_id([32]) - reg = T.alloc_buffer(shape, "float32", scope="local", layout=r_layout) + reg = T.alloc_tensor(shape, "float32", scope="local", layout=r_layout) Tx.warp.copy(reg, A, dispatch="vec_auto") Tx.warp.copy(B, reg, dispatch="vec_auto") # fmt: on @@ -345,12 +345,12 @@ def test_copy_g2l_l2g_vec_load(task, dtype): @T.prim_func def copy_sync( - A: T.Buffer(g_shape, dtype, layout=layoutA), B: T.Buffer(g_shape, dtype, layout=layoutB) + A: T.Tensor(g_shape, dtype, layout=layoutA), B: T.Tensor(g_shape, dtype, layout=layoutB) ) -> None: T.device_entry() T.cta_id([2]) T.thread_id([thread_cnt]) - A_local = T.alloc_buffer(l_shape, dtype, scope="local", layout=layoutLocal) + A_local = T.alloc_tensor(l_shape, dtype, scope="local", layout=layoutLocal) Tx.copy(A_local[r_lmem], A[r_gmem]) Tx.copy(B[r_gmem], A_local[r_lmem]) @@ -381,7 +381,7 @@ def run_and_check(): @T.prim_func -def _nc_strided_reg_copy(src: T.Buffer((1024,), "int32")) -> None: +def _nc_strided_reg_copy(src: T.Tensor((1024,), "int32")) -> None: T.device_entry() T.thread_id([128]) tid = T.thread_id_in_wg([128]) @@ -396,7 +396,7 @@ def _nc_strided_reg_copy(src: T.Buffer((1024,), "int32")) -> None: @T.prim_func -def _plain_strided_reg_copy(src: T.Buffer((1024,), "int32")) -> None: +def _plain_strided_reg_copy(src: T.Tensor((1024,), "int32")) -> None: T.device_entry() T.thread_id([128]) tid = T.thread_id_in_wg([128]) @@ -443,13 +443,13 @@ def test_reg_copy_linear_shared_hoists_thread_base(): linear_layout = TileLayout(S[shape]) @T.prim_func - def kernel(A: T.Buffer(shape, "float16", layout=linear_layout)) -> None: + def kernel(A: T.Tensor(shape, "float16", layout=linear_layout)) -> None: T.device_entry() T.cta_id([1]) T.thread_id([n_threads]) tid = T.thread_id_in_wg([n_threads]) - reg = T.alloc_buffer(shape, "float16", scope="local", layout=wg_local_layout(width)) - smem = T.alloc_buffer(shape, "float16", scope="shared", layout=linear_layout) + reg = T.alloc_tensor(shape, "float16", scope="local", layout=wg_local_layout(width)) + smem = T.alloc_tensor(shape, "float16", scope="shared", layout=linear_layout) reg_local = reg.local(width) for i in T.serial(width): @@ -499,15 +499,15 @@ def test_reg_copy_wg_local_to_swizzled_shared_uses_structured_compose_apply(): @T.prim_func def kernel( - A: T.Buffer(g_shape, "float16", layout=g_layout), - B: T.Buffer(g_shape, "float16", layout=g_layout), + A: T.Tensor(g_shape, "float16", layout=g_layout), + B: T.Tensor(g_shape, "float16", layout=g_layout), ) -> None: T.device_entry() T.cta_id([1]) T.thread_id([N_THREADS]) tid = T.thread_id_in_wg([N_THREADS]) - reg = T.alloc_buffer(g_shape, "float16", scope="local", layout=wg_local_layout(EPI_N)) - smem = T.alloc_buffer(g_shape, "float16", scope="shared", layout=smem_layout) + reg = T.alloc_tensor(g_shape, "float16", scope="local", layout=wg_local_layout(EPI_N)) + smem = T.alloc_tensor(g_shape, "float16", scope="shared", layout=smem_layout) # Populate the per-thread slice via .local() (decomposes the wg # thread-axis layout into a per-thread 1D view). @@ -554,11 +554,11 @@ def test_ptx_st_from_src_f32_vector_preserves_values(): """A vector store of f32 registers must preserve the values.""" @T.prim_func - def kernel(B: T.Buffer((4,), "float32")) -> None: + def kernel(B: T.Tensor((4,), "float32")) -> None: T.device_entry() T.cta_id([1]) T.thread_id([1]) - smem = T.alloc_buffer((4,), "float32", scope="shared") + smem = T.alloc_tensor((4,), "float32", scope="shared") reg = T.alloc_local((4,), "float32") out = T.alloc_local((4,), "float32") for i in range(4): @@ -586,7 +586,7 @@ def kernel(B: T.Buffer((4,), "float32")) -> None: def test_copy_fallback_handles_scalar_regions(): @T.prim_func - def kernel(B: T.Buffer((1,), "float32")) -> None: + def kernel(B: T.Tensor((1,), "float32")) -> None: T.device_entry() T.cta_id([1]) T.thread_id([1]) @@ -618,11 +618,11 @@ def kernel(B: T.Buffer((1,), "float32")) -> None: ) def test_copy_forced_vec_width_codegen(variant, dtype, n_elements, expected_st, expected_ld): @T.prim_func - def kernel(B: T.Buffer((n_elements,), dtype)) -> None: + def kernel(B: T.Tensor((n_elements,), dtype)) -> None: T.device_entry() T.cta_id([1]) T.thread_id([1]) - smem = T.alloc_buffer((n_elements,), dtype, scope="shared") + smem = T.alloc_tensor((n_elements,), dtype, scope="shared") reg = T.alloc_local((n_elements,), dtype) out = T.alloc_local((n_elements,), dtype) for i in range(n_elements): @@ -657,11 +657,11 @@ def test_copy_forced_vec_dynamic_swizzled_shared_uses_vector_ptx(): smem_layout = ComposeLayout(2, 3, 3, TileLayout(S[(64, 8, 32) : (32, 2048, 1)])) @T.prim_func - def kernel(B: T.Buffer((128, 4), "float32")) -> None: + def kernel(B: T.Tensor((128, 4), "float32")) -> None: T.device_entry() T.cta_id([1]) tid = T.thread_id([128]) - smem = T.alloc_buffer((64, 256), "float32", scope="shared", layout=smem_layout) + smem = T.alloc_tensor((64, 256), "float32", scope="shared", layout=smem_layout) reg = T.alloc_local((4,), "float32") out = T.alloc_local((4,), "float32") for i in range(4): @@ -695,13 +695,13 @@ def kernel(B: T.Buffer((128, 4), "float32")) -> None: @pytest.mark.skipif(not env.has_cuda_compute(9), reason="need cuda compute >= 9.0") def test_copy_explicit_vec_auto_uses_auto_family(): @T.prim_func - def kernel(B: T.Buffer((4,), "float32")) -> None: + def kernel(B: T.Tensor((4,), "float32")) -> None: T.device_entry() T.cta_id([1]) T.thread_id([1]) - smem = T.alloc_buffer((4,), "float32", scope="shared") - reg = T.alloc_buffer((4,), "float32", scope="local", layout=TileLayout(S[4])) - out = T.alloc_buffer((4,), "float32", scope="local", layout=TileLayout(S[4])) + smem = T.alloc_tensor((4,), "float32", scope="shared") + reg = T.alloc_tensor((4,), "float32", scope="local", layout=TileLayout(S[4])) + out = T.alloc_tensor((4,), "float32", scope="local", layout=TileLayout(S[4])) for i in range(4): reg[i] = T.cast(i + 1, "float32") Tx.copy(smem[:], reg[:], dispatch="vec_auto") @@ -726,11 +726,11 @@ def kernel(B: T.Buffer((4,), "float32")) -> None: @pytest.mark.parametrize("dispatch", ["reg", "gmem_smem"]) def test_copy_old_dispatch_names_are_not_registered(dispatch): @T.prim_func - def kernel(B: T.Buffer((4,), "float32")) -> None: + def kernel(B: T.Tensor((4,), "float32")) -> None: T.device_entry() T.cta_id([1]) T.thread_id([1]) - smem = T.alloc_buffer((4,), "float32", scope="shared") + smem = T.alloc_tensor((4,), "float32", scope="shared") reg = T.alloc_local((4,), "float32") Tx.copy(smem[:], reg[:], dispatch=dispatch) B[0] = T.cast(0, "float32") @@ -746,11 +746,11 @@ def kernel(B: T.Buffer((4,), "float32")) -> None: @pytest.mark.skipif(not env.has_cuda_compute(9), reason="need cuda compute >= 9.0") def test_copy_forced_vec_rejects_size_mismatch(): @T.prim_func - def kernel(B: T.Buffer((4,), "float32")) -> None: + def kernel(B: T.Tensor((4,), "float32")) -> None: T.device_entry() T.cta_id([1]) T.thread_id([1]) - smem = T.alloc_buffer((4,), "float32", scope="shared") + smem = T.alloc_tensor((4,), "float32", scope="shared") reg = T.alloc_local((4,), "float32") Tx.copy(smem[:], reg[:], dispatch="vec_64b") B[0] = T.cast(0, "float32") @@ -766,13 +766,13 @@ def kernel(B: T.Buffer((4,), "float32")) -> None: @pytest.mark.skipif(not env.has_cuda_compute(9), reason="need cuda compute >= 9.0") def test_copy_forced_vec_rejects_non_thread_scope(): @T.prim_func - def kernel(B: T.Buffer((4,), "float32")) -> None: + def kernel(B: T.Tensor((4,), "float32")) -> None: T.device_entry() T.cta_id([1]) T.lane_id([32]) T.thread_id([32]) - smem = T.alloc_buffer((4,), "float32", scope="shared") - reg = T.alloc_buffer((4,), "float32", scope="local", layout=TileLayout(S[4])) + smem = T.alloc_tensor((4,), "float32", scope="shared") + reg = T.alloc_tensor((4,), "float32", scope="local", layout=TileLayout(S[4])) Tx.warp.copy(smem[:], reg[:], dispatch="vec_128b") B[0] = T.cast(0, "float32") @@ -963,9 +963,9 @@ def _build_tcgen05_d_epilogue_deposit(): @T.prim_func def deposit( - d_reg: T.Buffer((m, n), _TCGEN05_D_DTYPE, scope="local", layout=reg_layout), + d_reg: T.Tensor((m, n), _TCGEN05_D_DTYPE, scope="local", layout=reg_layout), ) -> None: - smem_cd_mma = T.alloc_buffer((m, n), _TCGEN05_D_DTYPE, scope="shared", layout=smem_layout) + smem_cd_mma = T.alloc_tensor((m, n), _TCGEN05_D_DTYPE, scope="shared", layout=smem_layout) T.device_entry() T.cta_id([1]) T.warpgroup_id([1]) @@ -1074,7 +1074,7 @@ def _build_tcgen05_d_epilogue_deposit_roundtrip(): @T.prim_func def kernel( - A: T.Buffer((m, n), _TCGEN05_D_DTYPE), B: T.Buffer((m, n), _TCGEN05_D_DTYPE) + A: T.Tensor((m, n), _TCGEN05_D_DTYPE), B: T.Tensor((m, n), _TCGEN05_D_DTYPE) ) -> None: T.device_entry() T.cta_id([1]) @@ -1083,9 +1083,9 @@ def kernel( T.lane_id([32]) tid_wg = T.thread_id_in_wg([128]) lane = T.lane_id([32]) - d_reg = T.alloc_buffer((m, n), _TCGEN05_D_DTYPE, scope="local", layout=reg_layout) - d_reg_out = T.alloc_buffer((m, n), _TCGEN05_D_DTYPE, scope="local", layout=reg_layout) - smem_cd_mma = T.alloc_buffer((m, n), _TCGEN05_D_DTYPE, scope="shared", layout=smem_layout) + d_reg = T.alloc_tensor((m, n), _TCGEN05_D_DTYPE, scope="local", layout=reg_layout) + d_reg_out = T.alloc_tensor((m, n), _TCGEN05_D_DTYPE, scope="local", layout=reg_layout) + smem_cd_mma = T.alloc_tensor((m, n), _TCGEN05_D_DTYPE, scope="shared", layout=smem_layout) reg_in = d_reg.local(regs_per_thread) reg_out = d_reg_out.local(regs_per_thread) for r in T.serial(regs_per_thread): diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_vec_forced_cache.py b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_vec_forced_cache.py index b795105d8271..ef0e5cf579fa 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_vec_forced_cache.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_vec_forced_cache.py @@ -39,7 +39,7 @@ def _build_g2l2g_kernel(n_elements, dtype, dispatch, **copy_config): given config), then reg → B (global) via a plain forced-vec copy.""" @T.prim_func - def kernel(A: T.Buffer((n_elements,), dtype), B: T.Buffer((n_elements,), dtype)) -> None: + def kernel(A: T.Tensor((n_elements,), dtype), B: T.Tensor((n_elements,), dtype)) -> None: T.device_entry() T.cta_id([1]) T.thread_id([1]) @@ -55,11 +55,11 @@ def _build_g2s2g_kernel(n_elements, dtype, dispatch, **copy_config): local tmp inside the dispatch), then smem → B elementwise.""" @T.prim_func - def kernel(A: T.Buffer((n_elements,), dtype), B: T.Buffer((n_elements,), dtype)) -> None: + def kernel(A: T.Tensor((n_elements,), dtype), B: T.Tensor((n_elements,), dtype)) -> None: T.device_entry() T.cta_id([1]) T.thread_id([1]) - smem = T.alloc_buffer((n_elements,), dtype, scope="shared") + smem = T.alloc_tensor((n_elements,), dtype, scope="shared") Tx.copy(smem[:], A[:], dispatch=dispatch, **copy_config) for i in range(n_elements): B[i] = smem[i] @@ -147,11 +147,11 @@ def test_copy_vec_128b_nc_global_to_shared(): @pytest.mark.skipif(not env.has_cuda_compute(9), reason="need cuda compute >= 9.0") def test_copy_vec_nc_rejects_non_global_src(): @T.prim_func - def kernel(B: T.Buffer((4,), "float32")) -> None: + def kernel(B: T.Tensor((4,), "float32")) -> None: T.device_entry() T.cta_id([1]) T.thread_id([1]) - smem = T.alloc_buffer((4,), "float32", scope="shared") + smem = T.alloc_tensor((4,), "float32", scope="shared") reg = T.alloc_local((4,), "float32") Tx.copy(reg[:], smem[:], dispatch="vec_128b", cache="nc") B[0] = reg[0] diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_dsmem.py b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_dsmem.py index ef4d619fe590..87136a49ae3e 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_dsmem.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_dsmem.py @@ -46,8 +46,8 @@ def _make_dsmem_dispatch_call(shape, dtype, src_layout, dst_layout): from tvm.ir import Range from tvm.tirx.stmt import BufferRegion - src_buf = tvm.tirx.decl_buffer(shape, dtype, "A", scope="shared.dyn", layout=src_layout) - dst_buf = tvm.tirx.decl_buffer(shape, dtype, "B", scope="shared.dyn", layout=dst_layout) + src_buf = tvm.tirx.decl_tensor(shape, dtype, "A", scope="shared.dyn", layout=src_layout) + dst_buf = tvm.tirx.decl_tensor(shape, dtype, "B", scope="shared.dyn", layout=dst_layout) ranges = [Range.from_min_extent(0, s) for s in shape] config = {"mbar": Var("mbar", "handle"), "remote_cta_id": IntImm("int32", 1)} op_call = CopyAsync(BufferRegion(dst_buf, ranges), BufferRegion(src_buf, ranges), config=config) @@ -160,7 +160,7 @@ def test_dsmem(shape, dtype, src_spec, dst_spec, expected): # fmt: off @T.prim_func - def dsmem_copy(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: + def dsmem_copy(A: T.Tensor(shape, dtype), B: T.Tensor(shape, dtype)) -> None: T.device_entry() cbx = T.cta_id_in_cluster([CLUSTER_N]) @@ -169,13 +169,13 @@ def dsmem_copy(A: T.Buffer(shape, dtype), B: T.Buffer(shape, dtype)) -> None: pool = T.SMEMPool() # src_smem: CTA 0 writes here, dispatch reads from here src_raw = pool.alloc([src_phys], dtype, align=128) - src_smem = T.decl_buffer( + src_smem = T.decl_tensor( list(shape), dtype, src_raw.data, elem_offset=0, scope="shared.dyn", layout=src_layout, ) # dst_smem: dispatch writes here (on remote CTA), CTA 1 reads dst_raw = pool.alloc([dst_phys], dtype, align=128) - dst_smem = T.decl_buffer( + dst_smem = T.decl_tensor( list(shape), dtype, dst_raw.data, elem_offset=0, scope="shared.dyn", layout=dst_layout, ) @@ -233,7 +233,7 @@ def test_dsmem_dispatch_missing_config(): from tvm.tirx.stmt import BufferRegion layout = TileLayout(S[64]) - buf = tvm.tirx.decl_buffer((64,), "float16", "A", scope="shared.dyn", layout=layout) + buf = tvm.tirx.decl_tensor((64,), "float16", "A", scope="shared.dyn", layout=layout) br = BufferRegion(buf, [Range.from_min_extent(0, 64)]) target = tvm.target.Target({"kind": "cuda", "arch": "sm_90a"}) sctx = DispatchContext(target, ExecScope("thread"), {}, {}) diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_ldgsts.py b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_ldgsts.py index ee6982debcc0..3293f4ed3ba1 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_ldgsts.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_ldgsts.py @@ -80,13 +80,13 @@ def test_copy_g2s_s2g_cta_vec_load(task, dtype): # fmt: off @T.prim_func def copy_async( - A: T.Buffer(g_shape, dtype, layout=layoutA), B: T.Buffer(g_shape, dtype, layout=layoutB) + A: T.Tensor(g_shape, dtype, layout=layoutA), B: T.Tensor(g_shape, dtype, layout=layoutB) ) -> None: T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([thread_cnt]) - A_smem = T.alloc_buffer(s_shape, dtype, scope="shared", layout=layoutS) + A_smem = T.alloc_tensor(s_shape, dtype, scope="shared", layout=layoutS) Tx.cta.copy_async(A_smem[tuple(r_smem)], A[tuple(r_gmem)], dispatch="ldgsts") T.ptx.cp.async_.commit_group() @@ -123,10 +123,10 @@ def test_copy_ldgsts_predicate_zero_fill_codegen(): """ldgsts direct mode forwards predicate/zero-fill/prefetch without partition temps.""" @T.prim_func - def copy_async(A: T.Buffer((32, 16), "uint8", layout=TileLayout(S[32, 16]))) -> None: + def copy_async(A: T.Tensor((32, 16), "uint8", layout=TileLayout(S[32, 16]))) -> None: T.device_entry() tid = T.thread_id([32]) - A_smem = T.alloc_buffer((32, 16), "uint8", scope="shared", layout=TileLayout(S[32, 16])) + A_smem = T.alloc_tensor((32, 16), "uint8", scope="shared", layout=TileLayout(S[32, 16])) Tx.copy_async( A_smem[tid, :], diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tcgen05_cp.py b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tcgen05_cp.py index 5df0aede0fb2..73aaf5e9fc96 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tcgen05_cp.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tcgen05_cp.py @@ -150,13 +150,13 @@ def _make_cp_kernel( cfg.update(extra_cfg) @T.prim_func(check_well_formed=False) - def kernel(A: T.Buffer(s_full_shape, dtype), B: T.Buffer((128, W32), "uint32")): + def kernel(A: T.Tensor(s_full_shape, dtype), B: T.Tensor((128, W32), "uint32")): T.device_entry() warp_id = T.warp_id([4]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) T.lane_id([32]) - A_smem = T.alloc_buffer(s_full_shape, dtype, scope="shared", layout=s_full, align=1024) + A_smem = T.alloc_tensor(s_full_shape, dtype, scope="shared", layout=s_full, align=1024) tmem_addr = T.alloc_shared([1], "uint32") cp_mbar = T.alloc_shared([1], "uint64") if wg_id == 0: @@ -170,11 +170,11 @@ def kernel(A: T.Buffer(s_full_shape, dtype), B: T.Buffer((128, W32), "uint32")): T.cuda.cta_sync() Tx.cta.copy(A_smem[s_full_sl], A[s_full_sl]) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( t_full_shape, dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=t_full ) if pre_zero: - zero_reg = T.alloc_buffer((W32,), "uint32", scope="local") + zero_reg = T.alloc_tensor((W32,), "uint32", scope="local") for i in range(W32): zero_reg[i] = T.uint32(0) for i in range(W32): @@ -194,7 +194,7 @@ def kernel(A: T.Buffer(s_full_shape, dtype), B: T.Buffer((128, W32), "uint32")): T.ptx.tcgen05.fence__after_thread_sync() # Each of the 4 warps reads its own 32-lane slab (taddr lane 0 is # warp-slab-relative for .32x32b), covering all 128 TMEM lanes. - reg = T.alloc_buffer((W32,), "uint32", scope="local") + reg = T.alloc_tensor((W32,), "uint32", scope="local") for i in range(W32): T.ptx["tcgen05.ld.sync.aligned.32x32b.x1.b32"]( reg[i], T.cuda.get_tmem_addr(tmem_addr[0], 0, i) @@ -473,14 +473,14 @@ def _make_cp_kernel_cta2(s_full, s_shape, t_full, t_shape, dtype, cfg, W32, n_co t_sl = tuple(slice(0, e) for e in t_shape) @T.prim_func(check_well_formed=False) - def kernel(A: T.Buffer((2, *s_shape), dtype), B: T.Buffer((256, W32), "uint32")): + def kernel(A: T.Tensor((2, *s_shape), dtype), B: T.Tensor((256, W32), "uint32")): T.device_entry() warp_id = T.warp_id([4]) cbx, cby = T.cta_id_in_cluster([2, 1]) cta_id = T.cta_id([2]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(s_shape, dtype, scope="shared", layout=s_full, align=1024) + A_smem = T.alloc_tensor(s_shape, dtype, scope="shared", layout=s_full, align=1024) tmem_addr = T.alloc_shared([1], "uint32") cp_mbar = T.alloc_shared([1], "uint64") if tid_in_wg == 0: @@ -489,7 +489,7 @@ def kernel(A: T.Buffer((2, *s_shape), dtype), B: T.Buffer((256, W32), "uint32")) T.ptx.tcgen05.alloc.cta_group__2.sync.aligned.shared__cta.b32( T.address_of(tmem_addr), T.uint32(n_cols) ) - tmem = T.decl_buffer( + tmem = T.decl_tensor( t_shape, dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=t_full ) T.ptx.fence.mbarrier_init.release.cluster() @@ -498,7 +498,7 @@ def kernel(A: T.Buffer((2, *s_shape), dtype), B: T.Buffer((256, W32), "uint32")) T.cuda.cluster_sync() # Pre-zero both CTAs' tmem: alloc does not clear it, and the # 128x256b test asserts the odd CTA stays untouched. - zero_reg = T.alloc_buffer((W32,), "uint32", scope="local") + zero_reg = T.alloc_tensor((W32,), "uint32", scope="local") for i in range(W32): zero_reg[i] = T.uint32(0) for i in range(W32): @@ -518,7 +518,7 @@ def kernel(A: T.Buffer((2, *s_shape), dtype), B: T.Buffer((256, W32), "uint32")) T.cuda.mbarrier_wait(cp_mbar.ptr_to([0]), 0) T.cuda.cta_sync() T.ptx.tcgen05.fence__after_thread_sync() - reg = T.alloc_buffer((W32,), "uint32", scope="local") + reg = T.alloc_tensor((W32,), "uint32", scope="local") for i in range(W32): T.ptx["tcgen05.ld.sync.aligned.32x32b.x1.b32"]( reg[i], T.cuda.get_tmem_addr(tmem_addr[0], 0, i) @@ -641,13 +641,13 @@ def test_cp_default_32x128b_instruction_sequence_unchanged(): t_full = TileLayout(S[(4, 32, 16) : (16 @ TCol, 1 @ TLane, 1 @ TCol)] + R[4 : 32 @ TLane]) @T.prim_func(check_well_formed=False) - def kernel(A: T.Buffer((4, 32, 16), "uint8")): + def kernel(A: T.Tensor((4, 32, 16), "uint8")): T.device_entry() warp_id = T.warp_id([4]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) T.lane_id([32]) - A_smem = T.alloc_buffer((4, 32, 16), "uint8", scope="shared", layout=s_full, align=1024) + A_smem = T.alloc_tensor((4, 32, 16), "uint8", scope="shared", layout=s_full, align=1024) tmem_addr = T.alloc_shared([1], "uint32") if wg_id == 0: if warp_id == 0: @@ -657,7 +657,7 @@ def kernel(A: T.Buffer((4, 32, 16), "uint8")): T.cuda.cta_sync() Tx.cta.copy(A_smem[:, :, :], A[:, :, :]) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( (4, 32, 16), "uint8", scope="tmem", allocated_addr=tmem_addr[0], layout=t_full ) if tid_in_wg == 0: @@ -928,13 +928,13 @@ def _make_2d_kernel( OUT_BYTES = 16 @T.prim_func(check_well_formed=False) - def kernel(A: T.Buffer(s_full_shape, dtype), B: T.Buffer((OUT_LANES, OUT_BYTES), dtype)): + def kernel(A: T.Tensor(s_full_shape, dtype), B: T.Tensor((OUT_LANES, OUT_BYTES), dtype)): T.device_entry() warp_id = T.warp_id([4]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) lane_id = T.lane_id([32]) - A_smem = T.alloc_buffer(s_full_shape, dtype, scope="shared", layout=s_full, align=1024) + A_smem = T.alloc_tensor(s_full_shape, dtype, scope="shared", layout=s_full, align=1024) tmem_addr = T.alloc_shared([1], "uint32") cp_mbar = T.alloc_shared([1], "uint64") if wg_id == 0: @@ -948,7 +948,7 @@ def kernel(A: T.Buffer(s_full_shape, dtype), B: T.Buffer((OUT_LANES, OUT_BYTES), T.cuda.cta_sync() Tx.cta.copy(A_smem[:, :], A[:, :]) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( t_full_shape, dtype, scope="tmem", @@ -968,7 +968,7 @@ def kernel(A: T.Buffer(s_full_shape, dtype), B: T.Buffer((OUT_LANES, OUT_BYTES), T.cuda.cta_sync() T.ptx.tcgen05.fence__after_thread_sync() if warp_id == 0: - reg = T.alloc_buffer((4,), "uint32", scope="local") + reg = T.alloc_tensor((4,), "uint32", scope="local") for i in range(4): T.ptx["tcgen05.ld.sync.aligned.32x32b.x1.b32"]( reg[i], T.cuda.get_tmem_addr(tmem.allocated_addr[0], 0, i) @@ -991,13 +991,13 @@ def _make_3d_4tile_kernel(s_full, t_full, s_full_shape, t_full_shape, dtype, cta n_tmem_cols_total = max(32, t_full_shape[-1]) @T.prim_func(check_well_formed=False) - def kernel(A: T.Buffer(s_full_shape, dtype), B: T.Buffer((32, 16), dtype)): + def kernel(A: T.Tensor(s_full_shape, dtype), B: T.Tensor((32, 16), dtype)): T.device_entry() warp_id = T.warp_id([4]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) lane_id = T.lane_id([32]) - A_smem = T.alloc_buffer(s_full_shape, dtype, scope="shared", layout=s_full, align=1024) + A_smem = T.alloc_tensor(s_full_shape, dtype, scope="shared", layout=s_full, align=1024) tmem_addr = T.alloc_shared([1], "uint32") cp_mbar = T.alloc_shared([1], "uint64") if wg_id == 0: @@ -1011,7 +1011,7 @@ def kernel(A: T.Buffer(s_full_shape, dtype), B: T.Buffer((32, 16), dtype)): T.cuda.cta_sync() Tx.cta.copy(A_smem[:, :, :], A[:, :, :]) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( t_full_shape, dtype, scope="tmem", @@ -1031,7 +1031,7 @@ def kernel(A: T.Buffer(s_full_shape, dtype), B: T.Buffer((32, 16), dtype)): T.cuda.cta_sync() T.ptx.tcgen05.fence__after_thread_sync() if warp_id == 0: - reg = T.alloc_buffer((4,), "uint32", scope="local") + reg = T.alloc_tensor((4,), "uint32", scope="local") for i in range(4): T.ptx["tcgen05.ld.sync.aligned.32x32b.x1.b32"]( reg[i], T.cuda.get_tmem_addr(tmem.allocated_addr[0], 0, i) @@ -1169,13 +1169,13 @@ def test_align_middle_2_to_1_nvfp4_sfb(): n_tmem_cols_total = max(32, 32) # SFB occupies 32 cols total (8*4 elements / 4 epc) @T.prim_func(check_well_formed=False) - def kernel(A: T.Buffer(s_full_shape, "uint8"), B: T.Buffer((32, 16), "uint8")): + def kernel(A: T.Tensor(s_full_shape, "uint8"), B: T.Tensor((32, 16), "uint8")): T.device_entry() warp_id = T.warp_id([4]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) lane_id = T.lane_id([32]) - A_smem = T.alloc_buffer(s_full_shape, "uint8", scope="shared", layout=s_full, align=1024) + A_smem = T.alloc_tensor(s_full_shape, "uint8", scope="shared", layout=s_full, align=1024) tmem_addr = T.alloc_shared([1], "uint32") cp_mbar = T.alloc_shared([1], "uint64") if wg_id == 0: @@ -1189,7 +1189,7 @@ def kernel(A: T.Buffer(s_full_shape, "uint8"), B: T.Buffer((32, 16), "uint8")): T.cuda.cta_sync() Tx.cta.copy(A_smem[:, :], A[:, :]) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( t_full_shape, "uint8", scope="tmem", @@ -1205,7 +1205,7 @@ def kernel(A: T.Buffer(s_full_shape, "uint8"), B: T.Buffer((32, 16), "uint8")): T.cuda.cta_sync() T.ptx.tcgen05.fence__after_thread_sync() if warp_id == 0: - reg = T.alloc_buffer((4,), "uint32", scope="local") + reg = T.alloc_tensor((4,), "uint32", scope="local") for i in range(4): T.ptx["tcgen05.ld.sync.aligned.32x32b.x1.b32"]( reg[i], T.cuda.get_tmem_addr(tmem.allocated_addr[0], 0, i) diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tcgen05_ldst.py b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tcgen05_ldst.py index 2b4eb9bcbe70..618be43c5af8 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tcgen05_ldst.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tcgen05_ldst.py @@ -273,7 +273,7 @@ def _run_roundtrip_16b( @T.prim_func def kernel( - A: T.Buffer((128, per_thread_elems), dtype), B: T.Buffer((128, per_thread_elems), dtype) + A: T.Tensor((128, per_thread_elems), dtype), B: T.Tensor((128, per_thread_elems), dtype) ) -> None: # Per-thread input/output: A[tid_in_wg, i] feeds register slot i of the # warpgroup-collective fragment; B[tid_in_wg, i] is what comes back @@ -297,7 +297,7 @@ def kernel( T.tvm_storage_sync("shared") - tmem = T.decl_buffer( + tmem = T.decl_tensor( (tmem_rows, stage_width_elem), dtype, scope="tmem", @@ -513,10 +513,10 @@ def test_tcgen05_16xnb_sub_slab_view_read(shape, rep): @T.prim_func def kernel( - A: T.Buffer((128, regs128), dtype), - B128: T.Buffer((128, regs128), dtype), - B0: T.Buffer((128, regs64), dtype), - B1: T.Buffer((128, regs64), dtype), + A: T.Tensor((128, regs128), dtype), + B128: T.Tensor((128, regs128), dtype), + B0: T.Tensor((128, regs64), dtype), + B1: T.Tensor((128, regs64), dtype), ) -> None: T.device_entry() warp_id = T.warp_id([4]) @@ -532,21 +532,21 @@ def kernel( T.address_of(tmem_addr), T.uint32(tmem_cols) ) T.tvm_storage_sync("shared") - tmem_d = T.decl_buffer( + tmem_d = T.decl_tensor( (128, tmem_cols), dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=layout_d, ) - tmem_f0 = T.decl_buffer( + tmem_f0 = T.decl_tensor( (64, tmem_cols), dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=layout_f0, ) - tmem_f1 = T.decl_buffer( + tmem_f1 = T.decl_tensor( (64, tmem_cols), dtype, scope="tmem", @@ -644,7 +644,7 @@ def kernel() -> None: tmem_addr = T.alloc_shared([1], "uint32") if wg_id == 0: T.tvm_storage_sync("shared") - tmem = T.decl_buffer( + tmem = T.decl_tensor( (tmem_rows, stage_width_elem), "float32", scope="tmem", @@ -676,7 +676,7 @@ def kernel() -> None: T.warp_id_in_wg([4]) T.lane_id([32]) tmem_addr = T.alloc_shared([1], "uint32") - tmem = T.decl_buffer( + tmem = T.decl_tensor( (64, n_cols), "float32", scope="tmem", @@ -705,7 +705,7 @@ def kernel() -> None: T.warp_id_in_wg([4]) T.lane_id([32]) tmem_addr = T.alloc_shared([1], "uint32") - tmem = T.decl_buffer( + tmem = T.decl_tensor( (64, n_cols), "float32", scope="tmem", @@ -734,7 +734,7 @@ def kernel() -> None: T.warp_id_in_wg([4]) T.lane_id([32]) tmem_addr = T.alloc_shared([1], "uint32") - tmem = T.decl_buffer( + tmem = T.decl_tensor( (64, n_cols), "float32", scope="tmem", @@ -766,7 +766,7 @@ def test_datapath_B_ld_st_roundtrip(n_cols, col_offset): @T.prim_func def kernel( - A: T.Buffer((128, n_half), "float32"), B: T.Buffer((128, n_half), "float32") + A: T.Tensor((128, n_half), "float32"), B: T.Tensor((128, n_half), "float32") ) -> None: T.device_entry() warp_id = T.warp_id([4]) @@ -783,7 +783,7 @@ def kernel( T.address_of(tmem_addr), T.uint32(tmem_cols) ) T.tvm_storage_sync("shared") - tmem = T.decl_buffer( + tmem = T.decl_tensor( (64, n_cols), "float32", scope="tmem", @@ -872,7 +872,7 @@ def _run_load_test(shape: str, rep: int, dtype: str): @T.prim_func def kernel( - A: T.Buffer((128, stage_width_elem), dtype), B: T.Buffer((128, per_thread_elems), dtype) + A: T.Tensor((128, stage_width_elem), dtype), B: T.Tensor((128, per_thread_elems), dtype) ) -> None: # A is the host data we stage into TMEM via the standard .32x32b path. @@ -898,7 +898,7 @@ def kernel( T.tvm_storage_sync("shared") - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, stage_width_elem), dtype, scope="tmem", @@ -1040,7 +1040,7 @@ def test_tcgen05_st_16xnb_store(shape, rep, dtype): @T.prim_func def kernel( - A: T.Buffer((128, per_thread_elems), dtype), B: T.Buffer((128, stage_width_elem), dtype) + A: T.Tensor((128, per_thread_elems), dtype), B: T.Tensor((128, stage_width_elem), dtype) ) -> None: # A[tid_in_wg, i] is the i-th per-thread element to feed into the atom store. @@ -1066,7 +1066,7 @@ def kernel( T.tvm_storage_sync("shared") - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, stage_width_elem), dtype, scope="tmem", @@ -1174,7 +1174,7 @@ def test_alloc_tcgen05_frag_wrapper_compiles(shape, frag_rows, K_cols): and lowers to the correct tcgen05 atom for each supported instr_shape.""" @T.prim_func - def kernel(A: T.Buffer((128, K_cols), "float32")) -> None: + def kernel(A: T.Tensor((128, K_cols), "float32")) -> None: T.device_entry() warp_id = T.warp_id([4]) T.cta_id([2]) @@ -1190,7 +1190,7 @@ def kernel(A: T.Buffer((128, K_cols), "float32")) -> None: T.address_of(tmem_addr), T.uint32(max(32, K_cols)) ) T.tvm_storage_sync("shared") - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, K_cols), "float32", scope="tmem", @@ -1227,7 +1227,7 @@ def test_tcgen05_32x32b_float32_keeps_typed_register_operands(): K_cols = 32 @T.prim_func - def kernel(A: T.Buffer((128, K_cols), "float32")) -> None: + def kernel(A: T.Tensor((128, K_cols), "float32")) -> None: T.device_entry() warp_id = T.warp_id([4]) T.cta_id([2]) @@ -1243,7 +1243,7 @@ def kernel(A: T.Buffer((128, K_cols), "float32")) -> None: T.address_of(tmem_addr), T.uint32(K_cols) ) T.tvm_storage_sync("shared") - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, K_cols), "float32", scope="tmem", @@ -1293,7 +1293,7 @@ def kernel() -> None: T.thread_id([128]) if wg_id == 0: - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, K_cols), "float32", scope="tmem", @@ -1363,9 +1363,9 @@ def _run_sliced_vs_full_load(shape, full_rep, n_chunks): @T.prim_func def kernel( - A: T.Buffer((128, stage_width_elem), dtype), - Bf: T.Buffer((128, per_thread_elems), dtype), - Bs: T.Buffer((128, per_thread_elems), dtype), + A: T.Tensor((128, stage_width_elem), dtype), + Bf: T.Tensor((128, per_thread_elems), dtype), + Bs: T.Tensor((128, per_thread_elems), dtype), ) -> None: # full-load dump # sliced-load dump @@ -1386,7 +1386,7 @@ def kernel( T.address_of(tmem_addr), T.uint32(tmem_col_width_32b) ) T.tvm_storage_sync("shared") - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, stage_width_elem), dtype, scope="tmem", @@ -1505,7 +1505,7 @@ def next_power_of_2(x): # fmt: off @T.prim_func - def copy_async_test(A: T.Buffer((128, WIDTH), dtype), B: T.Buffer((128, WIDTH), dtype)) -> None: + def copy_async_test(A: T.Tensor((128, WIDTH), dtype), B: T.Tensor((128, WIDTH), dtype)) -> None: A_flat = A.view(-1) B_flat = B.view(-1) @@ -1526,7 +1526,7 @@ def copy_async_test(A: T.Buffer((128, WIDTH), dtype), B: T.Buffer((128, WIDTH), T.tvm_storage_sync("shared") - tmem = T.decl_buffer((128, WIDTH), dtype, scope="tmem", allocated_addr=tmem_addr[0], + tmem = T.decl_tensor((128, WIDTH), dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, WIDTH) : (1 @ TLane, 1 @ TCol)])) A_reg = T.alloc_local((WIDTH), dtype) @@ -1601,7 +1601,7 @@ def next_power_of_2(x): # fmt: off @T.prim_func - def copy_sync(A: T.Buffer((128, WIDTH), dtype), B: T.Buffer((128, WIDTH), dtype)) -> None: + def copy_sync(A: T.Tensor((128, WIDTH), dtype), B: T.Tensor((128, WIDTH), dtype)) -> None: A_flat = A.view(-1) B_flat = B.view(-1) @@ -1622,7 +1622,7 @@ def copy_sync(A: T.Buffer((128, WIDTH), dtype), B: T.Buffer((128, WIDTH), dtype) T.tvm_storage_sync("shared") - tmem = T.decl_buffer((128, OFFSET + WIDTH), dtype, scope="tmem", allocated_addr=tmem_addr[0], # noqa: E501 + tmem = T.decl_tensor((128, OFFSET + WIDTH), dtype, scope="tmem", allocated_addr=tmem_addr[0], # noqa: E501 layout=TileLayout(S[(128, OFFSET + WIDTH) : (1 @ TLane, 1 @ TCol)])) # noqa: E501 A_reg = T.alloc_local((WIDTH), dtype) @@ -1698,7 +1698,7 @@ def next_power_of_2(x): # fmt: off @T.prim_func - def copy_sync(A: T.Buffer((128, WIDTH), dtype), B: T.Buffer((128, WIDTH), dtype)) -> None: + def copy_sync(A: T.Tensor((128, WIDTH), dtype), B: T.Tensor((128, WIDTH), dtype)) -> None: A_flat = A.view(-1) B_flat = B.view(-1) @@ -1719,7 +1719,7 @@ def copy_sync(A: T.Buffer((128, WIDTH), dtype), B: T.Buffer((128, WIDTH), dtype) T.tvm_storage_sync("shared") - tmem = T.decl_buffer((128, WIDTH), dtype, scope="tmem", allocated_addr=tmem_addr[0], + tmem = T.decl_tensor((128, WIDTH), dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, WIDTH) : (1 @ TLane, 1 @ TCol)])) A_reg = T.alloc_local((TOTAL_LOCAL_WIDTH), dtype) diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py index ab2e7b49f857..1836719c1cf5 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tma.py @@ -199,7 +199,7 @@ def _make_op( g_layout = g_layout or _plain_layout(g_shape) s_layout = s_layout or _plain_layout(s_shape) s_dtype = s_dtype or dtype - g_buf = tvm.tirx.decl_buffer( + g_buf = tvm.tirx.decl_tensor( g_shape, dtype, "A", @@ -207,7 +207,7 @@ def _make_op( elem_offset=g_elem_offset, layout=g_layout, ) - s_buf = tvm.tirx.decl_buffer( + s_buf = tvm.tirx.decl_tensor( s_shape, s_dtype, "A_smem", @@ -302,7 +302,7 @@ def _ints(values): def _make_spec(**overrides): g_data = Var("A", PointerType(PrimType("float16"), "global")) - s_buf = tvm.tirx.decl_buffer( + s_buf = tvm.tirx.decl_tensor( (4, 8), "float16", "A_smem", @@ -1118,13 +1118,13 @@ def test_dispatch_propagates_flat_bind_to_auto_coordinate_proof(): func = _from_source( """ @T.prim_func -def bind_coordinate(D: T.Buffer((33360, 6144), 'bfloat16')): +def bind_coordinate(D: T.Tensor((33360, 6144), 'bfloat16')): T.device_entry() block = T.cta_id([192]) tid = T.thread_id([1]) tile_index = T.alloc_local((1,), "int32") - D_smem = T.alloc_buffer( + D_smem = T.alloc_tensor( (2, 16, 128), "bfloat16", scope="shared.dyn", @@ -1526,15 +1526,15 @@ def test_explicit_allows_different_operand_ranks_with_equal_payload_bytes(): source = _from_source( """ @T.prim_func -def rank_change(A: T.Buffer((8, 8), 'float16')): +def rank_change(A: T.Tensor((8, 8), 'float16')): T.device_entry() T.cta_id([1]) tid = T.thread_id([1]) - dyn = T.alloc_buffer((65,), "uint64", scope="shared.dyn") + dyn = T.alloc_tensor((65,), "uint64", scope="shared.dyn") T.attr({"tirx.dyn_smem_bytes": 65 * 8}) - A_smem = T.decl_buffer((64,), "float16", dyn.data, layout=T.TileLayout(T.S[64])) - mbar = T.decl_buffer((1,), "uint64", dyn.data, elem_offset=16) + A_smem = T.decl_tensor((64,), "float16", dyn.data, layout=T.TileLayout(T.S[64])) + mbar = T.decl_tensor((1,), "uint64", dyn.data, elem_offset=16) if tid == 0: Tx.copy_async( A_smem[:], A[:, :], dispatch="tma_explicit", mbar=mbar.ptr_to([0]) @@ -1548,8 +1548,8 @@ def rank_change(A: T.Buffer((8, 8), 'float16')): _SELECTOR_SOURCE = """ @T.prim_func def selector_gather( - A: T.Buffer((256, 64), 'bfloat16'), - B: T.Buffer((512, 80), 'bfloat16'), + A: T.Tensor((256, 64), 'bfloat16'), + B: T.Tensor((512, 80), 'bfloat16'), flag: T.int32, ): @@ -1558,12 +1558,12 @@ def selector_gather( T.device_entry() T.cta_id([1]) tid = T.thread_id([128]) - dyn = T.alloc_buffer((520,), "uint64", scope="shared.dyn") + dyn = T.alloc_tensor((520,), "uint64", scope="shared.dyn") T.attr({"tirx.dyn_smem_bytes": 520 * 8}) - A_smem = T.decl_buffer( + A_smem = T.decl_tensor( (4, 64), "bfloat16", dyn.data, layout=T.TileLayout(T.S[4, 64]) ) - mbar = T.decl_buffer((1,), "uint64", dyn.data, elem_offset=64) + mbar = T.decl_tensor((1,), "uint64", dyn.data, elem_offset=64) if tid == 0: T.ptx.mbarrier.init.shared.b64(mbar.ptr_to([0]), T.uint32(1)) Tx.copy_async( @@ -1698,8 +1698,8 @@ def _build_sparse_decode_qo_tma_regression(): # fmt: off @T.prim_func def kernel( - Q_storage: T.Buffer((64 * 576,), 'bfloat16'), - O_storage: T.Buffer((64 * 512,), 'bfloat16'), + Q_storage: T.Tensor((64 * 576,), 'bfloat16'), + O_storage: T.Tensor((64 * 512,), 'bfloat16'), q_stride_b: T.int64, q_stride_s: T.int64, q_stride_h: T.int64, @@ -1729,12 +1729,12 @@ def kernel( T.device_entry() T.cta_id([1]) tid = T.thread_id([128]) - dyn = T.alloc_buffer((shared_bytes + 8,), "uint8", scope="shared.dyn") + dyn = T.alloc_tensor((shared_bytes + 8,), "uint8", scope="shared.dyn") T.attr({"tirx.dyn_smem_bytes": shared_bytes + 8}) - q_smem = T.decl_buffer( + q_smem = T.decl_tensor( (64, 512), "bfloat16", dyn.data, scope="shared.dyn", layout=q_layout ) - q_tail_smem = T.decl_buffer( + q_tail_smem = T.decl_tensor( (64, 64), "bfloat16", dyn.data, @@ -1742,7 +1742,7 @@ def kernel( scope="shared.dyn", layout=q_tail_layout, ) - o_smem = T.decl_buffer( + o_smem = T.decl_tensor( (64, 512), "bfloat16", dyn.data, @@ -1750,7 +1750,7 @@ def kernel( scope="shared.dyn", layout=o_layout, ) - mbar = T.decl_buffer( + mbar = T.decl_tensor( (1,), "uint64", dyn.data, elem_offset=shared_bytes // 8, scope="shared.dyn" ) q_tail_smem_tma = q_tail_smem.view(64, 2, 32).permute(1, 0, 2) @@ -1994,21 +1994,21 @@ def _build_selector_gather_gpu_kernel(dtype="float16"): # fmt: off @T.prim_func def kernel( - A: T.Buffer((rows, cols), dtype), - B: T.Buffer((rows, cols), dtype), + A: T.Tensor((rows, cols), dtype), + B: T.Tensor((rows, cols), dtype), flag: T.int32, - Out: T.Buffer((4, cols), dtype), + Out: T.Tensor((4, cols), dtype), ): T.device_entry() T.cta_id([1]) tid = T.thread_id([128]) - dyn = T.alloc_buffer((shared_bytes + 64,), "uint8", scope="shared.dyn") + dyn = T.alloc_tensor((shared_bytes + 64,), "uint8", scope="shared.dyn") T.attr({"tirx.dyn_smem_bytes": shared_bytes + 64}) - A_smem = T.decl_buffer( + A_smem = T.decl_tensor( (4, cols), dtype, dyn.data, layout=T.TileLayout(T.S[4, cols]) ) - mbar = T.decl_buffer((1,), "uint64", dyn.data, elem_offset=shared_bytes // 8) + mbar = T.decl_tensor((1,), "uint64", dyn.data, elem_offset=shared_bytes // 8) mbar_ptr = T.meta_var(mbar.ptr_to([0])) if tid == 0: T.ptx.mbarrier.init.shared.b64(mbar_ptr, T.uint32(1)) diff --git a/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_binary.py b/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_binary.py index db003dfcf566..81d2f65beb8e 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_binary.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_binary.py @@ -91,14 +91,14 @@ def test_binary_op_shared(input, op_type, operands_type, dtype): # fmt: off @T.prim_func def binary_op_region_region( - A: T.Buffer(g_shape, dtype, layout=g_layout), B: T.Buffer(g_shape, dtype, layout=g_layout) + A: T.Tensor(g_shape, dtype, layout=g_layout), B: T.Tensor(g_shape, dtype, layout=g_layout) ) -> None: T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([thread_cnt]) - A_smem = T.alloc_buffer(g_shape, dtype, scope="shared", layout=s_layout) - B_smem = T.alloc_buffer(g_shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(g_shape, dtype, scope="shared", layout=s_layout) + B_smem = T.alloc_tensor(g_shape, dtype, scope="shared", layout=s_layout) Tx.cta.copy(A_smem[tuple(copy_slice)], A[tuple(copy_slice)]) Tx.cta.copy(B_smem[tuple(copy_slice)], B[tuple(copy_slice)]) @@ -124,13 +124,13 @@ def binary_op_region_region( @T.prim_func def binary_op_const_region_or_region_const( - A: T.Buffer(g_shape, dtype, layout=g_layout), _B: T.Buffer(g_shape, dtype, layout=g_layout) + A: T.Tensor(g_shape, dtype, layout=g_layout), _B: T.Tensor(g_shape, dtype, layout=g_layout) ) -> None: T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([thread_cnt]) - A_smem = T.alloc_buffer(g_shape, dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(g_shape, dtype, scope="shared", layout=s_layout) Tx.cta.copy(A_smem[tuple(copy_slice)], A[tuple(copy_slice)]) T.cuda.cta_sync() @@ -225,7 +225,7 @@ def bad_kernel() -> None: T.device_entry() _bx = T.cta_id([1]) _tid = T.thread_id([64]) - A_smem = T.alloc_buffer(shape, dtype, scope="shared", layout=layout) + A_smem = T.alloc_tensor(shape, dtype, scope="shared", layout=layout) if op_type == "sub": Tx.cta.sub(A_smem, const, A_smem) elif op_type == "fdiv": @@ -255,16 +255,16 @@ def test_binary_op_shared_subcta_scope(exec_scope, op_type): @T.prim_func def kernel( - A: T.Buffer(g_shape, dtype, layout=TileLayout(S[g_shape])), - B: T.Buffer(g_shape, dtype, layout=TileLayout(S[g_shape])), + A: T.Tensor(g_shape, dtype, layout=TileLayout(S[g_shape])), + B: T.Tensor(g_shape, dtype, layout=TileLayout(S[g_shape])), ) -> None: T.device_entry() warp_id = T.warp_id([(256) // 32]) wg_id = T.warpgroup_id([(256) // 128]) _bx = T.cta_id([1]) _tid = T.thread_id([256]) - A_smem = T.alloc_buffer(g_shape, dtype, scope="shared", layout=TileLayout(S[g_shape])) - B_smem = T.alloc_buffer(g_shape, dtype, scope="shared", layout=TileLayout(S[g_shape])) + A_smem = T.alloc_tensor(g_shape, dtype, scope="shared", layout=TileLayout(S[g_shape])) + B_smem = T.alloc_tensor(g_shape, dtype, scope="shared", layout=TileLayout(S[g_shape])) Tx.cta.copy(A_smem, A) Tx.cta.copy(B_smem, B) T.cuda.cta_sync() @@ -319,9 +319,9 @@ def test_binary_op_local_subcta_trivial(exec_scope, rhs_kind, op_type): @T.prim_func def kernel( - A: T.Buffer(a_shape, dtype, layout=TileLayout(S[a_shape])), - B: T.Buffer(b_shape, dtype, layout=TileLayout(S[b_shape])), - C: T.Buffer(c_shape, dtype, layout=TileLayout(S[c_shape])), + A: T.Tensor(a_shape, dtype, layout=TileLayout(S[a_shape])), + B: T.Tensor(b_shape, dtype, layout=TileLayout(S[b_shape])), + C: T.Tensor(c_shape, dtype, layout=TileLayout(S[c_shape])), ) -> None: T.device_entry() wg_id = T.warpgroup_id([(256) // 128]) @@ -330,9 +330,9 @@ def kernel( _tid = T.thread_id([256]) tid_in_scope = tid_in_scope_fn([n_threads]) b_n = T.meta_var(n if rhs_kind == "region" else 1) - A_local = T.alloc_buffer((m, n), dtype, scope="local", layout=TileLayout(S[(m, n)])) - C_local = T.alloc_buffer((m, n), dtype, scope="local", layout=TileLayout(S[(m, n)])) - B_local = T.alloc_buffer((m, b_n), dtype, scope="local", layout=TileLayout(S[(m, b_n)])) + A_local = T.alloc_tensor((m, n), dtype, scope="local", layout=TileLayout(S[(m, n)])) + C_local = T.alloc_tensor((m, n), dtype, scope="local", layout=TileLayout(S[(m, n)])) + B_local = T.alloc_tensor((m, b_n), dtype, scope="local", layout=TileLayout(S[(m, b_n)])) if thr_str <= _tid and _tid < thr_str + n_threads: for i in T.serial(m): @@ -430,16 +430,16 @@ def test_binary_op_vectorized(input, storage_scope, exec_scope, op_type, dtype): # fmt: off @T.prim_func def test_binary_cta( - A: T.Buffer(a_shape, dtype, layout=TileLayout(S[a_shape])), - B: T.Buffer(b_shape, dtype, layout=TileLayout(S[b_shape])), + A: T.Tensor(a_shape, dtype, layout=TileLayout(S[a_shape])), + B: T.Tensor(b_shape, dtype, layout=TileLayout(S[b_shape])), ) -> None: T.device_entry() _bx = T.cta_id([1]) tx = T.thread_id([thread_cnt]) if storage_scope == "shared": - A_smem = T.alloc_buffer(a_shape, dtype, scope="shared", layout=TileLayout(S[a_shape])) - B_smem = T.alloc_buffer(b_shape, dtype, scope="shared", layout=TileLayout(S[b_shape])) + A_smem = T.alloc_tensor(a_shape, dtype, scope="shared", layout=TileLayout(S[a_shape])) + B_smem = T.alloc_tensor(b_shape, dtype, scope="shared", layout=TileLayout(S[b_shape])) Tx.cta.copy(A_smem, A) Tx.cta.copy(B_smem, B) T.cuda.cta_sync() @@ -447,10 +447,10 @@ def test_binary_cta( T.cuda.cta_sync() Tx.cta.copy(A, A_smem) if storage_scope == "local": - A_local = T.alloc_buffer( + A_local = T.alloc_tensor( a_shape[1:], dtype, scope="local", layout=TileLayout(S[a_shape[1:]]) ) - B_local = T.alloc_buffer( + B_local = T.alloc_tensor( b_shape[1:], dtype, scope="local", layout=TileLayout(S[b_shape[1:]]) ) Tx.copy(A_local, A[tx]) @@ -460,16 +460,16 @@ def test_binary_cta( @T.prim_func def test_binary_thread( - A: T.Buffer(a_shape, dtype, layout=TileLayout(S[a_shape])), - B: T.Buffer(b_shape, dtype, layout=TileLayout(S[b_shape])), + A: T.Tensor(a_shape, dtype, layout=TileLayout(S[a_shape])), + B: T.Tensor(b_shape, dtype, layout=TileLayout(S[b_shape])), ) -> None: T.device_entry() _bx = T.cta_id([1]) tx = T.thread_id([thread_cnt]) if storage_scope == "shared": - A_smem = T.alloc_buffer(a_shape, dtype, scope="shared", layout=TileLayout(S[a_shape])) - B_smem = T.alloc_buffer(b_shape, dtype, scope="shared", layout=TileLayout(S[b_shape])) + A_smem = T.alloc_tensor(a_shape, dtype, scope="shared", layout=TileLayout(S[a_shape])) + B_smem = T.alloc_tensor(b_shape, dtype, scope="shared", layout=TileLayout(S[b_shape])) Tx.copy(A_smem, A) Tx.copy(B_smem, B) T.cuda.cta_sync() @@ -477,10 +477,10 @@ def test_binary_thread( T.cuda.cta_sync() Tx.copy(A, A_smem) elif storage_scope == "local": - A_local = T.alloc_buffer( + A_local = T.alloc_tensor( a_shape[1:], dtype, scope="local", layout=TileLayout(S[a_shape[1:]]) ) - B_local = T.alloc_buffer( + B_local = T.alloc_tensor( b_shape[1:], dtype, scope="local", layout=TileLayout(S[b_shape[1:]]) ) Tx.copy(A_local, A[tx]) @@ -539,16 +539,16 @@ def test_binary_op_packed_f32x2_auto_dispatch(op_type): @T.prim_func def test_func( - A: T.Buffer(a_shape, dtype, layout=TileLayout(S[a_shape])), - B: T.Buffer(b_shape, dtype, layout=TileLayout(S[b_shape])), + A: T.Tensor(a_shape, dtype, layout=TileLayout(S[a_shape])), + B: T.Tensor(b_shape, dtype, layout=TileLayout(S[b_shape])), ) -> None: T.device_entry() _bx = T.cta_id([1]) tx = T.thread_id([64]) - A_local = T.alloc_buffer( + A_local = T.alloc_tensor( a_shape[1:], dtype, scope="local", layout=TileLayout(S[a_shape[1:]]) ) - B_local = T.alloc_buffer( + B_local = T.alloc_tensor( b_shape[1:], dtype, scope="local", layout=TileLayout(S[b_shape[1:]]) ) Tx.copy(A_local, A[tx]) @@ -601,18 +601,18 @@ def test_binary_op_warpgroup_wg_local_layout(op_name): @T.prim_func def test_func( - A: T.Buffer((rows, cols), dtype, layout=TileLayout(S[rows, cols])), - B: T.Buffer((rows, cols), dtype, layout=TileLayout(S[rows, cols])), - C: T.Buffer((rows, cols), dtype, layout=TileLayout(S[rows, cols])), + A: T.Tensor((rows, cols), dtype, layout=TileLayout(S[rows, cols])), + B: T.Tensor((rows, cols), dtype, layout=TileLayout(S[rows, cols])), + C: T.Tensor((rows, cols), dtype, layout=TileLayout(S[rows, cols])), ) -> None: T.device_entry() _bx = T.cta_id([1]) wg_id = T.warpgroup_id([1]) tid = T.thread_id_in_wg([rows]) - lhs = T.alloc_buffer((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) - rhs = T.alloc_buffer((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) - out = T.alloc_buffer((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) + lhs = T.alloc_tensor((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) + rhs = T.alloc_tensor((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) + out = T.alloc_tensor((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) lhs_row = lhs.local(cols) rhs_row = rhs.local(cols) out_row = out.local(cols) @@ -678,18 +678,18 @@ def test_binary_op_warpgroup_wg_local_emits_packed_f32x2(op_name, ptx_op): @T.prim_func def test_func( - A: T.Buffer((rows, cols), dtype, layout=TileLayout(S[rows, cols])), - B: T.Buffer((rows, cols), dtype, layout=TileLayout(S[rows, cols])), - C: T.Buffer((rows, cols), dtype, layout=TileLayout(S[rows, cols])), + A: T.Tensor((rows, cols), dtype, layout=TileLayout(S[rows, cols])), + B: T.Tensor((rows, cols), dtype, layout=TileLayout(S[rows, cols])), + C: T.Tensor((rows, cols), dtype, layout=TileLayout(S[rows, cols])), ) -> None: T.device_entry() _bx = T.cta_id([1]) _wg_id = T.warpgroup_id([1]) tid = T.thread_id_in_wg([rows]) - lhs = T.alloc_buffer((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) - rhs = T.alloc_buffer((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) - out = T.alloc_buffer((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) + lhs = T.alloc_tensor((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) + rhs = T.alloc_tensor((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) + out = T.alloc_tensor((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) lhs_row = lhs.local(cols) rhs_row = rhs.local(cols) out_row = out.local(cols) @@ -733,15 +733,15 @@ def test_fma_warpgroup_wg_local_emits_packed_f32x2(): @T.prim_func def test_func( - A: T.Buffer((rows, cols), dtype, layout=TileLayout(S[rows, cols])), - C: T.Buffer((rows, cols), dtype, layout=TileLayout(S[rows, cols])), + A: T.Tensor((rows, cols), dtype, layout=TileLayout(S[rows, cols])), + C: T.Tensor((rows, cols), dtype, layout=TileLayout(S[rows, cols])), ) -> None: T.device_entry() _bx = T.cta_id([1]) _wg_id = T.warpgroup_id([1]) tid = T.thread_id_in_wg([rows]) - buf = T.alloc_buffer((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) + buf = T.alloc_tensor((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) buf_row = buf.local(cols) for i in T.serial(cols): buf_row[i] = A[tid, i] @@ -773,13 +773,13 @@ def test_binary_add_f32_sm100_packed_f32x2_dispatch(): @T.prim_func def k( - A: T.Buffer(shape, "float32", layout=lay), B: T.Buffer(shape, "float32", layout=lay) + A: T.Tensor(shape, "float32", layout=lay), B: T.Tensor(shape, "float32", layout=lay) ) -> None: T.device_entry() _bx = T.cta_id([1]) tx = T.thread_id([64]) - ra = T.alloc_buffer(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) - rb = T.alloc_buffer(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) + ra = T.alloc_tensor(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) + rb = T.alloc_tensor(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) Tx.copy(ra, A[tx]) Tx.copy(rb, B[tx]) Tx.add(ra, ra, rb) @@ -802,7 +802,7 @@ def test_binary_maximum_reg(): @T.prim_func def relu_max( - A: T.Buffer((N,), "float32"), B: T.Buffer((N,), "float32"), C: T.Buffer((N,), "float32") + A: T.Tensor((N,), "float32"), B: T.Tensor((N,), "float32"), C: T.Tensor((N,), "float32") ) -> None: T.device_entry() T.warp_id([4]) @@ -840,13 +840,13 @@ def test_binary_add_f16_scalar_fallback_dispatch(): @T.prim_func def k( - A: T.Buffer(shape, "float16", layout=lay), B: T.Buffer(shape, "float16", layout=lay) + A: T.Tensor(shape, "float16", layout=lay), B: T.Tensor(shape, "float16", layout=lay) ) -> None: T.device_entry() _bx = T.cta_id([1]) tx = T.thread_id([64]) - ra = T.alloc_buffer(shape[1:], "float16", scope="local", layout=TileLayout(S[shape[1:]])) - rb = T.alloc_buffer(shape[1:], "float16", scope="local", layout=TileLayout(S[shape[1:]])) + ra = T.alloc_tensor(shape[1:], "float16", scope="local", layout=TileLayout(S[shape[1:]])) + rb = T.alloc_tensor(shape[1:], "float16", scope="local", layout=TileLayout(S[shape[1:]])) Tx.copy(ra, A[tx]) Tx.copy(rb, B[tx]) Tx.add(ra, ra, rb) @@ -871,8 +871,8 @@ def test_mul_tcgen05_16x256b_atom_warpgroup_dispatch(): @T.prim_func def kernel( - A: T.Buffer((128, regs_per_thread), "float32"), - B: T.Buffer((128, regs_per_thread), "float32"), + A: T.Tensor((128, regs_per_thread), "float32"), + B: T.Tensor((128, regs_per_thread), "float32"), ) -> None: T.device_entry() T.cta_id([1]) @@ -880,8 +880,8 @@ def kernel( T.warp_id_in_wg([4]) T.lane_id([32]) tid = T.thread_id_in_wg([128]) - src = T.alloc_buffer((rows, cols), "float32", scope="local", layout=atom_layout) - dst = T.alloc_buffer((rows, cols), "float32", scope="local", layout=atom_layout) + src = T.alloc_tensor((rows, cols), "float32", scope="local", layout=atom_layout) + dst = T.alloc_tensor((rows, cols), "float32", scope="local", layout=atom_layout) src_local = src.local(regs_per_thread) dst_local = dst.local(regs_per_thread) for i in T.serial(regs_per_thread): diff --git a/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_fma.py b/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_fma.py index ac27cbd7125c..2f1adf9384e9 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_fma.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_fma.py @@ -57,11 +57,11 @@ def test_fma_scalar_scalar(): bias_val = -1.0 @T.prim_func - def test_func(A: T.Buffer((N,), dtype, layout=TileLayout(S[N]))) -> None: + def test_func(A: T.Tensor((N,), dtype, layout=TileLayout(S[N]))) -> None: T.device_entry() _bx = T.cta_id([1]) tx = T.thread_id([N]) - buf = T.alloc_buffer((1,), dtype, scope="local", layout=TileLayout(S[1])) + buf = T.alloc_tensor((1,), dtype, scope="local", layout=TileLayout(S[1])) Tx.copy(buf, A[tx : tx + 1]) Tx.fma(buf, buf, T.float32(scale_val), T.float32(bias_val)) Tx.copy(A[tx : tx + 1], buf) @@ -99,14 +99,14 @@ def test_fma_buffer_scale_scalar_bias(): @T.prim_func def test_func( - A: T.Buffer((N,), dtype, layout=TileLayout(S[N])), - B: T.Buffer((N,), dtype, layout=TileLayout(S[N])), + A: T.Tensor((N,), dtype, layout=TileLayout(S[N])), + B: T.Tensor((N,), dtype, layout=TileLayout(S[N])), ) -> None: T.device_entry() _bx = T.cta_id([1]) _tx = T.thread_id([1]) - acc = T.alloc_buffer((N,), dtype, scope="local", layout=TileLayout(S[N])) - frac = T.alloc_buffer((N,), dtype, scope="local", layout=TileLayout(S[N])) + acc = T.alloc_tensor((N,), dtype, scope="local", layout=TileLayout(S[N])) + frac = T.alloc_tensor((N,), dtype, scope="local", layout=TileLayout(S[N])) Tx.copy(acc, A[0:N]) Tx.copy(frac, B[0:N]) Tx.fma(acc, acc, frac, T.float32(coeff)) @@ -145,14 +145,14 @@ def test_mul_scalar_broadcast(): @T.prim_func def test_func( - A: T.Buffer((N,), dtype, layout=TileLayout(S[N])), - Scale: T.Buffer((1,), dtype, layout=TileLayout(S[1])), + A: T.Tensor((N,), dtype, layout=TileLayout(S[N])), + Scale: T.Tensor((1,), dtype, layout=TileLayout(S[1])), ) -> None: T.device_entry() _bx = T.cta_id([1]) _tx = T.thread_id([1]) - a_local = T.alloc_buffer((N,), dtype, scope="local", layout=TileLayout(S[N])) - s_local = T.alloc_buffer((1,), dtype, scope="local", layout=TileLayout(S[1])) + a_local = T.alloc_tensor((N,), dtype, scope="local", layout=TileLayout(S[N])) + s_local = T.alloc_tensor((1,), dtype, scope="local", layout=TileLayout(S[1])) Tx.copy(a_local, A[0:N]) Tx.copy(s_local, Scale[0:1]) Tx.mul(a_local, a_local, s_local[0]) @@ -192,11 +192,11 @@ def test_add_rounding_mode(): round_const = float(2**23 + 2**22) @T.prim_func - def test_func(A: T.Buffer((N,), dtype, layout=TileLayout(S[N]))) -> None: + def test_func(A: T.Tensor((N,), dtype, layout=TileLayout(S[N]))) -> None: T.device_entry() _bx = T.cta_id([1]) _tx = T.thread_id([1]) - buf = T.alloc_buffer((N,), dtype, scope="local", layout=TileLayout(S[N])) + buf = T.alloc_tensor((N,), dtype, scope="local", layout=TileLayout(S[N])) Tx.copy(buf, A[0:N]) Tx.add(buf, buf, T.float32(round_const), rounding_mode="rm") Tx.copy(A[0:N], buf) @@ -239,7 +239,7 @@ def test_fma_no_layout(): bias_val = 1.0 @T.prim_func - def test_func(A: T.Buffer((N,), dtype, layout=TileLayout(S[N]))) -> None: + def test_func(A: T.Tensor((N,), dtype, layout=TileLayout(S[N]))) -> None: T.device_entry() _bx = T.cta_id([1]) _tx = T.thread_id([1]) @@ -281,14 +281,14 @@ def test_sub_buffer_buffer_rounding(): @T.prim_func def test_func( - A: T.Buffer((N,), dtype, layout=TileLayout(S[N])), - B: T.Buffer((N,), dtype, layout=TileLayout(S[N])), + A: T.Tensor((N,), dtype, layout=TileLayout(S[N])), + B: T.Tensor((N,), dtype, layout=TileLayout(S[N])), ) -> None: T.device_entry() _bx = T.cta_id([1]) _tx = T.thread_id([1]) - a_buf = T.alloc_buffer((N,), dtype, scope="local", layout=TileLayout(S[N])) - b_buf = T.alloc_buffer((N,), dtype, scope="local", layout=TileLayout(S[N])) + a_buf = T.alloc_tensor((N,), dtype, scope="local", layout=TileLayout(S[N])) + b_buf = T.alloc_tensor((N,), dtype, scope="local", layout=TileLayout(S[N])) Tx.copy(a_buf, A[0:N]) Tx.copy(b_buf, B[0:N]) Tx.sub(a_buf, a_buf, b_buf, rounding_mode="rn") @@ -326,15 +326,15 @@ def test_fma_warpgroup_wg_local_layout(): @T.prim_func def test_func( - A: T.Buffer((rows, cols), dtype, layout=TileLayout(S[rows, cols])), - B: T.Buffer((rows, cols), dtype, layout=TileLayout(S[rows, cols])), + A: T.Tensor((rows, cols), dtype, layout=TileLayout(S[rows, cols])), + B: T.Tensor((rows, cols), dtype, layout=TileLayout(S[rows, cols])), ) -> None: T.device_entry() _bx = T.cta_id([1]) wg_id = T.warpgroup_id([1]) tid = T.thread_id_in_wg([rows]) - reg = T.alloc_buffer((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) + reg = T.alloc_tensor((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) reg_row = reg.local(cols) for i in T.serial(cols): reg_row[i] = A[tid, i] @@ -373,18 +373,18 @@ def test_fma_f32_sm100_packed_f32x2_dispatch(): @T.prim_func def k( - A: T.Buffer(shape, "float32", layout=lay), - B: T.Buffer(shape, "float32", layout=lay), - C: T.Buffer(shape, "float32", layout=lay), - D: T.Buffer(shape, "float32", layout=lay), + A: T.Tensor(shape, "float32", layout=lay), + B: T.Tensor(shape, "float32", layout=lay), + C: T.Tensor(shape, "float32", layout=lay), + D: T.Tensor(shape, "float32", layout=lay), ) -> None: T.device_entry() _bx = T.cta_id([1]) tx = T.thread_id([64]) - ra = T.alloc_buffer(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) - rb = T.alloc_buffer(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) - rc = T.alloc_buffer(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) - rd = T.alloc_buffer(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) + ra = T.alloc_tensor(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) + rb = T.alloc_tensor(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) + rc = T.alloc_tensor(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) + rd = T.alloc_tensor(shape[1:], "float32", scope="local", layout=TileLayout(S[shape[1:]])) Tx.copy(ra, A[tx]) Tx.copy(rb, B[tx]) Tx.copy(rc, C[tx]) diff --git a/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_unary.py b/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_unary.py index da73413b8ff4..0671c8a76f11 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_unary.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/elementwise/test_unary.py @@ -74,12 +74,12 @@ def test_unary_op_shared(input, op_type, src_dtype, dst_dtype): if in_place: # fmt: off @T.prim_func - def unary_op(A: T.Buffer(g_shape, src_dtype, layout=g_layout)) -> None: + def unary_op(A: T.Tensor(g_shape, src_dtype, layout=g_layout)) -> None: T.device_entry() _bx = T.cta_id([1]) _tx = T.thread_id([thread_cnt]) - A_smem = T.alloc_buffer(s_shape, src_dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(s_shape, src_dtype, scope="shared", layout=s_layout) Tx.cta.copy(A_smem[tuple(copy_slice)], A[tuple(copy_slice)]) T.cuda.cta_sync() if op_type == "zero": @@ -93,15 +93,15 @@ def unary_op(A: T.Buffer(g_shape, src_dtype, layout=g_layout)) -> None: # fmt: off @T.prim_func def unary_op( - A: T.Buffer(g_shape, src_dtype, layout=g_layout), - B: T.Buffer(g_shape, dst_dtype, layout=g_layout), + A: T.Tensor(g_shape, src_dtype, layout=g_layout), + B: T.Tensor(g_shape, dst_dtype, layout=g_layout), ) -> None: T.device_entry() _bx = T.cta_id([1]) _tx = T.thread_id([thread_cnt]) - A_smem = T.alloc_buffer(s_shape, src_dtype, scope="shared", layout=s_layout) - B_smem = T.alloc_buffer(s_shape, dst_dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(s_shape, src_dtype, scope="shared", layout=s_layout) + B_smem = T.alloc_tensor(s_shape, dst_dtype, scope="shared", layout=s_layout) Tx.cta.copy(A_smem[tuple(copy_slice)], A[tuple(copy_slice)]) T.cuda.cta_sync() if op_type == "zero": @@ -160,13 +160,13 @@ def test_unary_op_shared_subcta_scope(exec_scope): g_shape = (n_warps * 32, 8) @T.prim_func - def unary_op_subcta(A: T.Buffer(g_shape, dtype, layout=TileLayout(S[g_shape]))) -> None: + def unary_op_subcta(A: T.Tensor(g_shape, dtype, layout=TileLayout(S[g_shape]))) -> None: T.device_entry() warp_id = T.warp_id([(256) // 32]) wg_id = T.warpgroup_id([(256) // 128]) _bx = T.cta_id([1]) _tid = T.thread_id([256]) - A_smem = T.alloc_buffer(g_shape, dtype, scope="shared", layout=TileLayout(S[g_shape])) + A_smem = T.alloc_tensor(g_shape, dtype, scope="shared", layout=TileLayout(S[g_shape])) Tx.cta.copy(A_smem, A) T.cuda.cta_sync() if exec_scope == "warp": @@ -250,14 +250,14 @@ def test_unary_op_shared_with_bias_scale(input, op_type, bias_type, src_dtype, d @T.prim_func def unary_op_with_bias( - A: T.Buffer(g_shape, src_dtype, layout=g_layout), - bias: T.Buffer(g_shape, src_dtype, layout=g_layout), + A: T.Tensor(g_shape, src_dtype, layout=g_layout), + bias: T.Tensor(g_shape, src_dtype, layout=g_layout), ) -> None: T.device_entry() _bx = T.cta_id([1]) _tx = T.thread_id([thread_cnt]) - A_smem = T.alloc_buffer(s_shape, src_dtype, scope="shared", layout=s_layout) - bias_smem = T.alloc_buffer(s_shape, src_dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(s_shape, src_dtype, scope="shared", layout=s_layout) + bias_smem = T.alloc_tensor(s_shape, src_dtype, scope="shared", layout=s_layout) Tx.cta.copy(A_smem[tuple(copy_slice)], A[tuple(copy_slice)]) Tx.cta.copy(bias_smem[tuple(copy_slice)], bias[tuple(copy_slice)]) T.cuda.cta_sync() @@ -297,16 +297,16 @@ def unary_op_with_bias( @T.prim_func def unary_op_with_bias( - A: T.Buffer(g_shape, src_dtype, layout=g_layout), - B: T.Buffer(g_shape, dst_dtype, layout=g_layout), - bias: T.Buffer(g_shape, src_dtype, layout=g_layout), + A: T.Tensor(g_shape, src_dtype, layout=g_layout), + B: T.Tensor(g_shape, dst_dtype, layout=g_layout), + bias: T.Tensor(g_shape, src_dtype, layout=g_layout), ) -> None: T.device_entry() _bx = T.cta_id([1]) _tx = T.thread_id([thread_cnt]) - A_smem = T.alloc_buffer(s_shape, src_dtype, scope="shared", layout=s_layout) - B_smem = T.alloc_buffer(s_shape, dst_dtype, scope="shared", layout=s_layout) - bias_smem = T.alloc_buffer(s_shape, src_dtype, scope="shared", layout=s_layout) + A_smem = T.alloc_tensor(s_shape, src_dtype, scope="shared", layout=s_layout) + B_smem = T.alloc_tensor(s_shape, dst_dtype, scope="shared", layout=s_layout) + bias_smem = T.alloc_tensor(s_shape, src_dtype, scope="shared", layout=s_layout) Tx.cta.copy(A_smem[tuple(copy_slice)], A[tuple(copy_slice)]) Tx.cta.copy(bias_smem[tuple(copy_slice)], bias[tuple(copy_slice)]) T.cuda.cta_sync() @@ -465,8 +465,8 @@ def test_unary_op_local(input, op_type, src_dtype, dst_dtype): @T.prim_func def test_unary( - A: T.Buffer(g_shape_a, src_dtype, layout=g_layout_a), - B: T.Buffer(g_shape_b, dst_dtype, layout=g_layout_b), + A: T.Tensor(g_shape_a, src_dtype, layout=g_layout_a), + B: T.Tensor(g_shape_b, dst_dtype, layout=g_layout_b), ) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) @@ -479,13 +479,13 @@ def test_unary( warp_atom = atom.tile(warp_layout, (8, 4), (1, 2)) tile = T.TileLayout(T.S[(2, NUM_COL // 8) : (1, 2)]) acc_layout = warp_atom.tile(tile, (2, NUM_COL // 8), (8, 8)) - acc = T.alloc_buffer( + acc = T.alloc_tensor( [2, NUM_COL // 4], dtype=src_dtype, scope="local", layout=atom.tile(tile, (2, NUM_COL // 8), (1, 2)), ) - res = T.alloc_buffer( + res = T.alloc_tensor( [2, NUM_COL // 4], dtype=dst_dtype, scope="local", @@ -602,9 +602,9 @@ def test_unary_op_local_with_bias_scale(input, op_type, bias_type, src_dtype, ds @T.prim_func def test_unary_with_bias( - A: T.Buffer(g_shape_a, src_dtype, layout=g_layout_a), - B: T.Buffer(g_shape_b, dst_dtype, layout=g_layout_b), - bias: T.Buffer(g_shape_bias, src_dtype, layout=g_layout_bias), + A: T.Tensor(g_shape_a, src_dtype, layout=g_layout_a), + B: T.Tensor(g_shape_b, dst_dtype, layout=g_layout_b), + bias: T.Tensor(g_shape_bias, src_dtype, layout=g_layout_bias), ) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) @@ -617,19 +617,19 @@ def test_unary_with_bias( warp_atom = atom.tile(warp_layout, (8, 4), (1, 2)) tile = T.TileLayout(T.S[(2, NUM_COL // 8) : (1, 2)]) acc_layout = warp_atom.tile(tile, (2, NUM_COL // 8), (8, 8)) - acc = T.alloc_buffer( + acc = T.alloc_tensor( [2, NUM_COL // 4], dtype=src_dtype, scope="local", layout=atom.tile(tile, (2, NUM_COL // 8), (1, 2)), ) - bias_local = T.alloc_buffer( + bias_local = T.alloc_tensor( [2, NUM_COL // 4], dtype=src_dtype, scope="local", layout=atom.tile(tile, (2, NUM_COL // 8), (1, 2)), ) - res = T.alloc_buffer( + res = T.alloc_tensor( [2, NUM_COL // 4], dtype=dst_dtype, scope="local", @@ -728,32 +728,32 @@ def test_unary_op_vectorized(shape, op_type, exec_scope, storage_scope): # fmt: off @T.prim_func - def test_unary_thread(A: T.Buffer(shape, dtype, layout=TileLayout(S[shape]))) -> None: + def test_unary_thread(A: T.Tensor(shape, dtype, layout=TileLayout(S[shape]))) -> None: T.device_entry() _bx = T.cta_id([1]) tx = T.thread_id([128]) if storage_scope == "shared": - a_smem = T.alloc_buffer( + a_smem = T.alloc_tensor( shape, dtype=dtype, layout=TileLayout(S[shape]), scope="shared" ) Tx.fill(a_smem[tx], value) Tx.copy(A[tx], a_smem[tx]) elif storage_scope == "local": - a_local = T.alloc_buffer( + a_local = T.alloc_tensor( shape[1:], dtype=dtype, layout=TileLayout(S[shape[1:]]), scope="local" ) Tx.fill(a_local, value) Tx.copy(A[tx], a_local) @T.prim_func - def test_unary_cta(A: T.Buffer(shape, dtype, layout=TileLayout(S[shape]))) -> None: + def test_unary_cta(A: T.Tensor(shape, dtype, layout=TileLayout(S[shape]))) -> None: T.device_entry() _bx = T.cta_id([1]) _tid = T.thread_id([128]) if storage_scope == "shared": - a_smem = T.alloc_buffer( + a_smem = T.alloc_tensor( shape, dtype=dtype, layout=TileLayout(S[shape]), scope="shared" ) Tx.cta.fill(a_smem, value) @@ -786,11 +786,11 @@ def test_unary_op_local_thread_wise(op_type, dtype): local_shape = shape[1:] @T.prim_func - def kernel(A: T.Buffer(shape, dtype, layout=TileLayout(S[shape]))) -> None: + def kernel(A: T.Tensor(shape, dtype, layout=TileLayout(S[shape]))) -> None: T.device_entry() _bx = T.cta_id([1]) tid = T.thread_id([64]) - a_local = T.alloc_buffer( + a_local = T.alloc_tensor( local_shape, dtype, scope="local", layout=TileLayout(S[local_shape]) ) Tx.copy(a_local, A[tid]) @@ -848,8 +848,8 @@ def test_cast_thread_local(shape, A_dtype, B_dtype): # fmt: off @T.prim_func def test_cast( - A: T.Buffer(shape, A_dtype, layout=TileLayout(S[shape])), - B: T.Buffer(shape, B_dtype, layout=TileLayout(S[shape])), + A: T.Tensor(shape, A_dtype, layout=TileLayout(S[shape])), + B: T.Tensor(shape, B_dtype, layout=TileLayout(S[shape])), ) -> None: T.device_entry() @@ -903,16 +903,16 @@ def test_cast_warpgroup_local_view(A_dtype, B_dtype): # fmt: off @T.prim_func def test_cast( - A: T.Buffer(g_shape, A_dtype, layout=g_layout), - B: T.Buffer(g_shape, B_dtype, layout=g_layout), + A: T.Tensor(g_shape, A_dtype, layout=g_layout), + B: T.Tensor(g_shape, B_dtype, layout=g_layout), ) -> None: T.device_entry() cta_id = T.cta_id([1]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([N_THREADS]) - reg_src = T.alloc_buffer((LOCAL_LEN,), A_dtype, scope="local") - reg_dst = T.alloc_buffer((LOCAL_LEN,), B_dtype, scope="local") + reg_src = T.alloc_tensor((LOCAL_LEN,), A_dtype, scope="local") + reg_dst = T.alloc_tensor((LOCAL_LEN,), B_dtype, scope="local") for i in T.serial(LOCAL_LEN): reg_src[i] = A[tid_in_wg, i] reg_src_view = reg_src.view(N_THREADS, LOCAL_LEN, layout=cast_layout) @@ -961,8 +961,8 @@ def test_cast_warpgroup_src_layout_to_flat_uses_vec2_intrinsic(A_dtype, B_dtype) # fmt: off @T.prim_func def test_cast( - A: T.Buffer(g_shape, A_dtype, layout=g_layout), - B: T.Buffer(g_shape, B_dtype, layout=g_layout), + A: T.Tensor(g_shape, A_dtype, layout=g_layout), + B: T.Tensor(g_shape, B_dtype, layout=g_layout), ) -> None: T.device_entry() @@ -970,8 +970,8 @@ def test_cast( wg_id = T.warpgroup_id([1]) tid = T.thread_id_in_wg([N_THREADS]) for no in T.unroll(N_CHUNKS): - reg_src = T.alloc_buffer((LOCAL_LEN,), A_dtype, scope="local") - Dreg_chunk = T.alloc_buffer((LOCAL_LEN,), B_dtype, scope="local") + reg_src = T.alloc_tensor((LOCAL_LEN,), A_dtype, scope="local") + Dreg_chunk = T.alloc_tensor((LOCAL_LEN,), B_dtype, scope="local") for i in T.serial(LOCAL_LEN): reg_src[i] = A[tid, no * LOCAL_LEN + i] reg_src_view = reg_src.view(N_THREADS, LOCAL_LEN, layout=wg_local_layout(LOCAL_LEN)) @@ -1020,15 +1020,15 @@ def test_cast_cta_local_view(A_dtype, B_dtype): # fmt: off @T.prim_func def test_cast( - A: T.Buffer(g_shape, A_dtype, layout=g_layout), - B: T.Buffer(g_shape, B_dtype, layout=g_layout), + A: T.Tensor(g_shape, A_dtype, layout=g_layout), + B: T.Tensor(g_shape, B_dtype, layout=g_layout), ) -> None: T.device_entry() cta_id = T.cta_id([1]) tx_var = T.thread_id([N_THREADS]) - reg_src = T.alloc_buffer((LOCAL_LEN,), A_dtype, scope="local") - reg_dst = T.alloc_buffer((LOCAL_LEN,), B_dtype, scope="local") + reg_src = T.alloc_tensor((LOCAL_LEN,), A_dtype, scope="local") + reg_dst = T.alloc_tensor((LOCAL_LEN,), B_dtype, scope="local") for i in T.serial(LOCAL_LEN): reg_src[i] = A[tx_var, i] reg_src_view = reg_src.view(N_THREADS, LOCAL_LEN, layout=cast_layout) @@ -1072,15 +1072,15 @@ def test_cast_local_view_sliced(A_dtype, B_dtype, slice_start, slice_end): # fmt: off @T.prim_func def kernel( - A: T.Buffer(g_shape, A_dtype, layout=g_layout), - B: T.Buffer(g_shape, B_dtype, layout=g_layout), + A: T.Tensor(g_shape, A_dtype, layout=g_layout), + B: T.Tensor(g_shape, B_dtype, layout=g_layout), ) -> None: T.device_entry() _bx = T.cta_id([1]) tx = T.thread_id([N_THREADS]) - reg_src = T.alloc_buffer((LOCAL_LEN,), A_dtype, scope="local") - reg_dst = T.alloc_buffer((LOCAL_LEN,), B_dtype, scope="local") + reg_src = T.alloc_tensor((LOCAL_LEN,), A_dtype, scope="local") + reg_dst = T.alloc_tensor((LOCAL_LEN,), B_dtype, scope="local") for i in T.serial(LOCAL_LEN): reg_src[i] = A[tx, i] reg_src_view = reg_src.view(N_THREADS, LOCAL_LEN, layout=cast_layout) @@ -1184,15 +1184,15 @@ def test_cast_mixed_axes_and_subregion(slice_start, slice_end): @T.prim_func def kernel( - A: T.Buffer(full_shape, "float32", layout=g_layout), - B: T.Buffer(full_shape, "float16", layout=g_layout), + A: T.Tensor(full_shape, "float32", layout=g_layout), + B: T.Tensor(full_shape, "float16", layout=g_layout), ) -> None: T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([N_WARPS]) lane_id = T.lane_id([LANES]) - reg_src = T.alloc_buffer((LOCAL_LEN,), "float32", scope="local") - reg_dst = T.alloc_buffer((LOCAL_LEN,), "float16", scope="local") + reg_src = T.alloc_tensor((LOCAL_LEN,), "float32", scope="local") + reg_dst = T.alloc_tensor((LOCAL_LEN,), "float16", scope="local") j, k = lane_id // 4, lane_id % 4 for i in T.serial(LOCAL_LEN): reg_src[i] = A[j, warp_id, k, i] @@ -1266,15 +1266,15 @@ def test_cast_validate_extent_mismatch_rejected(): @T.prim_func def kernel( - A: T.Buffer(view_shape, "float32", layout=g_layout), - B: T.Buffer(view_shape, "float16", layout=g_layout), + A: T.Tensor(view_shape, "float32", layout=g_layout), + B: T.Tensor(view_shape, "float16", layout=g_layout), ) -> None: T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([2]) lane_id = T.lane_id([32]) - reg_src = T.alloc_buffer((8,), "float32", scope="local") - reg_dst = T.alloc_buffer((8,), "float16", scope="local") + reg_src = T.alloc_tensor((8,), "float32", scope="local") + reg_dst = T.alloc_tensor((8,), "float16", scope="local") j, k = lane_id // 4, lane_id % 4 for i in T.serial(8): reg_src[i] = A[warp_id, j, k, i] @@ -1308,13 +1308,13 @@ def test_unary_exp_f16_shared_scalar_fallback_dispatch(): @T.prim_func def k( - A: T.Buffer(shape, "float16", layout=lay), B: T.Buffer(shape, "float16", layout=lay) + A: T.Tensor(shape, "float16", layout=lay), B: T.Tensor(shape, "float16", layout=lay) ) -> None: T.device_entry() _bx = T.cta_id([1]) _tx = T.thread_id([64]) - sa = T.alloc_buffer(shape, "float16", scope="shared", layout=lay) - sb = T.alloc_buffer(shape, "float16", scope="shared", layout=lay) + sa = T.alloc_tensor(shape, "float16", scope="shared", layout=lay) + sb = T.alloc_tensor(shape, "float16", scope="shared", layout=lay) Tx.copy(sa, A) Tx.cta.exp(sb, sa) Tx.copy(B, sb) @@ -1341,13 +1341,13 @@ def test_cast_vec2_packed_dispatch(src_dtype, dst_dtype, intrinsic): @T.prim_func def k( - A: T.Buffer(shape, src_dtype, layout=lay), B: T.Buffer(shape, dst_dtype, layout=lay) + A: T.Tensor(shape, src_dtype, layout=lay), B: T.Tensor(shape, dst_dtype, layout=lay) ) -> None: T.device_entry() _bx = T.cta_id([1]) tx = T.thread_id([64]) - ra = T.alloc_buffer(shape[1:], src_dtype, scope="local", layout=TileLayout(S[shape[1:]])) - rb = T.alloc_buffer(shape[1:], dst_dtype, scope="local", layout=TileLayout(S[shape[1:]])) + ra = T.alloc_tensor(shape[1:], src_dtype, scope="local", layout=TileLayout(S[shape[1:]])) + rb = T.alloc_tensor(shape[1:], dst_dtype, scope="local", layout=TileLayout(S[shape[1:]])) Tx.copy(ra, A[tx]) Tx.cast(rb, ra) Tx.copy(B[tx], rb) @@ -1380,20 +1380,20 @@ def test_cast_wg_rejects_thread_local_view(): @T.prim_func def kernel( - A: T.Buffer((_SL_ROWS, _SL_COLS), "float32", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), - B: T.Buffer((_SL_ROWS, _SL_COLS), "float16", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), + A: T.Tensor((_SL_ROWS, _SL_COLS), "float32", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), + B: T.Tensor((_SL_ROWS, _SL_COLS), "float16", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), ) -> None: T.device_entry() _bx = T.cta_id([1]) _wg = T.warpgroup_id([1]) tid = T.thread_id_in_wg([_SL_ROWS]) - src = T.alloc_buffer( + src = T.alloc_tensor( (_SL_ROWS, _SL_COLS), "float32", scope="local", layout=TileLayout(S[(_SL_ROWS, _SL_COLS) : (1 @ tid_in_wg, 1)]), ) - dst = T.alloc_buffer( + dst = T.alloc_tensor( (_SL_ROWS, _SL_COLS), "float16", scope="local", @@ -1416,19 +1416,19 @@ def test_cast_cta_rejects_thread_local_view(): @T.prim_func def kernel( - A: T.Buffer((_SL_ROWS, _SL_COLS), "float32", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), - B: T.Buffer((_SL_ROWS, _SL_COLS), "float16", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), + A: T.Tensor((_SL_ROWS, _SL_COLS), "float32", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), + B: T.Tensor((_SL_ROWS, _SL_COLS), "float16", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), ) -> None: T.device_entry() _bx = T.cta_id([1]) tx_var = T.thread_id([_SL_ROWS]) - src = T.alloc_buffer( + src = T.alloc_tensor( (_SL_ROWS, _SL_COLS), "float32", scope="local", layout=TileLayout(S[(_SL_ROWS, _SL_COLS) : (1 @ tx, 1)]), ) - dst = T.alloc_buffer( + dst = T.alloc_tensor( (_SL_ROWS, _SL_COLS), "float16", scope="local", @@ -1452,20 +1452,20 @@ def test_cast_wg_rejects_partial_thread_coverage(): @T.prim_func def kernel( - A: T.Buffer((half, _SL_COLS), "float32", layout=TileLayout(S[half, _SL_COLS])), - B: T.Buffer((half, _SL_COLS), "float16", layout=TileLayout(S[half, _SL_COLS])), + A: T.Tensor((half, _SL_COLS), "float32", layout=TileLayout(S[half, _SL_COLS])), + B: T.Tensor((half, _SL_COLS), "float16", layout=TileLayout(S[half, _SL_COLS])), ) -> None: T.device_entry() _bx = T.cta_id([1]) _wg = T.warpgroup_id([1]) tid = T.thread_id_in_wg([_SL_ROWS]) - src = T.alloc_buffer( + src = T.alloc_tensor( (half, _SL_COLS), "float32", scope="local", layout=TileLayout(S[(half, _SL_COLS) : (1 @ tid_in_wg, 1)]), ) - dst = T.alloc_buffer( + dst = T.alloc_tensor( (half, _SL_COLS), "float16", scope="local", @@ -1488,20 +1488,20 @@ def test_cast_wg_accepts_wg_level_layout(): @T.prim_func def kernel( - A: T.Buffer((_SL_ROWS, _SL_COLS), "float32", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), - B: T.Buffer((_SL_ROWS, _SL_COLS), "float16", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), + A: T.Tensor((_SL_ROWS, _SL_COLS), "float32", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), + B: T.Tensor((_SL_ROWS, _SL_COLS), "float16", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), ) -> None: T.device_entry() _bx = T.cta_id([1]) _wg = T.warpgroup_id([1]) tid = T.thread_id_in_wg([_SL_ROWS]) - src = T.alloc_buffer( + src = T.alloc_tensor( (_SL_ROWS, _SL_COLS), "float32", scope="local", layout=TileLayout(S[(_SL_ROWS, _SL_COLS) : (1 @ tid_in_wg, 1)]), ) - dst = T.alloc_buffer( + dst = T.alloc_tensor( (_SL_ROWS, _SL_COLS), "float16", scope="local", @@ -1523,16 +1523,16 @@ def test_cast_thread_accepts_local_view(): @T.prim_func def kernel( - A: T.Buffer((_SL_ROWS, _SL_COLS), "float32", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), - B: T.Buffer((_SL_ROWS, _SL_COLS), "float16", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), + A: T.Tensor((_SL_ROWS, _SL_COLS), "float32", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), + B: T.Tensor((_SL_ROWS, _SL_COLS), "float16", layout=TileLayout(S[_SL_ROWS, _SL_COLS])), ) -> None: T.device_entry() _bx = T.cta_id([1]) tx_var = T.thread_id([_SL_ROWS]) - src = T.alloc_buffer( + src = T.alloc_tensor( (_SL_COLS,), "float32", scope="local", layout=TileLayout(S[(_SL_COLS,)]) ) - dst = T.alloc_buffer( + dst = T.alloc_tensor( (_SL_COLS,), "float16", scope="local", layout=TileLayout(S[(_SL_COLS,)]) ) for i in T.serial(_SL_COLS): @@ -1575,7 +1575,7 @@ def _tcgen05_cast_warpgroup_kernel(): m, k = _TCGEN05_M, _TCGEN05_K @T.prim_func - def kernel(A: T.Buffer((m, k), "bfloat16"), B: T.Buffer((m, k), "float32")) -> None: + def kernel(A: T.Tensor((m, k), "bfloat16"), B: T.Tensor((m, k), "float32")) -> None: T.device_entry() T.cta_id([1]) T.warpgroup_id([1]) @@ -1583,8 +1583,8 @@ def kernel(A: T.Buffer((m, k), "bfloat16"), B: T.Buffer((m, k), "float32")) -> N T.lane_id([32]) tid_wg = T.thread_id_in_wg([128]) lane = T.lane_id([32]) - a_bf16 = T.alloc_buffer((m, k), "bfloat16", scope="local", layout=atom_layout) - a_fp32 = T.alloc_buffer((m, k), "float32", scope="local", layout=atom_layout) + a_bf16 = T.alloc_tensor((m, k), "bfloat16", scope="local", layout=atom_layout) + a_fp32 = T.alloc_tensor((m, k), "float32", scope="local", layout=atom_layout) reg_in = a_bf16.local(_TCGEN05_REGS_PER_THREAD) reg_out = a_fp32.local(_TCGEN05_REGS_PER_THREAD) for r in T.serial(_TCGEN05_REGS_PER_THREAD): diff --git a/tests/python/tirx/operator/tile_primitive/cuda/gemm/test_gemm_mma_m16n8k_.py b/tests/python/tirx/operator/tile_primitive/cuda/gemm/test_gemm_mma_m16n8k_.py index c8b4200f427a..89b46da5f42b 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/gemm/test_gemm_mma_m16n8k_.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/gemm/test_gemm_mma_m16n8k_.py @@ -110,24 +110,24 @@ def gemm(): _cta = T.cta_id([1]) _warp = T.warp_id([1]) _lane = T.lane_id([32]) - A = T.alloc_buffer((M, K), dtype, scope="local", layout=Al) - B = T.alloc_buffer((K, N), dtype, scope="local", layout=Bl) - C = T.alloc_buffer((M, N), "float32", scope="local", layout=Dl) - D = T.alloc_buffer((M, N), "float32", scope="local", layout=Dl) + A = T.alloc_tensor((M, K), dtype, scope="local", layout=Al) + B = T.alloc_tensor((K, N), dtype, scope="local", layout=Bl) + C = T.alloc_tensor((M, N), "float32", scope="local", layout=Dl) + D = T.alloc_tensor((M, N), "float32", scope="local", layout=Dl) Tx.warp.gemm(D, A, B, C, transpose_A=False, transpose_B=False, alpha=1.0, beta=beta) return gemm @T.prim_func - def gemm(D_g: T.Buffer((M, N), "float32")): + def gemm(D_g: T.Tensor((M, N), "float32")): T.device_entry() _cta = T.cta_id([1]) _warp = T.warp_id([1]) lane = T.lane_id([32]) - A = T.alloc_buffer((M, K), dtype, scope="local", layout=Al) - B = T.alloc_buffer((K, N), dtype, scope="local", layout=Bl) - C = T.alloc_buffer((M, N), "float32", scope="local", layout=Dl) - D = T.alloc_buffer((M, N), "float32", scope="local", layout=Dl) + A = T.alloc_tensor((M, K), dtype, scope="local", layout=Al) + B = T.alloc_tensor((K, N), dtype, scope="local", layout=Bl) + C = T.alloc_tensor((M, N), "float32", scope="local", layout=Dl) + D = T.alloc_tensor((M, N), "float32", scope="local", layout=Dl) Tx.warp.gemm(D, A, B, C, transpose_A=False, transpose_B=False, alpha=1.0, beta=beta) # Decode D's per-thread registers (c = ((mt*Nt + nt)*2 + rM)*2 + rN) # back to logical (M, N) and store, exercising the whole tiling. @@ -150,10 +150,10 @@ def gemm_min(): T.device_entry() _cta = T.cta_id([1]) _tid = T.thread_id([32]) - D = T.alloc_buffer((16, 8), "float32", scope="local", layout=D_FRAG) - C = T.alloc_buffer((16, 8), "float32", scope="local", layout=D_FRAG) - A = T.alloc_buffer((16, 16), dtype, scope="local", layout=A_FRAG) - B = T.alloc_buffer((16, 8), dtype, scope="local", layout=B_FRAG) + D = T.alloc_tensor((16, 8), "float32", scope="local", layout=D_FRAG) + C = T.alloc_tensor((16, 8), "float32", scope="local", layout=D_FRAG) + A = T.alloc_tensor((16, 16), dtype, scope="local", layout=A_FRAG) + B = T.alloc_tensor((16, 8), dtype, scope="local", layout=B_FRAG) Tx.warp.gemm(D, A, B, C, transpose_A=False, transpose_B=False, alpha=alpha, beta=beta) return gemm_min @@ -174,10 +174,10 @@ def gemm(): _cta = T.cta_id([1]) _warp = T.warp_id([1]) _lane = T.lane_id([32]) - A = T.alloc_buffer(A_shape, "float16", scope="local", layout=Al) - B = T.alloc_buffer(B_shape, "float16", scope="local", layout=Bl) - C = T.alloc_buffer((16, 8), "float32", scope="local", layout=D_FRAG) - D = T.alloc_buffer((16, 8), "float32", scope="local", layout=D_FRAG) + A = T.alloc_tensor(A_shape, "float16", scope="local", layout=Al) + B = T.alloc_tensor(B_shape, "float16", scope="local", layout=Bl) + C = T.alloc_tensor((16, 8), "float32", scope="local", layout=D_FRAG) + D = T.alloc_tensor((16, 8), "float32", scope="local", layout=D_FRAG) Tx.warp.gemm( D, A, @@ -192,15 +192,15 @@ def gemm(): return gemm @T.prim_func - def gemm(D_g: T.Buffer((16, 8), "float32")): + def gemm(D_g: T.Tensor((16, 8), "float32")): T.device_entry() _cta = T.cta_id([1]) _warp = T.warp_id([1]) lane = T.lane_id([32]) - A = T.alloc_buffer(A_shape, "float16", scope="local", layout=Al) - B = T.alloc_buffer(B_shape, "float16", scope="local", layout=Bl) - C = T.alloc_buffer((16, 8), "float32", scope="local", layout=D_FRAG) - D = T.alloc_buffer((16, 8), "float32", scope="local", layout=D_FRAG) + A = T.alloc_tensor(A_shape, "float16", scope="local", layout=Al) + B = T.alloc_tensor(B_shape, "float16", scope="local", layout=Bl) + C = T.alloc_tensor((16, 8), "float32", scope="local", layout=D_FRAG) + D = T.alloc_tensor((16, 8), "float32", scope="local", layout=D_FRAG) Tx.warp.gemm( D, A, @@ -226,10 +226,10 @@ def gemm_min(): T.device_entry() _cta = T.cta_id([1]) _tid = T.thread_id([32]) - D = T.alloc_buffer((16, 8), d_dtype, scope="local", layout=D_FRAG) - C = T.alloc_buffer((16, 8), c_dtype, scope="local", layout=D_FRAG) - A = T.alloc_buffer((16, 16), a_dtype, scope="local", layout=A_FRAG) - B = T.alloc_buffer((16, 8), b_dtype, scope="local", layout=B_FRAG) + D = T.alloc_tensor((16, 8), d_dtype, scope="local", layout=D_FRAG) + C = T.alloc_tensor((16, 8), c_dtype, scope="local", layout=D_FRAG) + A = T.alloc_tensor((16, 16), a_dtype, scope="local", layout=A_FRAG) + B = T.alloc_tensor((16, 8), b_dtype, scope="local", layout=B_FRAG) Tx.warp.gemm(D, A, B, C, transpose_A=False, transpose_B=False, alpha=1.0, beta=0.0) return gemm_min @@ -253,19 +253,19 @@ def _build_tiled_numeric(Mt, Nt, Kt, kinst, beta, dtype): @T.prim_func def gemm( - A_g: T.Buffer((M, K), dtype), - B_g: T.Buffer((K, N), dtype), - C_g: T.Buffer((M, N), "float32"), - D_g: T.Buffer((M, N), "float32"), + A_g: T.Tensor((M, K), dtype), + B_g: T.Tensor((K, N), dtype), + C_g: T.Tensor((M, N), "float32"), + D_g: T.Tensor((M, N), "float32"), ): T.device_entry() _cta = T.cta_id([1]) _warp = T.warp_id([1]) lane = T.lane_id([32]) - A_f = T.alloc_buffer((M, K), dtype, scope="local", layout=Al) - B_f = T.alloc_buffer((K, N), dtype, scope="local", layout=Bl) - C_f = T.alloc_buffer((M, N), "float32", scope="local", layout=Dl) - D_f = T.alloc_buffer((M, N), "float32", scope="local", layout=Dl) + A_f = T.alloc_tensor((M, K), dtype, scope="local", layout=Al) + B_f = T.alloc_tensor((K, N), dtype, scope="local", layout=Bl) + C_f = T.alloc_tensor((M, N), "float32", scope="local", layout=Dl) + D_f = T.alloc_tensor((M, N), "float32", scope="local", layout=Dl) A_reg = A_f.local(Mt, Kt, kHi_n, 2, KP) for mt, kt, kHi, rM, kp in T.grid(Mt, Kt, kHi_n, 2, KP): A_reg[mt, kt, kHi, rM, kp] = A_g[ @@ -306,17 +306,17 @@ def _build_transpose_numeric(transpose_A, transpose_B, dtype="float16"): @T.prim_func def gemm( - A_g: T.Buffer(A_shape, dtype), - B_g: T.Buffer(B_shape, dtype), - D_g: T.Buffer((16, 8), "float32"), + A_g: T.Tensor(A_shape, dtype), + B_g: T.Tensor(B_shape, dtype), + D_g: T.Tensor((16, 8), "float32"), ): T.device_entry() _cta = T.cta_id([1]) _warp = T.warp_id([1]) lane = T.lane_id([32]) - A_f = T.alloc_buffer(A_shape, dtype, scope="local", layout=Al) - B_f = T.alloc_buffer(B_shape, dtype, scope="local", layout=Bl) - D_f = T.alloc_buffer((16, 8), "float32", scope="local", layout=D_FRAG) + A_f = T.alloc_tensor(A_shape, dtype, scope="local", layout=Al) + B_f = T.alloc_tensor(B_shape, dtype, scope="local", layout=Bl) + D_f = T.alloc_tensor((16, 8), "float32", scope="local", layout=D_FRAG) A_reg = A_f.local(2, 2, 2) if transpose_A: # A_KM_FRAG: buffer is [K, M]. @@ -441,17 +441,17 @@ def test_cuda_gemm_mma_numerical(dtype): @T.prim_func def gemm( - A_g: T.Buffer((16, 16), dtype), - B_g: T.Buffer((16, 8), dtype), - D_g: T.Buffer((16, 8), "float32"), + A_g: T.Tensor((16, 16), dtype), + B_g: T.Tensor((16, 8), dtype), + D_g: T.Tensor((16, 8), "float32"), ): T.device_entry() _cta = T.cta_id([1]) _warp = T.warp_id([1]) lane = T.lane_id([32]) - A_f = T.alloc_buffer((16, 16), dtype, scope="local", layout=A_FRAG) - B_f = T.alloc_buffer((16, 8), dtype, scope="local", layout=B_FRAG) - D_f = T.alloc_buffer((16, 8), "float32", scope="local", layout=D_FRAG) + A_f = T.alloc_tensor((16, 16), dtype, scope="local", layout=A_FRAG) + B_f = T.alloc_tensor((16, 8), dtype, scope="local", layout=B_FRAG) + D_f = T.alloc_tensor((16, 8), "float32", scope="local", layout=D_FRAG) A_reg = A_f.local(8) for s in T.unroll(8): # Physical register order: s = 4*kHi + 2*rM + kp. diff --git a/tests/python/tirx/operator/tile_primitive/cuda/gemm_async/test_gemm_async.py b/tests/python/tirx/operator/tile_primitive/cuda/gemm_async/test_gemm_async.py index 8b4abe8339e6..1c7914045d00 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/gemm_async/test_gemm_async.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/gemm_async/test_gemm_async.py @@ -239,7 +239,7 @@ def test_gemm_tcgen05_cta_group_1(task): # fmt: off @T.prim_func def gemm_async( - A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), C: T.Buffer(C_shape, C_dtype) + A: T.Tensor(A_shape, A_dtype), B: T.Tensor(B_shape, B_dtype), C: T.Tensor(C_shape, C_dtype) ) -> None: T.device_entry() @@ -248,8 +248,8 @@ def gemm_async( wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape, A_dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape, B_dtype, scope="shared", layout=B_layout) + A_smem = T.alloc_tensor(A_shape, A_dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape, B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") @@ -265,7 +265,7 @@ def gemm_async( T.address_of(tmem_addr), T.uint32(cols_alloc) ) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, C_shape[1]), C_dtype, scope="tmem", @@ -366,7 +366,7 @@ def test_gemm_tcgen05_cta_group_1_layout_f_m64(): # fmt: off @T.prim_func def gemm_layout_f( - A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), C: T.Buffer(C_shape, C_dtype) + A: T.Tensor(A_shape, A_dtype), B: T.Tensor(B_shape, B_dtype), C: T.Tensor(C_shape, C_dtype) ) -> None: T.device_entry() @@ -376,8 +376,8 @@ def gemm_layout_f( tid_in_wg = T.thread_id_in_wg([128]) lane_id = T.lane_id([32]) - A_smem = T.alloc_buffer(A_shape, A_dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape, B_dtype, scope="shared", layout=B_layout) + A_smem = T.alloc_tensor(A_shape, A_dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape, B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") @@ -394,7 +394,7 @@ def gemm_layout_f( ) T.cuda.cta_sync() # Layout F C operand — the path under test. - tmem = T.decl_buffer( + tmem = T.decl_tensor( (64, N), C_dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=c_layout ) @@ -515,7 +515,7 @@ def test_gemm_tcgen05_cta_group_2(task): # fmt: off @T.prim_func def gemm_async( - A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), C: T.Buffer(C_shape, C_dtype) + A: T.Tensor(A_shape, A_dtype), B: T.Tensor(B_shape, B_dtype), C: T.Tensor(C_shape, C_dtype) ) -> None: T.device_entry() @@ -525,8 +525,8 @@ def gemm_async( wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape_per_cta, A_dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape_per_cta, B_dtype, scope="shared", layout=B_layout) + A_smem = T.alloc_tensor(A_shape_per_cta, A_dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape_per_cta, B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") @@ -534,7 +534,7 @@ def gemm_async( ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64"), "shared"))] = T.reinterpret( PointerType(PrimType("uint64"), "shared"), _mapa(tma_mbar.ptr_to([0]), 0) ) - tma_mbar_cta_0 = T.decl_buffer([1], "uint64", data=ptr, scope="shared") + tma_mbar_cta_0 = T.decl_tensor([1], "uint64", data=ptr, scope="shared") if tid_in_wg == 0: T.ptx.mbarrier.init.shared.b64(tma_mbar.ptr_to([0]), T.uint32(1)) @@ -544,7 +544,7 @@ def gemm_async( T.ptx.tcgen05.alloc.cta_group__2.sync.aligned.shared__cta.b32( T.address_of(tmem_addr), T.uint32(cols_alloc) ) - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, C_shape[1]), C_dtype, scope="tmem", @@ -684,9 +684,9 @@ def test_gemm_tcgen05_cta_group_2_layout_b(): # fmt: off @T.prim_func def gemm_async( - A: T.Buffer((M_per_cta * 2, K), A_dtype), - B: T.Buffer((N_logical, K), B_dtype), - C: T.Buffer(C_shape, C_dtype), + A: T.Tensor((M_per_cta * 2, K), A_dtype), + B: T.Tensor((N_logical, K), B_dtype), + C: T.Tensor(C_shape, C_dtype), ) -> None: T.device_entry() @@ -696,8 +696,8 @@ def gemm_async( wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape, A_dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape, B_dtype, scope="shared", layout=B_layout) + A_smem = T.alloc_tensor(A_shape, A_dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape, B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") @@ -705,7 +705,7 @@ def gemm_async( ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64"), "shared"))] = T.reinterpret( PointerType(PrimType("uint64"), "shared"), _mapa(tma_mbar.ptr_to([0]), 0) ) - tma_mbar_cta_0 = T.decl_buffer([1], "uint64", data=ptr, scope="shared") + tma_mbar_cta_0 = T.decl_tensor([1], "uint64", data=ptr, scope="shared") if tid_in_wg == 0: T.ptx.mbarrier.init.shared.b64(tma_mbar.ptr_to([0]), T.uint32(1)) @@ -715,7 +715,7 @@ def gemm_async( T.ptx.tcgen05.alloc.cta_group__2.sync.aligned.shared__cta.b32( T.address_of(tmem_addr), T.uint32(cols_alloc) ) - tmem = T.decl_buffer( + tmem = T.decl_tensor( (M_per_cta, N_logical), C_dtype, scope="tmem", @@ -723,7 +723,7 @@ def gemm_async( layout=TileLayout(S[(M_per_cta, 2, N_half) : (1 @ TLane, 64 @ TLane, 1 @ TCol)]), ) # Physical TMEM view for readback: (128, N_half) standard layout - tmem_phys = T.decl_buffer( + tmem_phys = T.decl_tensor( (128, N_half), C_dtype, scope="tmem", @@ -842,9 +842,9 @@ def test_gemm_tcgen05_cta_group_2_datapath_b_readback(): # fmt: off @T.prim_func def gemm_async( - A: T.Buffer((m_per_cta * 2, k), input_dtype), - B: T.Buffer((n_logical, k), input_dtype), - C: T.Buffer(c_shape, c_dtype), + A: T.Tensor((m_per_cta * 2, k), input_dtype), + B: T.Tensor((n_logical, k), input_dtype), + C: T.Tensor(c_shape, c_dtype), ) -> None: T.device_entry() @@ -854,8 +854,8 @@ def gemm_async( wg_id = T.warpgroup_id([1]) tid = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(a_shape, a_dtype, scope="shared", layout=a_layout) - B_smem = T.alloc_buffer(b_shape, b_dtype, scope="shared", layout=b_layout) + A_smem = T.alloc_tensor(a_shape, a_dtype, scope="shared", layout=a_layout) + B_smem = T.alloc_tensor(b_shape, b_dtype, scope="shared", layout=b_layout) tmem_addr = T.alloc_shared([1], "uint32") mma_mbar = T.alloc_shared([1], "uint64") @@ -868,7 +868,7 @@ def gemm_async( T.ptx.fence.mbarrier_init.release.cluster() T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( (m_per_cta, n_logical), c_dtype, scope="tmem", @@ -1025,7 +1025,7 @@ def test_gemm_block_scaled_fp8_cta_group_1(task): # fmt: off @T.prim_func - def gemm_async_fn(A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), C: T.Buffer(C_shape, C_dtype), SFA_in: T.Buffer((128,), 'uint32'), SFB_in: T.Buffer((128,), 'uint32')) -> None: # noqa: E501 + def gemm_async_fn(A: T.Tensor(A_shape, A_dtype), B: T.Tensor(B_shape, B_dtype), C: T.Tensor(C_shape, C_dtype), SFA_in: T.Tensor((128,), 'uint32'), SFB_in: T.Tensor((128,), 'uint32')) -> None: # noqa: E501 T.device_entry() warp_id = T.warp_id([(1) * 4]) @@ -1033,17 +1033,17 @@ def gemm_async_fn(A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape, A_dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape, B_dtype, scope="shared", layout=B_layout) - SFA_smem = T.alloc_buffer((4, 32), "uint32", scope="shared", layout=SF_smem_layout) - SFB_smem = T.alloc_buffer((4, 32), "uint32", scope="shared", layout=SF_smem_layout) + A_smem = T.alloc_tensor(A_shape, A_dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape, B_dtype, scope="shared", layout=B_layout) + SFA_smem = T.alloc_tensor((4, 32), "uint32", scope="shared", layout=SF_smem_layout) + SFB_smem = T.alloc_tensor((4, 32), "uint32", scope="shared", layout=SF_smem_layout) SFA_smem_post = SFA_smem.view(4, 32, layout=SF_smem_post_layout) SFB_smem_post = SFB_smem.view(4, 32, layout=SF_smem_post_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") - descSFA = T.alloc_buffer((1,), "uint64", scope="local") - descSFB = T.alloc_buffer((1,), "uint64", scope="local") + descSFA = T.alloc_tensor((1,), "uint64", scope="local") + descSFB = T.alloc_tensor((1,), "uint64", scope="local") if tid_in_wg == 0: T.ptx.mbarrier.init.shared.b64(tma_mbar.ptr_to([0]), T.uint32(1)) @@ -1057,9 +1057,9 @@ def gemm_async_fn(A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), ) T.cuda.cta_sync() - tmem = T.decl_buffer((128, C_shape[1]), C_dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, C_shape[1]) : (1 @ TLane, 1 @ TCol)])) # noqa: E501 - sfa_tmem = T.decl_buffer((M, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFA_TMEM_START, layout=sfa_layout) # noqa: E501 - sfb_tmem = T.decl_buffer((N, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFB_TMEM_START, layout=sfb_layout) # noqa: E501 + tmem = T.decl_tensor((128, C_shape[1]), C_dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, C_shape[1]) : (1 @ TLane, 1 @ TCol)])) # noqa: E501 + sfa_tmem = T.decl_tensor((M, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFA_TMEM_START, layout=sfa_layout) # noqa: E501 + sfb_tmem = T.decl_tensor((N, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFB_TMEM_START, layout=sfb_layout) # noqa: E501 # TMA load A and B from global to shared if tid_in_wg == 0: @@ -1221,7 +1221,7 @@ def test_gemm_block_scaled_fp8_cta_group_2(task): # fmt: off @T.prim_func - def gemm_async_fn(A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), C: T.Buffer(C_shape, C_dtype), SFA_in: T.Buffer((M_total,), 'uint32'), SFB_in: T.Buffer((128,), 'uint32')) -> None: # noqa: E501 + def gemm_async_fn(A: T.Tensor(A_shape, A_dtype), B: T.Tensor(B_shape, B_dtype), C: T.Tensor(C_shape, C_dtype), SFA_in: T.Tensor((M_total,), 'uint32'), SFB_in: T.Tensor((128,), 'uint32')) -> None: # noqa: E501 T.device_entry() warp_id = T.warp_id([(1) * 4]) @@ -1230,20 +1230,20 @@ def gemm_async_fn(A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape_per_cta, A_dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape_per_cta, B_dtype, scope="shared", layout=B_layout) - SFA_smem = T.alloc_buffer((4, 32), "uint32", scope="shared", layout=SF_smem_layout) - SFB_smem = T.alloc_buffer((4, 32), "uint32", scope="shared", layout=SF_smem_layout) + A_smem = T.alloc_tensor(A_shape_per_cta, A_dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape_per_cta, B_dtype, scope="shared", layout=B_layout) + SFA_smem = T.alloc_tensor((4, 32), "uint32", scope="shared", layout=SF_smem_layout) + SFB_smem = T.alloc_tensor((4, 32), "uint32", scope="shared", layout=SF_smem_layout) SFA_smem_post = SFA_smem.view(4, 32, layout=SF_smem_post_layout) SFB_smem_post = SFB_smem.view(4, 32, layout=SF_smem_post_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") - descSFA = T.alloc_buffer((1,), "uint64", scope="local") - descSFB = T.alloc_buffer((1,), "uint64", scope="local") + descSFA = T.alloc_tensor((1,), "uint64", scope="local") + descSFB = T.alloc_tensor((1,), "uint64", scope="local") ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64"), "shared"))] = T.reinterpret(PointerType(PrimType("uint64"), "shared"), _mapa(tma_mbar.ptr_to([0]), 0)) # noqa: E501 - tma_mbar_cta_0 = T.decl_buffer([1], "uint64", data=ptr, scope="shared") + tma_mbar_cta_0 = T.decl_tensor([1], "uint64", data=ptr, scope="shared") if tid_in_wg == 0: T.ptx.mbarrier.init.shared.b64(tma_mbar.ptr_to([0]), T.uint32(1)) @@ -1253,10 +1253,10 @@ def gemm_async_fn(A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), T.ptx.tcgen05.alloc.cta_group__2.sync.aligned.shared__cta.b32( T.address_of(tmem_addr), T.uint32(cols_alloc) ) - tmem = T.decl_buffer((128, C_shape[1]), C_dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, C_shape[1]) : (1 @ TLane, 1 @ TCol)])) # noqa: E501 + tmem = T.decl_tensor((128, C_shape[1]), C_dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, C_shape[1]) : (1 @ TLane, 1 @ TCol)])) # noqa: E501 - sfa_tmem = T.decl_buffer((128, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFA_TMEM_START, layout=sf_layout) # noqa: E501 - sfb_tmem = T.decl_buffer((128, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFB_TMEM_START, layout=sf_layout) # noqa: E501 + sfa_tmem = T.decl_tensor((128, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFA_TMEM_START, layout=sf_layout) # noqa: E501 + sfb_tmem = T.decl_tensor((128, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFB_TMEM_START, layout=sf_layout) # noqa: E501 T.ptx.fence.mbarrier_init.release.cluster() T.ptx.fence.proxy.async_.shared__cta() @@ -1419,7 +1419,7 @@ def test_gemm_block_scaled_nvfp4_cta_group_1(): # fmt: off @T.prim_func - def gemm_async_fn(A_packed: T.Buffer(A_packed_shape, 'uint8'), B_packed: T.Buffer(B_packed_shape, 'uint8'), C: T.Buffer(C_shape, C_dtype), SFA_in: T.Buffer((128,), 'uint32'), SFB_in: T.Buffer((128,), 'uint32')) -> None: # noqa: E501 + def gemm_async_fn(A_packed: T.Tensor(A_packed_shape, 'uint8'), B_packed: T.Tensor(B_packed_shape, 'uint8'), C: T.Tensor(C_shape, C_dtype), SFA_in: T.Tensor((128,), 'uint32'), SFB_in: T.Tensor((128,), 'uint32')) -> None: # noqa: E501 T.device_entry() warp_id = T.warp_id([(1) * 4]) @@ -1427,20 +1427,20 @@ def gemm_async_fn(A_packed: T.Buffer(A_packed_shape, 'uint8'), B_packed: T.Buffe wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem_packed = T.alloc_buffer(A_packed_shape, "uint8", scope="shared", layout=A_uint8_layout) # noqa: E501 - B_smem_packed = T.alloc_buffer(B_packed_shape, "uint8", scope="shared", layout=B_uint8_layout) # noqa: E501 - A_smem = T.decl_buffer(A_fp4_shape, "float4_e2m1fn", data=A_smem_packed.data, scope="shared", layout=A_fp4_layout) # noqa: E501 - B_smem = T.decl_buffer(B_fp4_shape, "float4_e2m1fn", data=B_smem_packed.data, scope="shared", layout=B_fp4_layout) # noqa: E501 + A_smem_packed = T.alloc_tensor(A_packed_shape, "uint8", scope="shared", layout=A_uint8_layout) # noqa: E501 + B_smem_packed = T.alloc_tensor(B_packed_shape, "uint8", scope="shared", layout=B_uint8_layout) # noqa: E501 + A_smem = T.decl_tensor(A_fp4_shape, "float4_e2m1fn", data=A_smem_packed.data, scope="shared", layout=A_fp4_layout) # noqa: E501 + B_smem = T.decl_tensor(B_fp4_shape, "float4_e2m1fn", data=B_smem_packed.data, scope="shared", layout=B_fp4_layout) # noqa: E501 - SFA_smem = T.alloc_buffer((4, 32), "uint32", scope="shared", layout=SF_smem_layout) - SFB_smem = T.alloc_buffer((4, 32), "uint32", scope="shared", layout=SF_smem_layout) + SFA_smem = T.alloc_tensor((4, 32), "uint32", scope="shared", layout=SF_smem_layout) + SFB_smem = T.alloc_tensor((4, 32), "uint32", scope="shared", layout=SF_smem_layout) SFA_smem_post = SFA_smem.view(4, 32, layout=SF_smem_post_layout) SFB_smem_post = SFB_smem.view(4, 32, layout=SF_smem_post_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") - descSFA = T.alloc_buffer((1,), "uint64", scope="local") - descSFB = T.alloc_buffer((1,), "uint64", scope="local") + descSFA = T.alloc_tensor((1,), "uint64", scope="local") + descSFB = T.alloc_tensor((1,), "uint64", scope="local") if tid_in_wg == 0: T.ptx.mbarrier.init.shared.b64(tma_mbar.ptr_to([0]), T.uint32(1)) @@ -1454,9 +1454,9 @@ def gemm_async_fn(A_packed: T.Buffer(A_packed_shape, 'uint8'), B_packed: T.Buffe ) T.cuda.cta_sync() - tmem = T.decl_buffer((128, C_shape[1]), C_dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, C_shape[1]) : (1 @ TLane, 1 @ TCol)])) # noqa: E501 - sfa_tmem = T.decl_buffer((M, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFA_TMEM_START, layout=sfa_layout) # noqa: E501 - sfb_tmem = T.decl_buffer((N, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFB_TMEM_START, layout=sfb_layout) # noqa: E501 + tmem = T.decl_tensor((128, C_shape[1]), C_dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, C_shape[1]) : (1 @ TLane, 1 @ TCol)])) # noqa: E501 + sfa_tmem = T.decl_tensor((M, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFA_TMEM_START, layout=sfa_layout) # noqa: E501 + sfb_tmem = T.decl_tensor((N, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFB_TMEM_START, layout=sfb_layout) # noqa: E501 # TMA load A and B as uint8 if tid_in_wg == 0: @@ -1599,7 +1599,7 @@ def test_gemm_block_scaled_nvfp4_cta_group_2(): # fmt: off @T.prim_func - def gemm_async_fn(A_packed: T.Buffer(A_packed_shape, 'uint8'), B_packed: T.Buffer(B_packed_shape, 'uint8'), C: T.Buffer(C_shape, C_dtype), SFA_in: T.Buffer((M_total,), 'uint32'), SFB_in: T.Buffer((128,), 'uint32')) -> None: # noqa: E501 + def gemm_async_fn(A_packed: T.Tensor(A_packed_shape, 'uint8'), B_packed: T.Tensor(B_packed_shape, 'uint8'), C: T.Tensor(C_shape, C_dtype), SFA_in: T.Tensor((M_total,), 'uint32'), SFB_in: T.Tensor((128,), 'uint32')) -> None: # noqa: E501 T.device_entry() warp_id = T.warp_id([(1) * 4]) @@ -1608,23 +1608,23 @@ def gemm_async_fn(A_packed: T.Buffer(A_packed_shape, 'uint8'), B_packed: T.Buffe wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem_packed = T.alloc_buffer(A_packed_per_cta, "uint8", scope="shared", layout=A_uint8_layout) # noqa: E501 - B_smem_packed = T.alloc_buffer(B_packed_per_cta, "uint8", scope="shared", layout=B_uint8_layout) # noqa: E501 - A_smem = T.decl_buffer(A_fp4_per_cta, "float4_e2m1fn", data=A_smem_packed.data, scope="shared", layout=A_fp4_layout) # noqa: E501 - B_smem = T.decl_buffer(B_fp4_per_cta, "float4_e2m1fn", data=B_smem_packed.data, scope="shared", layout=B_fp4_layout) # noqa: E501 + A_smem_packed = T.alloc_tensor(A_packed_per_cta, "uint8", scope="shared", layout=A_uint8_layout) # noqa: E501 + B_smem_packed = T.alloc_tensor(B_packed_per_cta, "uint8", scope="shared", layout=B_uint8_layout) # noqa: E501 + A_smem = T.decl_tensor(A_fp4_per_cta, "float4_e2m1fn", data=A_smem_packed.data, scope="shared", layout=A_fp4_layout) # noqa: E501 + B_smem = T.decl_tensor(B_fp4_per_cta, "float4_e2m1fn", data=B_smem_packed.data, scope="shared", layout=B_fp4_layout) # noqa: E501 - SFA_smem = T.alloc_buffer((4, 32), "uint32", scope="shared", layout=SF_smem_layout) - SFB_smem = T.alloc_buffer((4, 32), "uint32", scope="shared", layout=SF_smem_layout) + SFA_smem = T.alloc_tensor((4, 32), "uint32", scope="shared", layout=SF_smem_layout) + SFB_smem = T.alloc_tensor((4, 32), "uint32", scope="shared", layout=SF_smem_layout) SFA_smem_post = SFA_smem.view(4, 32, layout=SF_smem_post_layout) SFB_smem_post = SFB_smem.view(4, 32, layout=SF_smem_post_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") - descSFA = T.alloc_buffer((1,), "uint64", scope="local") - descSFB = T.alloc_buffer((1,), "uint64", scope="local") + descSFA = T.alloc_tensor((1,), "uint64", scope="local") + descSFB = T.alloc_tensor((1,), "uint64", scope="local") ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64"), "shared"))] = T.reinterpret(PointerType(PrimType("uint64"), "shared"), _mapa(tma_mbar.ptr_to([0]), 0)) # noqa: E501 - tma_mbar_cta_0 = T.decl_buffer([1], "uint64", data=ptr, scope="shared") + tma_mbar_cta_0 = T.decl_tensor([1], "uint64", data=ptr, scope="shared") if tid_in_wg == 0: T.ptx.mbarrier.init.shared.b64(tma_mbar.ptr_to([0]), T.uint32(1)) @@ -1634,10 +1634,10 @@ def gemm_async_fn(A_packed: T.Buffer(A_packed_shape, 'uint8'), B_packed: T.Buffe T.ptx.tcgen05.alloc.cta_group__2.sync.aligned.shared__cta.b32( T.address_of(tmem_addr), T.uint32(cols_alloc) ) - tmem = T.decl_buffer((128, C_shape[1]), C_dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, C_shape[1]) : (1 @ TLane, 1 @ TCol)])) # noqa: E501 + tmem = T.decl_tensor((128, C_shape[1]), C_dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, C_shape[1]) : (1 @ TLane, 1 @ TCol)])) # noqa: E501 - sfa_tmem = T.decl_buffer((M_per_cta, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFA_TMEM_START, layout=sfa_layout) # noqa: E501 - sfb_tmem = T.decl_buffer((N_total, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFB_TMEM_START, layout=sfb_layout) # noqa: E501 + sfa_tmem = T.decl_tensor((M_per_cta, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFA_TMEM_START, layout=sfa_layout) # noqa: E501 + sfb_tmem = T.decl_tensor((N_total, sf_mma_k), SF_dtype, scope="tmem", allocated_addr=SFB_TMEM_START, layout=sfb_layout) # noqa: E501 T.ptx.fence.mbarrier_init.release.cluster() T.ptx.fence.proxy.async_.shared__cta() @@ -1808,7 +1808,7 @@ def test_gemm_block_scaled_fp8_sf_id(): # fmt: off @T.prim_func - def gemm_async_fn(A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), C: T.Buffer(C_shape, C_dtype), SFA_in: T.Buffer((128,), 'uint32'), SFB_in: T.Buffer((128,), 'uint32')) -> None: # noqa: E501 + def gemm_async_fn(A: T.Tensor(A_shape, A_dtype), B: T.Tensor(B_shape, B_dtype), C: T.Tensor(C_shape, C_dtype), SFA_in: T.Tensor((128,), 'uint32'), SFB_in: T.Tensor((128,), 'uint32')) -> None: # noqa: E501 T.device_entry() warp_id = T.warp_id([(1) * 4]) @@ -1816,17 +1816,17 @@ def gemm_async_fn(A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape, A_dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape, B_dtype, scope="shared", layout=B_layout) - SFA_smem = T.alloc_buffer((4, 32), "uint32", scope="shared", layout=SF_smem_layout) - SFB_smem = T.alloc_buffer((4, 32), "uint32", scope="shared", layout=SF_smem_layout) + A_smem = T.alloc_tensor(A_shape, A_dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape, B_dtype, scope="shared", layout=B_layout) + SFA_smem = T.alloc_tensor((4, 32), "uint32", scope="shared", layout=SF_smem_layout) + SFB_smem = T.alloc_tensor((4, 32), "uint32", scope="shared", layout=SF_smem_layout) SFA_smem_post = SFA_smem.view(4, 32, layout=SF_smem_post_layout) SFB_smem_post = SFB_smem.view(4, 32, layout=SF_smem_post_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") - descSFA = T.alloc_buffer((1,), "uint64", scope="local") - descSFB = T.alloc_buffer((1,), "uint64", scope="local") + descSFA = T.alloc_tensor((1,), "uint64", scope="local") + descSFB = T.alloc_tensor((1,), "uint64", scope="local") if tid_in_wg == 0: T.ptx.mbarrier.init.shared.b64(tma_mbar.ptr_to([0]), T.uint32(1)) @@ -1840,9 +1840,9 @@ def gemm_async_fn(A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), ) T.cuda.cta_sync() - tmem = T.decl_buffer(C_shape, C_dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, C_shape[1]) : (1 @ TLane, 1 @ TCol)])) # noqa: E501 - sfa_tmem = T.decl_buffer((M, sf_mma_k * num_ki), SF_dtype, scope="tmem", allocated_addr=SFA_TMEM_START, layout=sfa_layout) # noqa: E501 - sfb_tmem = T.decl_buffer((N, sf_mma_k * num_ki), SF_dtype, scope="tmem", allocated_addr=SFB_TMEM_START, layout=sfb_layout) # noqa: E501 + tmem = T.decl_tensor(C_shape, C_dtype, scope="tmem", allocated_addr=tmem_addr[0], layout=TileLayout(S[(128, C_shape[1]) : (1 @ TLane, 1 @ TCol)])) # noqa: E501 + sfa_tmem = T.decl_tensor((M, sf_mma_k * num_ki), SF_dtype, scope="tmem", allocated_addr=SFA_TMEM_START, layout=sfa_layout) # noqa: E501 + sfb_tmem = T.decl_tensor((N, sf_mma_k * num_ki), SF_dtype, scope="tmem", allocated_addr=SFB_TMEM_START, layout=sfb_layout) # noqa: E501 # TMA load A and B from global to shared if tid_in_wg == 0: @@ -2157,9 +2157,9 @@ def test_gemm_tcgen05_arbitrary_tiles(task): # fmt: off @T.prim_func def gemm_async( - A: T.Buffer(A_shape, A_dtype, **A_gmem_kw), - B: T.Buffer(B_shape, B_dtype, **B_gmem_kw), - C: T.Buffer(C_shape, C_dtype), + A: T.Tensor(A_shape, A_dtype, **A_gmem_kw), + B: T.Tensor(B_shape, B_dtype, **B_gmem_kw), + C: T.Tensor(C_shape, C_dtype), ) -> None: T.device_entry() @@ -2168,8 +2168,8 @@ def gemm_async( wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape, A_dtype, scope="shared", layout=A_layout, align=1024) - B_smem = T.alloc_buffer(B_shape, B_dtype, scope="shared", layout=B_layout, align=1024) + A_smem = T.alloc_tensor(A_shape, A_dtype, scope="shared", layout=A_layout, align=1024) + B_smem = T.alloc_tensor(B_shape, B_dtype, scope="shared", layout=B_layout, align=1024) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") @@ -2185,7 +2185,7 @@ def gemm_async( T.address_of(tmem_addr), T.uint32(cols_alloc) ) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( (M, C_shape[1]), C_dtype, scope="tmem", @@ -2289,21 +2289,21 @@ def test_gemm_tcgen05_no_swizzle_smem_descriptor_codegen(a_layout_kind): @T.prim_func def gemm_async_no_swizzle( - A: T.Buffer((M, K), dtype, layout=A_layout), B: T.Buffer((K, B_N), dtype, layout=B_layout) + A: T.Tensor((M, K), dtype, layout=A_layout), B: T.Tensor((K, B_N), dtype, layout=B_layout) ) -> None: T.device_entry() warp_id = T.warp_id([4]) T.thread_id([128]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer((M, K), dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer((K, B_N), dtype, scope="shared", layout=B_layout) + A_smem = T.alloc_tensor((M, K), dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor((K, B_N), dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") if warp_id == 0: T.ptx.tcgen05.alloc.cta_group__2.sync.aligned.shared__cta.b32( T.address_of(tmem_addr[0]), T.uint32(256) ) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( (M, N), "float32", scope="tmem", @@ -2355,21 +2355,21 @@ def gemm_async_replicated_a() -> None: warp_id = T.warp_id([4]) T.thread_id([128]) tid_in_wg = T.thread_id_in_wg([128]) - B_smem = T.alloc_buffer((N_half, K), dtype, scope="shared", layout=B_layout) + B_smem = T.alloc_tensor((N_half, K), dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") if warp_id == 0: T.ptx.tcgen05.alloc.cta_group__2.sync.aligned.shared__cta.b32( T.address_of(tmem_addr[0]), T.uint32(128) ) T.cuda.cta_sync() - C_tmem = T.decl_buffer( + C_tmem = T.decl_tensor( (M, N), "float32", scope="tmem", allocated_addr=tmem_addr[0], layout=C_layout, ) - A_tmem = T.decl_buffer( + A_tmem = T.decl_tensor( (M, K), dtype, scope="tmem", @@ -2417,21 +2417,21 @@ def gemm_async_flat_a() -> None: warp_id = T.warp_id([4]) T.thread_id([128]) tid_in_wg = T.thread_id_in_wg([128]) - B_smem = T.alloc_buffer((N_half, K), dtype, scope="shared", layout=B_layout) + B_smem = T.alloc_tensor((N_half, K), dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") if warp_id == 0: T.ptx.tcgen05.alloc.cta_group__2.sync.aligned.shared__cta.b32( T.address_of(tmem_addr[0]), T.uint32(128) ) T.cuda.cta_sync() - C_tmem = T.decl_buffer( + C_tmem = T.decl_tensor( (M, N), "float32", scope="tmem", allocated_addr=tmem_addr[0], layout=C_layout, ) - A_tmem = T.decl_buffer( + A_tmem = T.decl_tensor( (M, K), dtype, scope="tmem", @@ -2487,17 +2487,17 @@ def test_gemm_tcgen05_no_swizzle_col_major_a_ws_local_idesc(): # fmt: off @T.prim_func def gemm_ws( - A: T.Buffer((M, K), dtype), - B: T.Buffer((K, N), dtype), - C: T.Buffer((128, N // 2), "float32"), + A: T.Tensor((M, K), dtype), + B: T.Tensor((K, N), dtype), + C: T.Tensor((128, N // 2), "float32"), ) -> None: T.device_entry() warp_id = T.warp_id([4]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer((M, K), dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer((K, N), dtype, scope="shared", layout=B_layout) + A_smem = T.alloc_tensor((M, K), dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor((K, N), dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") mma_mbar = T.alloc_shared([1], "uint64") if tid_in_wg == 0: @@ -2510,7 +2510,7 @@ def gemm_ws( T.cuda.cta_sync() # M=64 .ws accumulates via datapath E (FlashMLA head64's tmem_o layout): # lane = m + 64*(n >= N/2), col = n % (N/2). - tmem = T.decl_buffer( + tmem = T.decl_tensor( (M, N), "float32", scope="tmem", @@ -2518,7 +2518,7 @@ def gemm_ws( layout=TileLayout(S[(M, 2, N // 2) : (1 @ TLane, 64 @ TLane, 1 @ TCol)]), ) # Identity overlay of the physical 128x128 TMEM footprint for readback. - tmem_ldst = T.decl_buffer( + tmem_ldst = T.decl_tensor( (128, N // 2), "float32", scope="tmem", @@ -2620,15 +2620,15 @@ def test_gemm_tcgen05_contiguous_kslice_partial_k(k_lo, k_hi): # fmt: off @T.prim_func def gemm_async( - A: T.Buffer(A_shape, dtype), B: T.Buffer(B_shape, dtype), C: T.Buffer(C_shape, "float32") + A: T.Tensor(A_shape, dtype), B: T.Tensor(B_shape, dtype), C: T.Tensor(C_shape, "float32") ) -> None: T.device_entry() warp_id = T.warp_id([4]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape, dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape, dtype, scope="shared", layout=B_layout) + A_smem = T.alloc_tensor(A_shape, dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape, dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") @@ -2642,7 +2642,7 @@ def gemm_async( T.address_of(tmem_addr), T.uint32(128) ) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, N), "float32", scope="tmem", @@ -2722,15 +2722,15 @@ def _run_dense_gemm( @T.prim_func def gemm_async( - A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), C: T.Buffer(C_shape, C_dtype) + A: T.Tensor(A_shape, A_dtype), B: T.Tensor(B_shape, B_dtype), C: T.Tensor(C_shape, C_dtype) ) -> None: T.device_entry() warp_id = T.warp_id([4]) T.cta_id([1]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape, A_dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape, B_dtype, scope="shared", layout=B_layout) + A_smem = T.alloc_tensor(A_shape, A_dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape, B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") @@ -2744,7 +2744,7 @@ def gemm_async( T.address_of(tmem_addr), T.uint32(cols_alloc) ) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, N), C_dtype, scope="tmem", @@ -2845,15 +2845,15 @@ def _run_dense_gemm( @T.prim_func def gemm_async( - A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), C: T.Buffer(C_shape, C_dtype) + A: T.Tensor(A_shape, A_dtype), B: T.Tensor(B_shape, B_dtype), C: T.Tensor(C_shape, C_dtype) ) -> None: T.device_entry() warp_id = T.warp_id([4]) T.cta_id([1]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape, A_dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape, B_dtype, scope="shared", layout=B_layout) + A_smem = T.alloc_tensor(A_shape, A_dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape, B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") @@ -2867,7 +2867,7 @@ def gemm_async( T.address_of(tmem_addr), T.uint32(cols_alloc) ) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, N), C_dtype, scope="tmem", @@ -2940,7 +2940,7 @@ def _build_smem_desc_kernel(smem_desc, weight_stationary=False, pass_descI=False # fmt: off @T.prim_func def gemm_async( - A: T.Buffer(A_shape, A_dtype), B: T.Buffer(B_shape, B_dtype), C: T.Buffer(C_shape, C_dtype) + A: T.Tensor(A_shape, A_dtype), B: T.Tensor(B_shape, B_dtype), C: T.Tensor(C_shape, C_dtype) ) -> None: T.device_entry() @@ -2948,8 +2948,8 @@ def gemm_async( cta_id = T.cta_id([1]) wg_id = T.warpgroup_id([1]) tid_in_wg = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape, A_dtype, scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape, B_dtype, scope="shared", layout=B_layout) + A_smem = T.alloc_tensor(A_shape, A_dtype, scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape, B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") tma_mbar = T.alloc_shared([1], "uint64") mma_mbar = T.alloc_shared([1], "uint64") @@ -2963,7 +2963,7 @@ def gemm_async( T.address_of(tmem_addr), T.uint32(128) ) T.cuda.cta_sync() - tmem = T.decl_buffer( + tmem = T.decl_tensor( (128, C_shape[1]), C_dtype, scope="tmem", @@ -3055,9 +3055,9 @@ def kernel() -> None: cbx, cby = T.cta_id_in_cluster([2, 1]) thread_id = T.thread_id([128]) tid = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape, "float16", scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape, "float16", scope="shared", layout=B_layout) - C_tmem = T.decl_buffer((M_per_cta, N), "float32", scope="tmem", allocated_addr=0, layout=C_layout) # noqa: E501 + A_smem = T.alloc_tensor(A_shape, "float16", scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape, "float16", scope="shared", layout=B_layout) + C_tmem = T.decl_tensor((M_per_cta, N), "float32", scope="tmem", allocated_addr=0, layout=C_layout) # noqa: E501 if tid == 0: Tx.gemm_async(C_tmem[:, :], A_smem[:, :], B_smem[:, :], dispatch="tcgen05", cta_group=2, mma_m=mma_m, mma_n=N) # noqa: E501 # fmt: on @@ -3083,11 +3083,11 @@ def kernel() -> None: cbx, cby = T.cta_id_in_cluster([2, 1]) thread_id = T.thread_id([128]) tid = T.thread_id_in_wg([128]) - A_smem = T.alloc_buffer(A_shape, "float8_e4m3fn", scope="shared", layout=A_layout) - B_smem = T.alloc_buffer(B_shape, "float8_e4m3fn", scope="shared", layout=B_layout) - C_tmem = T.decl_buffer((M, N), "float32", scope="tmem", allocated_addr=0, layout=C_layout) - SFA_tmem = T.decl_buffer((M, 1), "float8_e8m0fnu", scope="tmem", allocated_addr=N, layout=sf_layout) # noqa: E501 - SFB_tmem = T.decl_buffer((N, 1), "float8_e8m0fnu", scope="tmem", allocated_addr=N + 4, layout=sf_layout) # noqa: E501 + A_smem = T.alloc_tensor(A_shape, "float8_e4m3fn", scope="shared", layout=A_layout) + B_smem = T.alloc_tensor(B_shape, "float8_e4m3fn", scope="shared", layout=B_layout) + C_tmem = T.decl_tensor((M, N), "float32", scope="tmem", allocated_addr=0, layout=C_layout) + SFA_tmem = T.decl_tensor((M, 1), "float8_e8m0fnu", scope="tmem", allocated_addr=N, layout=sf_layout) # noqa: E501 + SFB_tmem = T.decl_tensor((N, 1), "float8_e8m0fnu", scope="tmem", allocated_addr=N + 4, layout=sf_layout) # noqa: E501 if tid == 0: Tx.gemm_async(C_tmem[:, :], A_smem[:, :], B_smem[:, :], SFA=SFA_tmem[:, :], SFB=SFB_tmem[:, :], dispatch="tcgen05", cta_group=2, mma_m=256, mma_n=64) # noqa: E501 # fmt: on @@ -3289,12 +3289,12 @@ def _build_cta1_m64_packed_c_kernel(weight_stationary=None, mma_config=None): C_layout = TileLayout(S[(M, 2, N // 2) : (1 @ TLane, 64 @ TLane, 1 @ TCol)]) @T.prim_func - def gemm_packed_c(B: T.Buffer((N, K), B_dtype)) -> None: + def gemm_packed_c(B: T.Tensor((N, K), B_dtype)) -> None: T.device_entry() warp_id = T.warp_id([4]) T.thread_id([128]) tid = T.thread_id_in_wg([128]) - B_smem = T.alloc_buffer((N, K), B_dtype, scope="shared", layout=B_layout) + B_smem = T.alloc_tensor((N, K), B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") if warp_id == 0: T.ptx.tcgen05.alloc.cta_group__1.sync.aligned.shared__cta.b32( @@ -3303,14 +3303,14 @@ def gemm_packed_c(B: T.Buffer((N, K), B_dtype)) -> None: T.cuda.cta_sync() # A-in-TMEM .ws reads A from both 64-lane halves, so A is declared in # the honest batched A[2, M, K] fold (the A-side). - A_tmem = T.decl_buffer( + A_tmem = T.decl_tensor( (2, M, K), A_dtype, scope="tmem", allocated_addr=256, layout=TileLayout(S[(2, M, K) : (64 @ TLane, 1 @ TLane, 1 @ TCol)]), ) - C_tmem = T.decl_buffer( + C_tmem = T.decl_tensor( (M, N), C_dtype, scope="tmem", @@ -3377,12 +3377,12 @@ def _build_cta1_m64_batched_c_kernel(): C_layout = TileLayout(S[(2, M, N // 2) : (64 @ TLane, 1 @ TLane, 1 @ TCol)]) @T.prim_func - def gemm_batched_c(B: T.Buffer((N, K), B_dtype)) -> None: + def gemm_batched_c(B: T.Tensor((N, K), B_dtype)) -> None: T.device_entry() warp_id = T.warp_id([4]) T.thread_id([128]) tid = T.thread_id_in_wg([128]) - B_smem = T.alloc_buffer((N, K), B_dtype, scope="shared", layout=B_layout) + B_smem = T.alloc_tensor((N, K), B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") if warp_id == 0: T.ptx.tcgen05.alloc.cta_group__1.sync.aligned.shared__cta.b32( @@ -3391,14 +3391,14 @@ def gemm_batched_c(B: T.Buffer((N, K), B_dtype)) -> None: T.cuda.cta_sync() # Honest batched A[2, M, K] fold (A-side), matching the # batched C below: both banks of the M=64 .ws are explicit. - A_tmem = T.decl_buffer( + A_tmem = T.decl_tensor( (2, M, K), B_dtype, scope="tmem", allocated_addr=256, layout=TileLayout(S[(2, M, K) : (64 @ TLane, 1 @ TLane, 1 @ TCol)]), ) - C_tmem = T.decl_buffer( + C_tmem = T.decl_tensor( (2, M, N // 2), "float32", scope="tmem", @@ -3430,26 +3430,26 @@ def _build_cta1_m64_identity_c_ws_kernel(): B_layout = mma_shared_layout(B_dtype, SwizzleMode.SWIZZLE_32B_ATOM, (N, K)) @T.prim_func - def gemm_identity_c(B: T.Buffer((N, K), B_dtype)) -> None: + def gemm_identity_c(B: T.Tensor((N, K), B_dtype)) -> None: T.device_entry() warp_id = T.warp_id([4]) T.thread_id([128]) tid = T.thread_id_in_wg([128]) - B_smem = T.alloc_buffer((N, K), B_dtype, scope="shared", layout=B_layout) + B_smem = T.alloc_tensor((N, K), B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") if warp_id == 0: T.ptx.tcgen05.alloc.cta_group__1.sync.aligned.shared__cta.b32( T.address_of(tmem_addr[0]), T.uint32(512) ) T.cuda.cta_sync() - A_tmem = T.decl_buffer( + A_tmem = T.decl_tensor( (M, K), B_dtype, scope="tmem", allocated_addr=256, layout=TileLayout(S[(M, K) : (1 @ TLane, 1 @ TCol)]), ) - C_tmem = T.decl_buffer( + C_tmem = T.decl_tensor( (M, N), "float32", scope="tmem", @@ -3497,26 +3497,26 @@ def _build_cta1_m64_flat_a_ws_kernel(): C_layout = TileLayout(S[(M, 2, N // 2) : (1 @ TLane, 64 @ TLane, 1 @ TCol)]) @T.prim_func - def gemm_flat_a(B: T.Buffer((N, K), B_dtype)) -> None: + def gemm_flat_a(B: T.Tensor((N, K), B_dtype)) -> None: T.device_entry() warp_id = T.warp_id([4]) T.thread_id([128]) tid = T.thread_id_in_wg([128]) - B_smem = T.alloc_buffer((N, K), B_dtype, scope="shared", layout=B_layout) + B_smem = T.alloc_tensor((N, K), B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") if warp_id == 0: T.ptx.tcgen05.alloc.cta_group__1.sync.aligned.shared__cta.b32( T.address_of(tmem_addr[0]), T.uint32(512) ) T.cuda.cta_sync() - A_tmem = T.decl_buffer( + A_tmem = T.decl_tensor( (M, K), B_dtype, scope="tmem", allocated_addr=256, layout=TileLayout(S[(M, K) : (1 @ TLane, 1 @ TCol)]), ) - C_tmem = T.decl_buffer( + C_tmem = T.decl_tensor( (M, N), "float32", scope="tmem", @@ -3564,26 +3564,26 @@ def _build_m128_batched_a_kernel(): B_layout = mma_shared_layout(B_dtype, SwizzleMode.SWIZZLE_32B_ATOM, (N, K)) @T.prim_func - def gemm_m128_batched_a(B: T.Buffer((N, K), B_dtype)) -> None: + def gemm_m128_batched_a(B: T.Tensor((N, K), B_dtype)) -> None: T.device_entry() warp_id = T.warp_id([4]) T.thread_id([128]) tid = T.thread_id_in_wg([128]) - B_smem = T.alloc_buffer((N, K), B_dtype, scope="shared", layout=B_layout) + B_smem = T.alloc_tensor((N, K), B_dtype, scope="shared", layout=B_layout) tmem_addr = T.alloc_shared([1], "uint32") if warp_id == 0: T.ptx.tcgen05.alloc.cta_group__1.sync.aligned.shared__cta.b32( T.address_of(tmem_addr[0]), T.uint32(512) ) T.cuda.cta_sync() - A_tmem = T.decl_buffer( + A_tmem = T.decl_tensor( (2, M, K), B_dtype, scope="tmem", allocated_addr=256, layout=TileLayout(S[(2, M, K) : (64 @ TLane, 1 @ TLane, 1 @ TCol)]), ) - C_tmem = T.decl_buffer( + C_tmem = T.decl_tensor( (M, N), "float32", scope="tmem", @@ -3705,13 +3705,13 @@ def full_region(buf): A_shape = (M, K) if not transA else (K, M) B_shape = (K, N) if transB else (N, K) - A_buf = tvm.tirx.decl_buffer(A_shape, dtype, "A", scope=A_scope, layout=A_layout) + A_buf = tvm.tirx.decl_tensor(A_shape, dtype, "A", scope=A_scope, layout=A_layout) if A_scope == "tmem": A_buf = A_buf.with_allocated_addr([tvm.tirx.IntImm("uint32", A_allocated_addr)]) - B_buf = tvm.tirx.decl_buffer(B_shape, dtype, "B_smem", scope="shared.dyn", layout=B_layout) + B_buf = tvm.tirx.decl_tensor(B_shape, dtype, "B_smem", scope="shared.dyn", layout=B_layout) if C_layout is None: C_layout = TileLayout(S[(M, N) : (1 @ TLane, 1 @ TCol)]) - C_buf = tvm.tirx.decl_buffer((M, N), "float32", "C_tmem", scope="tmem", layout=C_layout) + C_buf = tvm.tirx.decl_tensor((M, N), "float32", "C_tmem", scope="tmem", layout=C_layout) C_buf = C_buf.with_allocated_addr([tvm.tirx.IntImm("uint32", C_allocated_addr)]) call = GemmAsync( full_region(C_buf), @@ -3785,21 +3785,21 @@ def full_region(buf): M, N, K = 128, 64, 64 data_dtype = "float4_e2m1fn" sf_dtype = "float8_e4m3fn" - A = tvm.tirx.decl_buffer( + A = tvm.tirx.decl_tensor( (M, K), data_dtype, "A_smem", scope="shared.dyn", layout=mma_shared_layout(data_dtype, SwizzleMode.SWIZZLE_32B_ATOM, (M, K)), ) - B = tvm.tirx.decl_buffer( + B = tvm.tirx.decl_tensor( (N, K), data_dtype, "B_smem", scope="shared.dyn", layout=mma_shared_layout(data_dtype, SwizzleMode.SWIZZLE_32B_ATOM, (N, K)), ) - C = tvm.tirx.decl_buffer( + C = tvm.tirx.decl_tensor( (M, N), "float32", "C_tmem", @@ -3812,10 +3812,10 @@ def full_region(buf): sfb_base = sf_tmem_layout(N, SF_K=sf_per_mma, sf_per_mma=sf_per_mma) sfa_layout = TileLayout.from_iters(sfa_base.shard, sfa_base.replica, {TLane: 1}) sfb_layout = TileLayout.from_iters(sfb_base.shard, sfb_base.replica, {TLane: 2}) - SFA = tvm.tirx.decl_buffer( + SFA = tvm.tirx.decl_tensor( (M, sf_per_mma), sf_dtype, "SFA_tmem", scope="tmem", layout=sfa_layout ).with_allocated_addr([tvm.tirx.IntImm("uint32", 256)]) - SFB = tvm.tirx.decl_buffer( + SFB = tvm.tirx.decl_tensor( (N, sf_per_mma), sf_dtype, "SFB_tmem", scope="tmem", layout=sfb_layout ).with_allocated_addr([tvm.tirx.IntImm("uint32", 320)]) diff --git a/tests/python/tirx/operator/tile_primitive/cuda/permute_layout/test_permute_layout.py b/tests/python/tirx/operator/tile_primitive/cuda/permute_layout/test_permute_layout.py index e9bbeb5ffcf2..5d04d9898e3b 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/permute_layout/test_permute_layout.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/permute_layout/test_permute_layout.py @@ -201,7 +201,7 @@ def test_sf_blockwise_transpose(name, pipe, blk, dtype): # fmt: off @T.prim_func - def f(A_buf: T.Buffer(shape, dtype, layout=pre), B_buf: T.Buffer(shape, dtype, layout=post)): + def f(A_buf: T.Tensor(shape, dtype, layout=pre), B_buf: T.Tensor(shape, dtype, layout=post)): T.device_entry() T.cta_id([1]) @@ -249,8 +249,8 @@ def test_identity_passes_through_as_copy(): # fmt: off @T.prim_func def f( - A_buf: T.Buffer(shape, "uint32", layout=layout), - B_buf: T.Buffer(shape, "uint32", layout=layout), + A_buf: T.Tensor(shape, "uint32", layout=layout), + B_buf: T.Tensor(shape, "uint32", layout=layout), ): T.device_entry() @@ -287,7 +287,7 @@ def test_generic_transpose(shape, src_strides, dst_strides, dtype): # fmt: off @T.prim_func - def f(A_buf: T.Buffer(shape, dtype, layout=pre), B_buf: T.Buffer(shape, dtype, layout=post)): + def f(A_buf: T.Tensor(shape, dtype, layout=pre), B_buf: T.Tensor(shape, dtype, layout=post)): T.device_entry() T.cta_id([1]) @@ -314,8 +314,8 @@ def _build_and_assert_rejected(shape, src_layout, dst_layout, dtype, msg_substr) # fmt: off @T.prim_func def f( - A_buf: T.Buffer(shape, dtype, layout=src_layout), - B_buf: T.Buffer(shape, dtype, layout=dst_layout), + A_buf: T.Tensor(shape, dtype, layout=src_layout), + B_buf: T.Tensor(shape, dtype, layout=dst_layout), ): T.device_entry() @@ -340,8 +340,8 @@ def test_reject_dtype_mismatch(): # fmt: off @T.prim_func def f( - A_buf: T.Buffer(shape, "uint32", layout=layout), - B_buf: T.Buffer(shape, "uint16", layout=layout), + A_buf: T.Tensor(shape, "uint32", layout=layout), + B_buf: T.Tensor(shape, "uint16", layout=layout), ): T.device_entry() @@ -363,8 +363,8 @@ def test_reject_shape_mismatch(): # fmt: off @T.prim_func def f( - A_buf: T.Buffer((4, 32), "uint32", layout=src_layout), - B_buf: T.Buffer((8, 16), "uint32", layout=dst_layout), + A_buf: T.Tensor((4, 32), "uint32", layout=src_layout), + B_buf: T.Tensor((8, 16), "uint32", layout=dst_layout), ): T.device_entry() @@ -397,8 +397,8 @@ def test_reject_swizzle_layout(): # fmt: off @T.prim_func def f( - A_buf: T.Buffer((4, 32), "uint32", layout=swizzled), - B_buf: T.Buffer((4, 32), "uint32", layout=plain), + A_buf: T.Tensor((4, 32), "uint32", layout=swizzled), + B_buf: T.Tensor((4, 32), "uint32", layout=plain), ): T.device_entry() @@ -420,8 +420,8 @@ def test_reject_non_warp_scope(): # fmt: off @T.prim_func def f( - A_buf: T.Buffer((4, 32), "uint32", layout=layout_pre), - B_buf: T.Buffer((4, 32), "uint32", layout=layout_post), + A_buf: T.Tensor((4, 32), "uint32", layout=layout_pre), + B_buf: T.Tensor((4, 32), "uint32", layout=layout_post), ): T.device_entry() @@ -456,13 +456,13 @@ def test_shared_to_shared_uses_direct_ldst(dtype): # fmt: off @T.prim_func - def f(A_buf: T.Buffer(shape, dtype, layout=pre), B_buf: T.Buffer(shape, dtype, layout=post)): + def f(A_buf: T.Tensor(shape, dtype, layout=pre), B_buf: T.Tensor(shape, dtype, layout=post)): T.device_entry() T.cta_id([1]) tid = T.thread_id([32]) - sA = T.alloc_buffer(shape, dtype, scope="shared", layout=pre) - sB = T.alloc_buffer(shape, dtype, scope="shared", layout=post) + sA = T.alloc_tensor(shape, dtype, scope="shared", layout=pre) + sB = T.alloc_tensor(shape, dtype, scope="shared", layout=post) Tx.cta.copy(sA[:, :], A_buf[:, :]) T.cuda.cta_sync() Tx.warp.permute_layout(sB[:, :], sA[:, :]) diff --git a/tests/python/tirx/operator/tile_primitive/cuda/reduction/test_reduction.py b/tests/python/tirx/operator/tile_primitive/cuda/reduction/test_reduction.py index 134d79d4b657..83e496bd51cf 100644 --- a/tests/python/tirx/operator/tile_primitive/cuda/reduction/test_reduction.py +++ b/tests/python/tirx/operator/tile_primitive/cuda/reduction/test_reduction.py @@ -70,15 +70,15 @@ def test_reduction_shared( # fmt: off @T.prim_func def test_reduction( - A: T.Buffer(src_shape, dtype, layout=g_layout_src), - B: T.Buffer(dst_shape, dtype, layout=g_layout_dst), + A: T.Tensor(src_shape, dtype, layout=g_layout_src), + B: T.Tensor(dst_shape, dtype, layout=g_layout_dst), ) -> None: T.device_entry() _bx = T.cta_id([1]) _tid = T.thread_id([thread_cnt]) - A_smem = T.alloc_buffer(s_shape_src, dtype, scope="shared", layout=s_layout_src) - B_smem = T.alloc_buffer(s_shape_dst, dtype, scope="shared", layout=s_layout_dst) + A_smem = T.alloc_tensor(s_shape_src, dtype, scope="shared", layout=s_layout_src) + B_smem = T.alloc_tensor(s_shape_dst, dtype, scope="shared", layout=s_layout_dst) Tx.cta.copy(A_smem[tuple(copy_slice_src)], A[tuple(copy_slice_src)]) if accum: @@ -170,16 +170,16 @@ def test_reduction_shared_subscope(exec_scope, op_type, accum): if exec_scope == "warp": @T.prim_func def test_func( - A: T.Buffer(src_shape, dtype, layout=g_layout_src), - B: T.Buffer(dst_shape, dtype, layout=g_layout_dst), + A: T.Tensor(src_shape, dtype, layout=g_layout_src), + B: T.Tensor(dst_shape, dtype, layout=g_layout_dst), ) -> None: T.device_entry() warp_id = T.warp_id([(256) // 32]) _bx = T.cta_id([1]) _tid = T.thread_id([256]) - A_smem = T.alloc_buffer(list(src_shape), dtype, scope="shared", layout=s_layout_src) - B_smem = T.alloc_buffer(list(dst_shape), dtype, scope="shared", layout=s_layout_dst) + A_smem = T.alloc_tensor(list(src_shape), dtype, scope="shared", layout=s_layout_src) + B_smem = T.alloc_tensor(list(dst_shape), dtype, scope="shared", layout=s_layout_dst) Tx.cta.copy(A_smem, A) if accum: Tx.cta.copy(B_smem, B) @@ -197,16 +197,16 @@ def test_func( elif exec_scope == "warpgroup": @T.prim_func def test_func( - A: T.Buffer(src_shape, dtype, layout=g_layout_src), - B: T.Buffer(dst_shape, dtype, layout=g_layout_dst), + A: T.Tensor(src_shape, dtype, layout=g_layout_src), + B: T.Tensor(dst_shape, dtype, layout=g_layout_dst), ) -> None: T.device_entry() wg_id = T.warpgroup_id([(256) // 128]) _bx = T.cta_id([1]) _tid = T.thread_id([256]) - A_smem = T.alloc_buffer(list(src_shape), dtype, scope="shared", layout=s_layout_src) - B_smem = T.alloc_buffer(list(dst_shape), dtype, scope="shared", layout=s_layout_dst) + A_smem = T.alloc_tensor(list(src_shape), dtype, scope="shared", layout=s_layout_src) + B_smem = T.alloc_tensor(list(dst_shape), dtype, scope="shared", layout=s_layout_dst) Tx.cta.copy(A_smem, A) if accum: Tx.cta.copy(B_smem, B) @@ -224,15 +224,15 @@ def test_func( elif exec_scope == "thread": @T.prim_func def test_func( - A: T.Buffer(src_shape, dtype, layout=g_layout_src), - B: T.Buffer(dst_shape, dtype, layout=g_layout_dst), + A: T.Tensor(src_shape, dtype, layout=g_layout_src), + B: T.Tensor(dst_shape, dtype, layout=g_layout_dst), ) -> None: T.device_entry() _bx = T.cta_id([1]) _tid = T.thread_id([256]) - A_smem = T.alloc_buffer(list(src_shape), dtype, scope="shared", layout=s_layout_src) - B_smem = T.alloc_buffer(list(dst_shape), dtype, scope="shared", layout=s_layout_dst) + A_smem = T.alloc_tensor(list(src_shape), dtype, scope="shared", layout=s_layout_src) + B_smem = T.alloc_tensor(list(dst_shape), dtype, scope="shared", layout=s_layout_dst) Tx.cta.copy(A_smem, A) if accum: Tx.cta.copy(B_smem, B) @@ -325,15 +325,15 @@ def decompose_flat(flat_idx, shape): # fmt: off @T.prim_func def test_func( - A: T.Buffer(list(src_shape), dtype, layout=TileLayout(S[src_shape])), - B: T.Buffer(list(dst_shape), dtype, layout=TileLayout(S[dst_shape])), + A: T.Tensor(list(src_shape), dtype, layout=TileLayout(S[src_shape])), + B: T.Tensor(list(dst_shape), dtype, layout=TileLayout(S[dst_shape])), ) -> None: T.device_entry() _bx = T.cta_id([1]) _tid = T.thread_id([1]) - A_local = T.alloc_buffer(list(src_shape), dtype, scope="local") - B_local = T.alloc_buffer(list(dst_shape), dtype, scope="local") + A_local = T.alloc_tensor(list(src_shape), dtype, scope="local") + B_local = T.alloc_tensor(list(dst_shape), dtype, scope="local") for i in T.serial(src_total): idx = T.meta_var(decompose_flat(i, src_shape)) @@ -454,8 +454,8 @@ def decompose_flat(flat_idx, shape): # fmt: off @T.prim_func def test_func( - A: T.Buffer(list(src_shape), dtype, layout=g_layout_a), - B: T.Buffer(list(dst_shape), dtype, layout=g_layout_b), + A: T.Tensor(list(src_shape), dtype, layout=g_layout_a), + B: T.Tensor(list(dst_shape), dtype, layout=g_layout_b), ) -> None: T.device_entry() @@ -463,8 +463,8 @@ def test_func( _warp_id = T.warp_id([1]) lane_id = T.lane_id([thread_cnt]) - acc = T.alloc_buffer(list((1, *inner_dims)), dtype=dtype, scope="local", layout=g_layout_a) - red = T.alloc_buffer(list((1, *dst_dims)), dtype=dtype, scope="local", layout=g_layout_b) + acc = T.alloc_tensor(list((1, *inner_dims)), dtype=dtype, scope="local", layout=g_layout_a) + red = T.alloc_tensor(list((1, *dst_dims)), dtype=dtype, scope="local", layout=g_layout_b) for i in T.serial(src_local_total): idx = T.meta_var(decompose_flat(i, inner_dims)) acc[(0, *list(idx))] = A[(lane_id, *list(idx))] @@ -557,8 +557,8 @@ def test_reduction_local_view_complex(n_groups, n_warps, op_type, dtype, shuffle # fmt: off @T.prim_func def test_func( - A: T.Buffer(g_shape_a, dtype, layout=g_layout_a), - B: T.Buffer(g_shape_b, dtype, layout=g_layout_b), + A: T.Tensor(g_shape_a, dtype, layout=g_layout_a), + B: T.Tensor(g_shape_b, dtype, layout=g_layout_b), ) -> None: T.device_entry() @@ -572,7 +572,7 @@ def test_func( warp_atom = atom.tile(warp_layout, (8, 4), (1, 2)) tile = T.TileLayout(T.S[(2, NUM_COL // 8) : (1, 2)]) acc_layout = warp_atom.tile(tile, (2, NUM_COL // 8), (8, 8)) - acc = T.alloc_buffer( + acc = T.alloc_tensor( [2, NUM_COL // 4], dtype=dtype, scope="local", @@ -584,7 +584,7 @@ def test_func( red_warp_atom = red_atom.tile(warp_layout, (8, 4), (1, 1)) red_tile = T.TileLayout(T.S[(2, 1) : (1, 1)]) red_layout = red_warp_atom.tile(red_tile, (2, 1), (8, 4)) - red = T.alloc_buffer( + red = T.alloc_tensor( [2], dtype=dtype, scope="local", @@ -683,15 +683,15 @@ def test_reduction_local_optimized_3input_maxmin(reduction_len, op_type, accum): # fmt: off @T.prim_func def test_func( - A: T.Buffer([reduction_len], dtype, layout=TileLayout(S[reduction_len])), - B: T.Buffer([1], dtype, layout=TileLayout(S[1])), + A: T.Tensor([reduction_len], dtype, layout=TileLayout(S[reduction_len])), + B: T.Tensor([1], dtype, layout=TileLayout(S[1])), ) -> None: T.device_entry() _bx = T.cta_id([1]) _tid = T.thread_id([1]) - A_local = T.alloc_buffer([reduction_len], dtype, scope="local") - B_local = T.alloc_buffer([1], dtype, scope="local") + A_local = T.alloc_tensor([reduction_len], dtype, scope="local") + B_local = T.alloc_tensor([1], dtype, scope="local") # Load from global to local for i in T.serial(reduction_len): @@ -756,15 +756,15 @@ def test_reduction_local_optimized_packed_add_sum(reduction_len, accum): # fmt: off @T.prim_func def test_func( - A: T.Buffer([reduction_len], dtype, layout=TileLayout(S[reduction_len])), - B: T.Buffer([1], dtype, layout=TileLayout(S[1])), + A: T.Tensor([reduction_len], dtype, layout=TileLayout(S[reduction_len])), + B: T.Tensor([1], dtype, layout=TileLayout(S[1])), ) -> None: T.device_entry() _bx = T.cta_id([1]) _tid = T.thread_id([1]) - A_local = T.alloc_buffer([reduction_len], dtype, scope="local") - B_local = T.alloc_buffer([1], dtype, scope="local") + A_local = T.alloc_tensor([reduction_len], dtype, scope="local") + B_local = T.alloc_tensor([1], dtype, scope="local") # Load from global to local for i in T.serial(reduction_len): @@ -833,15 +833,15 @@ def test_reduction_op_warp_shuffle(op_type, dtype): # fmt: off @T.prim_func def test_func( - A: T.Buffer(g_shape, dtype, layout=g_layout), B: T.Buffer(g_shape, dtype, layout=g_layout) + A: T.Tensor(g_shape, dtype, layout=g_layout), B: T.Tensor(g_shape, dtype, layout=g_layout) ) -> None: T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([1]) lane_id = T.lane_id([32]) - src_local = T.alloc_buffer([1], dtype, scope="local") - dst_local = T.alloc_buffer([1], dtype, scope="local") + src_local = T.alloc_tensor([1], dtype, scope="local") + dst_local = T.alloc_tensor([1], dtype, scope="local") src_local[0] = A[lane_id] src_view = src_local.view(N, layout=src_layout) dst_view = dst_local.view(1, layout=dst_layout) @@ -903,16 +903,16 @@ def test_reduction_op_warp_shuffle_multi_elem(op_type, dtype): dst_lay = TileLayout(S[ELEMS_PER_THREAD]) @T.prim_func def test_func( - A: T.Buffer(g_shape, dtype, layout=g_layout), - B: T.Buffer([ELEMS_PER_THREAD], dtype, layout=dst_lay), + A: T.Tensor(g_shape, dtype, layout=g_layout), + B: T.Tensor([ELEMS_PER_THREAD], dtype, layout=dst_lay), ) -> None: T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([1]) lane_id = T.lane_id([32]) - src_local = T.alloc_buffer([ELEMS_PER_THREAD], dtype, scope="local") - dst_local = T.alloc_buffer([ELEMS_PER_THREAD], dtype, scope="local") + src_local = T.alloc_tensor([ELEMS_PER_THREAD], dtype, scope="local") + dst_local = T.alloc_tensor([ELEMS_PER_THREAD], dtype, scope="local") for i in T.serial(ELEMS_PER_THREAD): src_local[i] = A[lane_id * ELEMS_PER_THREAD + i] src_view = src_local.view(TOTAL, layout=src_layout) @@ -972,16 +972,16 @@ def test_reduction_op_warp_shuffle_gapped_permuted_storage(): # fmt: off @T.prim_func def test_func( - A: T.Buffer(src_shape, "float32", layout=TileLayout(S[src_shape])), - B: T.Buffer(local_shape, "float32", layout=TileLayout(S[local_shape])), + A: T.Tensor(src_shape, "float32", layout=TileLayout(S[src_shape])), + B: T.Tensor(local_shape, "float32", layout=TileLayout(S[local_shape])), ) -> None: T.device_entry() _cta_id = T.cta_id([1]) _warp_id = T.warp_id([1]) lane_id = T.lane_id([n_lanes]) - src_local = T.alloc_buffer([7], "float32", scope="local") - dst_local = T.alloc_buffer([7], "float32", scope="local") + src_local = T.alloc_tensor([7], "float32", scope="local") + dst_local = T.alloc_tensor([7], "float32", scope="local") src_view = src_local.view(*src_shape, layout=src_layout) dst_view = dst_local.view(*local_shape, layout=dst_layout) for i, j in T.grid(*local_shape): @@ -1029,8 +1029,8 @@ def test_reduction_warp_shuffle_multi_warp_loop(): # fmt: off @T.prim_func def test_func( - A: T.Buffer([N_ITER, N], "float32", scope="global"), - B: T.Buffer([N_ITER], "float32", scope="global"), + A: T.Tensor([N_ITER, N], "float32", scope="global"), + B: T.Tensor([N_ITER], "float32", scope="global"), ) -> None: T.device_entry() @@ -1041,10 +1041,10 @@ def test_func( pool = T.SMEMPool() sum_smem = pool.alloc([BDY], "float32") pool.commit() - partial_buf = T.alloc_buffer([1], "float32", scope="local") - result_buf = T.alloc_buffer([1], "float32", scope="local") - cross_buf = T.alloc_buffer([1], "float32", scope="local") - cross_res = T.alloc_buffer([1], "float32", scope="local") + partial_buf = T.alloc_tensor([1], "float32", scope="local") + result_buf = T.alloc_tensor([1], "float32", scope="local") + cross_buf = T.alloc_tensor([1], "float32", scope="local") + cross_res = T.alloc_tensor([1], "float32", scope="local") for it in T.serial(N_ITER): partial_buf[0] = A[it, thread_id] @@ -1102,16 +1102,16 @@ def test_reduction_warpgroup_wg_local_layout(op_name): @T.prim_func def test_func( - A: T.Buffer((rows, cols), dtype, layout=TileLayout(S[rows, cols])), - B: T.Buffer((rows, 1), dtype, layout=TileLayout(S[rows, 1])), + A: T.Tensor((rows, cols), dtype, layout=TileLayout(S[rows, cols])), + B: T.Tensor((rows, 1), dtype, layout=TileLayout(S[rows, 1])), ) -> None: T.device_entry() _bx = T.cta_id([1]) wg_id = T.warpgroup_id([1]) tid = T.thread_id_in_wg([rows]) - src = T.alloc_buffer((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) - dst = T.alloc_buffer((rows, 1), dtype, scope="local", layout=wg_local_layout(1)) + src = T.alloc_tensor((rows, cols), dtype, scope="local", layout=wg_local_layout(cols)) + dst = T.alloc_tensor((rows, 1), dtype, scope="local", layout=wg_local_layout(1)) src_local = src.local(cols) for i in T.serial(cols): src_local[i] = A[tid, i] diff --git a/tests/python/tirx/operator/tile_primitive/test_dispatcher.py b/tests/python/tirx/operator/tile_primitive/test_dispatcher.py index ac21a2d7f3d9..45e940dda63e 100644 --- a/tests/python/tirx/operator/tile_primitive/test_dispatcher.py +++ b/tests/python/tirx/operator/tile_primitive/test_dispatcher.py @@ -140,13 +140,13 @@ def test_dispatch_prints_real_opcall_ir(): """Create a real TilePrimitiveCall via BufferRegions and ensure its IR is in the table.""" _import_and_register() from tvm.ir import Op - from tvm.tirx.buffer import decl_buffer + from tvm.tirx.buffer import decl_tensor from tvm.tirx.operator.tile_primitive.dispatcher import run_dispatch from tvm.tirx.tile_primitive import TilePrimitiveCall # Build a real TIRx TilePrimitiveCall: tirx.tile.copy(A[0:64], B[0:64]) - A = decl_buffer((64,), "float32", scope="global") - B = decl_buffer((64,), "float32", scope="shared") + A = decl_tensor((64,), "float32", scope="global") + B = decl_tensor((64,), "float32", scope="shared") real_opcall = TilePrimitiveCall( A[0:64], B[0:64], op=Op.get("tirx.tile.copy"), workspace={}, config={} ) diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py b/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py index 4ddcd5f6bc7b..9759510130cf 100644 --- a/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py +++ b/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py @@ -76,9 +76,9 @@ def test_simple_binary(op_type, operands_type): @T.prim_func def binary() ->None: T.device_entry() - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) if T.constexpr( operands_type == "region_region" or operands_type.startswith("region_broadcast") ): @@ -91,9 +91,9 @@ def binary() ->None: @T.prim_func def expected(): T.func_attr({"global_symbol": "binary"}) - A_sbuf = T.alloc_buffer(src1_shape, scope="trn.sbuf") - B_sbuf = T.alloc_buffer(src2_shape, scope="trn.sbuf") - C_sbuf = T.alloc_buffer(dst_shape, scope="trn.sbuf") + A_sbuf = T.alloc_tensor(src1_shape, scope="trn.sbuf") + B_sbuf = T.alloc_tensor(src2_shape, scope="trn.sbuf") + C_sbuf = T.alloc_tensor(dst_shape, scope="trn.sbuf") for b_loop in T.serial(0, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -151,9 +151,9 @@ def test_binary_complex(op_type, operands_type): @T.prim_func def binary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) A_sbuf_view = A_sbuf.view(*src1_view_shape) B_sbuf_view = B_sbuf.view(*src2_view_shape) C_sbuf_view = C_sbuf.view(*dst_view_shape) @@ -175,12 +175,12 @@ def binary() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "binary"}) - A_sbuf = T.alloc_buffer(src1_layout_data_iter, scope="trn.sbuf") - B_sbuf = T.alloc_buffer(src2_layout_data_iter, scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - A_sbuf_view = T.decl_buffer(src1_layout_data_iter, data=A_sbuf.data, scope="trn.sbuf", layout=None) # noqa: E501 - B_sbuf_view = T.decl_buffer(src2_layout_data_iter, data=B_sbuf.data, scope="trn.sbuf", layout=None) # noqa: E501 - C_sbuf_view = T.decl_buffer((128, 2048), data=C_sbuf.data, scope="trn.sbuf", layout=None) + A_sbuf = T.alloc_tensor(src1_layout_data_iter, scope="trn.sbuf") + B_sbuf = T.alloc_tensor(src2_layout_data_iter, scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + A_sbuf_view = T.decl_tensor(src1_layout_data_iter, data=A_sbuf.data, scope="trn.sbuf", layout=None) # noqa: E501 + B_sbuf_view = T.decl_tensor(src2_layout_data_iter, data=B_sbuf.data, scope="trn.sbuf", layout=None) # noqa: E501 + C_sbuf_view = T.decl_tensor((128, 2048), data=C_sbuf.data, scope="trn.sbuf", layout=None) for i, b_loop in T.grid(4, b_extent): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -216,17 +216,17 @@ def test_binary_broadcast1(): @T.prim_func def binary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.add(C_sbuf, A_sbuf, B_sbuf) @T.prim_func def expected(): T.func_attr({"global_symbol": "binary"}) - A_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") for b_loop in T.serial(0, 512): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -252,17 +252,17 @@ def test_binary_broadcast2(): @T.prim_func def binary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.add(C_sbuf, A_sbuf, B_sbuf) @T.prim_func def expected(): T.func_attr({"global_symbol": "binary"}) - A_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") for b_loop in T.serial(0, 128): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -288,17 +288,17 @@ def test_binary_broadcast3(): @T.prim_func def binary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.add(C_sbuf, A_sbuf, B_sbuf[0]) @T.prim_func def expected(): T.func_attr({"global_symbol": "binary"}) - A_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") for b_loop in T.serial(0, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -324,18 +324,18 @@ def test_binary_with_guard(): @T.prim_func def binary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for j in range(4): Tx.add(C_sbuf[:, :, 0:j*128], A_sbuf[:, :, 0:j*128], B_sbuf[:, 0:j*128]) @T.prim_func def expected(): T.func_attr({"global_symbol": "binary"}) - A_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") for j, b_loop in T.grid(4, 96): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_compose_op_trn.py b/tests/python/tirx/operator/tile_primitive/trn/test_compose_op_trn.py index 4e52cf5f8c98..b01f6a4931ad 100644 --- a/tests/python/tirx/operator/tile_primitive/trn/test_compose_op_trn.py +++ b/tests/python/tirx/operator/tile_primitive/trn/test_compose_op_trn.py @@ -60,23 +60,23 @@ def test_simple_activation_reduce(): @T.prim_func def activation_reduce(): T.device_entry() - A = T.alloc_buffer(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) - B = T.alloc_buffer(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) - C = T.alloc_buffer(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) + A = T.alloc_tensor(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) + B = T.alloc_tensor(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) + C = T.alloc_tensor(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) Tx.unary_reduce(B, C, A, "sqrt", "sum", reduce_axes=1) @T.prim_func def expected(): T.func_attr({"global_symbol": "activation_reduce"}) - const_bias = T.alloc_buffer((128, 512), scope="trn.sbuf") + const_bias = T.alloc_tensor((128, 512), scope="trn.sbuf") with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for f_loop in T.serial(512, annotations={"nki_dim": "F"}): T.nki.memset(const_bias[p_loop, f_loop], T.float32(0.0)) - A = T.alloc_buffer((128, 512), scope="trn.sbuf") - B = T.alloc_buffer((128, 512), scope="trn.sbuf") - C = T.alloc_buffer((128, 1), scope="trn.sbuf") + A = T.alloc_tensor((128, 512), scope="trn.sbuf") + B = T.alloc_tensor((128, 512), scope="trn.sbuf") + C = T.alloc_tensor((128, 1), scope="trn.sbuf") for b_loop in range(1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -102,23 +102,23 @@ def test_activation_reduce_in_loop(): @T.prim_func def activation_reduce(): T.device_entry() - A = T.alloc_buffer(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) - B = T.alloc_buffer(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) - C = T.alloc_buffer(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) + A = T.alloc_tensor(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) + B = T.alloc_tensor(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) + C = T.alloc_tensor(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) for i in range(2): Tx.unary_reduce(B, C, A[i*16:i*16+16], "sqrt", "sum", reduce_axes=1) @T.prim_func def expected(): T.func_attr({"global_symbol": "activation_reduce"}) - const_bias = T.alloc_buffer((128, 512), scope="trn.sbuf") + const_bias = T.alloc_tensor((128, 512), scope="trn.sbuf") with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for f_loop in T.serial(512, annotations={"nki_dim": "F"}): T.nki.memset(const_bias[p_loop, f_loop], T.float32(0.0)) - A = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B = T.alloc_buffer((128, 8192), scope="trn.sbuf") - C = T.alloc_buffer((128, 16), scope="trn.sbuf") + A = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B = T.alloc_tensor((128, 8192), scope="trn.sbuf") + C = T.alloc_tensor((128, 16), scope="trn.sbuf") for i, b_loop in T.grid(2, 16): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -144,23 +144,23 @@ def test_activation_reduce_in_loop2(): @T.prim_func def activation_reduce(): T.device_entry() - A = T.alloc_buffer(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) - B = T.alloc_buffer(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) - C = T.alloc_buffer(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) + A = T.alloc_tensor(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) + B = T.alloc_tensor(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) + C = T.alloc_tensor(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) for i in range(2): Tx.unary_reduce(B, C, A[i*16:i*16+16], "sqrt", "sum", reduce_axes=1) @T.prim_func def expected(): T.func_attr({"global_symbol": "activation_reduce"}) - const_bias = T.alloc_buffer((128, 512), scope="trn.sbuf") + const_bias = T.alloc_tensor((128, 512), scope="trn.sbuf") with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for f_loop in T.serial(512, annotations={"nki_dim": "F"}): T.nki.memset(const_bias[p_loop, f_loop], T.float32(0.0)) - A = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B = T.alloc_buffer((128, 8192), scope="trn.sbuf") - C = T.alloc_buffer((128, 16), scope="trn.sbuf") + A = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B = T.alloc_tensor((128, 8192), scope="trn.sbuf") + C = T.alloc_tensor((128, 16), scope="trn.sbuf") for i, b_loop in T.grid(2, 16): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -186,24 +186,24 @@ def test_activation_reduce_two_stage(): @T.prim_func def activation_reduce(): T.device_entry() - A = T.alloc_buffer(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) - B = T.alloc_buffer(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) - C = T.alloc_buffer(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) + A = T.alloc_tensor(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) + B = T.alloc_tensor(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) + C = T.alloc_tensor(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) for i in range(2): Tx.unary_reduce(B, C, A[i*16:i*16+16], "sqrt", "sum", reduce_axes=(0,1)) @T.prim_func def expected(): T.func_attr({"global_symbol": "activation_reduce"}) - partial_reduce = T.alloc_buffer((128, 8), scope="trn.sbuf") - const_bias = T.alloc_buffer((128, 1024), scope="trn.sbuf") + partial_reduce = T.alloc_tensor((128, 8), scope="trn.sbuf") + const_bias = T.alloc_tensor((128, 1024), scope="trn.sbuf") with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for f_loop in T.serial(1024, annotations={"nki_dim": "F"}): T.nki.memset(const_bias[p_loop, f_loop], T.float32(0.0)) - A = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B = T.alloc_buffer((128, 8192), scope="trn.sbuf") - C = T.alloc_buffer((128, 1), scope="trn.sbuf") + A = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B = T.alloc_tensor((128, 8192), scope="trn.sbuf") + C = T.alloc_tensor((128, 1), scope="trn.sbuf") for i, b_loop in T.grid(2, 1): for reduction_b_loop in range(8): T.attr(0, "tensorized_nki_instruction", 1) @@ -236,20 +236,20 @@ def test_activation_reduce_with_bias_scale(): @T.prim_func def activation_reduce(): T.device_entry() - A = T.alloc_buffer(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) - B = T.alloc_buffer(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) - C = T.alloc_buffer(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) - bias = T.alloc_buffer(bias_shape, dtype="float32", scope="trn.sbuf", layout=bias_layout) + A = T.alloc_tensor(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) + B = T.alloc_tensor(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) + C = T.alloc_tensor(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) + bias = T.alloc_tensor(bias_shape, dtype="float32", scope="trn.sbuf", layout=bias_layout) for i in range(2): Tx.unary_reduce(B, C, A[i*16:i*16+16], "sqrt", "sum", reduce_axes=1, bias=bias, scale=2.0) # noqa: E501 @T.prim_func def expected(): T.func_attr({"global_symbol": "activation_reduce"}) - A = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B = T.alloc_buffer((128, 8192), scope="trn.sbuf") - C = T.alloc_buffer((128, 16), scope="trn.sbuf") - bias = T.alloc_buffer((128, 1), scope="trn.sbuf") + A = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B = T.alloc_tensor((128, 8192), scope="trn.sbuf") + C = T.alloc_tensor((128, 16), scope="trn.sbuf") + bias = T.alloc_tensor((128, 1), scope="trn.sbuf") for i, b_loop in T.grid(2, 16): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -274,17 +274,17 @@ def test_simple_tensor_scalar_reduce(): @T.prim_func def tensor_scalar_reduce(): T.device_entry() - A = T.alloc_buffer(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) - B = T.alloc_buffer(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) - C = T.alloc_buffer(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) + A = T.alloc_tensor(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) + B = T.alloc_tensor(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) + C = T.alloc_tensor(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) Tx.binary_reduce(B, C, A, 1.0, "add", "sum", reduce_axes=1) @T.prim_func def expected(): T.func_attr({"global_symbol": "tensor_scalar_reduce"}) - A = T.alloc_buffer((128, 512), scope="trn.sbuf") - B = T.alloc_buffer((128, 512), scope="trn.sbuf") - C = T.alloc_buffer((128, 1), scope="trn.sbuf") + A = T.alloc_tensor((128, 512), scope="trn.sbuf") + B = T.alloc_tensor((128, 512), scope="trn.sbuf") + C = T.alloc_tensor((128, 1), scope="trn.sbuf") for b_loop in range(1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -311,10 +311,10 @@ def test_tensor_tensor_reduce_fail(): @T.prim_func def tensor_scalar_reduce(): T.device_entry() - A = T.alloc_buffer(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) - B = T.alloc_buffer(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) - C = T.alloc_buffer(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) - D = T.alloc_buffer(D_shape, dtype="float32", scope="trn.sbuf", layout=D_layout) + A = T.alloc_tensor(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) + B = T.alloc_tensor(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) + C = T.alloc_tensor(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) + D = T.alloc_tensor(D_shape, dtype="float32", scope="trn.sbuf", layout=D_layout) Tx.binary_reduce(B, C, A, D, "add", "sum", reduce_axes=1) # fmt: off @@ -338,19 +338,19 @@ def test_tensor_scalar_reduce_complex(): @T.prim_func def tensor_scalar_reduce() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) - D_sbuf = T.alloc_buffer(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + D_sbuf = T.alloc_tensor(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 Tx.binary_reduce(C_sbuf, D_sbuf, B_sbuf, A_sbuf, "add", "sum", reduce_axes=0) @T.prim_func def expected(): T.func_attr({"global_symbol": "tensor_scalar_reduce"}) - A_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - D_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + D_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") for b_loop in range(512): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -375,18 +375,18 @@ def test_tensor_scalar_reduce_two_stage(): @T.prim_func def tensor_scalar_reduce() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(dst1_shape, "float32", scope="trn.sbuf", layout=dst1_layout) - C_sbuf = T.alloc_buffer(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(dst1_shape, "float32", scope="trn.sbuf", layout=dst1_layout) + C_sbuf = T.alloc_tensor(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 Tx.binary_reduce(B_sbuf, C_sbuf, A_sbuf, 1.0, "add", "sum", reduce_axes=(1, 2)) @T.prim_func def expected(): T.func_attr({"global_symbol": "tensor_scalar_reduce"}) - partial_reduce = T.alloc_buffer((128, 4), scope="trn.sbuf") - A_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") + partial_reduce = T.alloc_tensor((128, 4), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") for b_loop in range(4): for reduction_b_loop in range(4): T.attr(0, "tensorized_nki_instruction", 1) @@ -419,21 +419,21 @@ def test_vector_chain(): @T.prim_func def binary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) - _C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) - D_sbuf = T.alloc_buffer(src3_shape, "float32", scope="trn.sbuf", layout=src3_layout) - E_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) + _C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + D_sbuf = T.alloc_tensor(src3_shape, "float32", scope="trn.sbuf", layout=src3_layout) + E_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.binary_chain(E_sbuf, A_sbuf, B_sbuf, D_sbuf, "add", "add", reverse1=True) @T.prim_func def expected(): T.func_attr({"global_symbol": "binary"}) - A_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") - _C_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - D_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") - E_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") + _C_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + D_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") + E_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") for b_loop in T.serial(0, 512): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -461,21 +461,21 @@ def test_vector_chain_2(): @T.prim_func def binary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) - _C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) - D_sbuf = T.alloc_buffer(src3_shape, "float32", scope="trn.sbuf", layout=src3_layout) - E_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) + _C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + D_sbuf = T.alloc_tensor(src3_shape, "float32", scope="trn.sbuf", layout=src3_layout) + E_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.binary_chain(E_sbuf, A_sbuf, B_sbuf, D_sbuf, "add", "add", reverse1=True) @T.prim_func def expected(): T.func_attr({"global_symbol": "binary"}) - A_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") - _C_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - D_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - E_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") + _C_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + D_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + E_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") for b_loop in T.serial(0, 512): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -499,16 +499,16 @@ def test_reduce_negate(): @T.prim_func def reduction(): T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(4): Tx.reduce_negate(B_sbuf[:, i], A_sbuf[:, :, i], reduce_op="sum", reduce_axes=-2) @T.prim_func def expected(): T.func_attr({"global_symbol": "reduction"}) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") for i, b_loop in T.grid(4, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -533,9 +533,9 @@ def test_binary_reduce_guard(): @T.prim_func def binary_reduce() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) - C_sbuf = T.alloc_buffer(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + C_sbuf = T.alloc_tensor(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 for j in range(4): for i in range(4): Tx.binary_reduce(B_sbuf[0:128*(j+1), 0:128*(i+1)], C_sbuf[0:128*(j+1)], A_sbuf[0:128*(j+1), 0:128*(i+1)], 0.0, "add", "sum", [-1]) # noqa: E501 @@ -543,9 +543,9 @@ def binary_reduce() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "binary_reduce"}) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") for j, i, b_loop in T.grid(4, 4, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -572,9 +572,9 @@ def test_unary_reduce_guard(): @T.prim_func def unary_reduce() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) - C_sbuf = T.alloc_buffer(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + C_sbuf = T.alloc_tensor(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 for j in range(4): for i in range(4): Tx.unary_reduce(B_sbuf[0:128*(j+1), 0:128*(i+1)], C_sbuf[0:128*(j+1)], A_sbuf[0:128*(j+1), 0:128*(i+1)], "sqrt", "sum", reduce_axes=[-1]) # noqa: E501 @@ -582,14 +582,14 @@ def unary_reduce() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "unary_reduce"}) - const_bias = T.alloc_buffer((128, 512), scope="trn.sbuf") + const_bias = T.alloc_tensor((128, 512), scope="trn.sbuf") with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for f_loop in T.serial(512, annotations={"nki_dim": "F"}): T.nki.memset(const_bias[p_loop, f_loop], T.float32(0.0)) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") for j, i, b_loop in T.grid(4, 4, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(128, annotations={"nki_dim": "P"}): @@ -618,18 +618,18 @@ def test_binary_chain_guard(): @T.prim_func def binary_chain() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for j in range(4): for i in range(4): Tx.binary_chain(C_sbuf[0:128*(j+1), 0:128*(i+1)], A_sbuf[0:128*(j+1), 0:128*(i+1)], B_sbuf[0:128*(j+1), 0], 1.0, "add", "sub", reverse1=True) # noqa: E501 @T.prim_func def expected(): T.func_attr({"global_symbol": "binary_chain"}) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") for j, i, b_loop in T.grid(4, 4, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -656,25 +656,25 @@ def test_activation_reduce_two_stage_workspace(): @T.prim_func def activation_reduce(): T.device_entry() - intermediate_buffer = T.alloc_buffer((128, 16), scope="trn.sbuf") - A = T.alloc_buffer(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) - B = T.alloc_buffer(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) - C = T.alloc_buffer(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) + intermediate_buffer = T.alloc_tensor((128, 16), scope="trn.sbuf") + A = T.alloc_tensor(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) + B = T.alloc_tensor(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) + C = T.alloc_tensor(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) for i in range(2): Tx.unary_reduce(B, C, A[i*16:i*16+16], "sqrt", "sum", reduce_axes=(0,1), workspace={"partial_reduce": intermediate_buffer}) # noqa: E501 @T.prim_func def expected(): T.func_attr({"global_symbol": "activation_reduce"}) - const_bias = T.alloc_buffer((128, 1024), scope="trn.sbuf") + const_bias = T.alloc_tensor((128, 1024), scope="trn.sbuf") with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for f_loop in T.serial(1024, annotations={"nki_dim": "F"}): T.nki.memset(const_bias[p_loop, f_loop], T.float32(0.0)) - intermediate_buffer = T.alloc_buffer((128, 16), scope="trn.sbuf") - A = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B = T.alloc_buffer((128, 8192), scope="trn.sbuf") - C = T.alloc_buffer((128, 1), scope="trn.sbuf") + intermediate_buffer = T.alloc_tensor((128, 16), scope="trn.sbuf") + A = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B = T.alloc_tensor((128, 8192), scope="trn.sbuf") + C = T.alloc_tensor((128, 1), scope="trn.sbuf") for i, b_loop in T.grid(2, 1): for reduction_b_loop in range(8): T.attr(0, "tensorized_nki_instruction", 1) @@ -706,19 +706,19 @@ def test_tensor_scalar_reduce_two_stage_workspace(): @T.prim_func def tensor_scalar_reduce() -> None: T.device_entry() - intermediate_buffer = T.alloc_buffer((128, 8), scope="trn.sbuf") - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(dst1_shape, "float32", scope="trn.sbuf", layout=dst1_layout) - C_sbuf = T.alloc_buffer(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 + intermediate_buffer = T.alloc_tensor((128, 8), scope="trn.sbuf") + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(dst1_shape, "float32", scope="trn.sbuf", layout=dst1_layout) + C_sbuf = T.alloc_tensor(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 Tx.binary_reduce(B_sbuf, C_sbuf, A_sbuf, 1.0, "add", "sum", reduce_axes=(1, 2), workspace={"partial_reduce": intermediate_buffer}) # noqa: E501 @T.prim_func def expected(): T.func_attr({"global_symbol": "tensor_scalar_reduce"}) - intermediate_buffer = T.alloc_buffer((128, 8), scope="trn.sbuf") - A_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") + intermediate_buffer = T.alloc_tensor((128, 8), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") for b_loop in range(4): for reduction_b_loop in range(4): T.attr(0, "tensorized_nki_instruction", 1) @@ -741,20 +741,20 @@ def test_unary_reduce_complex(): @T.prim_func def unary_reduce(): T.device_entry() - p = T.alloc_buffer((128, 8192), "float16", scope="trn.sbuf", layout="PF") - rowsum_p = T.alloc_buffer((2, 128, 1), scope="trn.sbuf", layout="FPF") - qk = T.alloc_buffer((2, 128, 8192), scope="trn.sbuf", layout="FPF") - running_max = T.alloc_buffer((16384, 1), dtype="float32", scope="trn.sbuf", layout="PF") + p = T.alloc_tensor((128, 8192), "float16", scope="trn.sbuf", layout="PF") + rowsum_p = T.alloc_tensor((2, 128, 1), scope="trn.sbuf", layout="FPF") + qk = T.alloc_tensor((2, 128, 8192), scope="trn.sbuf", layout="FPF") + running_max = T.alloc_tensor((16384, 1), dtype="float32", scope="trn.sbuf", layout="PF") for i in range(4): Tx.unary_reduce(p[0:128, 0:8192], rowsum_p[i % 2, 0:128, 0], qk[i % 2, 0:128, 0:8192], "exp", "sum", bias=running_max[i * 128:i * 128 + 128, 0]) # noqa: E501 @T.prim_func def expected(): T.func_attr({"global_symbol": "unary_reduce"}) - p = T.alloc_buffer((128, 8192), "float16", scope="trn.sbuf") - rowsum_p = T.alloc_buffer((128, 2), scope="trn.sbuf") - qk = T.alloc_buffer((128, 16384), scope="trn.sbuf") - running_max = T.alloc_buffer((128, 128), scope="trn.sbuf") + p = T.alloc_tensor((128, 8192), "float16", scope="trn.sbuf") + rowsum_p = T.alloc_tensor((128, 2), scope="trn.sbuf") + qk = T.alloc_tensor((128, 16384), scope="trn.sbuf") + running_max = T.alloc_tensor((128, 128), scope="trn.sbuf") for i, b_loop in T.grid(4, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(128, annotations={"nki_dim": "P"}): diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_copy_trn.py b/tests/python/tirx/operator/tile_primitive/trn/test_copy_trn.py index cc0c06b5192f..3d0953337c6a 100644 --- a/tests/python/tirx/operator/tile_primitive/trn/test_copy_trn.py +++ b/tests/python/tirx/operator/tile_primitive/trn/test_copy_trn.py @@ -55,17 +55,17 @@ def test_simple_copy(): dst_layout = TileLayout(S[(128, 512) : (1 @ P, 1 @ F)]) @T.prim_func - def copy(A: T.Buffer(src_shape, "float32", layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, "float32", layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.copy(A_sbuf, A) @T.prim_func - def expected(A: T.Buffer((128, 512), layout=None)): + def expected(A: T.Tensor((128, 512), layout=None)): T.func_attr({"global_symbol": "copy"}) - A_1 = T.decl_buffer((65536,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") + A_1 = T.decl_tensor((65536,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") for b_loop in T.serial(0, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim": "P"}): @@ -86,17 +86,17 @@ def test_simple_copy_2(): dst_layout = TileLayout(S[(128, 4, 128) : (4 @ F, 1 @ F, 1 @ P)]) @T.prim_func - def copy(A: T.Buffer(src_shape, "float32", layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, "float32", layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.copy(A_sbuf, A) @T.prim_func - def expected(A: T.Buffer((128, 512), layout=None)): + def expected(A: T.Tensor((128, 512), layout=None)): T.func_attr({"global_symbol": "copy"}) - A_1 = T.decl_buffer((65536,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") + A_1 = T.decl_tensor((65536,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") for b_loop in T.serial(0, 512): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim": "P"}): @@ -116,18 +116,18 @@ def test_copy_in_a_loop(): dst_layout = TileLayout(S[(4, 128, 512) : (512 @ F, 1 @ P, 1 @ F)]) @T.prim_func - def copy(A: T.Buffer(src_shape, "float32", layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, "float32", layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(4): Tx.copy(A_sbuf[i * 128 : i * 128 + 128, :], A[i * 128 : i * 128 + 128, :]) @T.prim_func - def expected(A: T.Buffer((512, 512), layout=None)): + def expected(A: T.Tensor((512, 512), layout=None)): T.func_attr({"global_symbol": "copy"}) - A_1 = T.decl_buffer((262144,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") + A_1 = T.decl_tensor((262144,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") for i, b_loop in T.grid(4, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim": "P"}): @@ -149,22 +149,22 @@ def test_copy_in_a_loop_2(): dst_layout = TileLayout(S[(128, 2048) : (1 @ P, 1 @ F)]) @T.prim_func - def copy(A: T.Buffer(src_shape, "float32", layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, "float32", layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) A_sbuf_view = A_sbuf.view(128, 4, 512) A_view = A.view(128, 4, 512) for i in range(4): Tx.copy(A_sbuf_view[:, i, :], A_view[:, i, :]) @T.prim_func - def expected(A: T.Buffer((512, 512), layout=None)): + def expected(A: T.Tensor((512, 512), layout=None)): T.func_attr({"global_symbol": "copy"}) - _A_flat = T.decl_buffer((262144,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - A_sbuf_view = T.decl_buffer((128, 2048), data=A_sbuf.data, scope="trn.sbuf", layout=None) - A_view = T.decl_buffer((262144,), data=_A_flat.data, layout=None) + _A_flat = T.decl_tensor((262144,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + A_sbuf_view = T.decl_tensor((128, 2048), data=A_sbuf.data, scope="trn.sbuf", layout=None) + A_view = T.decl_tensor((262144,), data=_A_flat.data, layout=None) for i, b_loop in T.grid(4, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim": "P"}): @@ -191,21 +191,21 @@ def test_copy_transpose(): @T.prim_func def copy() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.copy(B_sbuf, A_sbuf) @T.prim_func def expected(): T.func_attr({"global_symbol": "copy"}) - identity = T.alloc_buffer((128, 128), scope="trn.sbuf") - acc_psum = T.alloc_buffer((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) + identity = T.alloc_tensor((128, 128), scope="trn.sbuf") + acc_psum = T.alloc_tensor((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for rhs_f_loop in T.serial(128, annotations={"nki_dim": "F"}): T.nki.identity(identity[p_loop, rhs_f_loop], 128) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") for b_loop in range(16): for extend_b_loop in range(1): T.attr(0, "tensorized_nki_instruction", 1) @@ -237,22 +237,22 @@ def test_copy_transpose_2(): @T.prim_func def copy() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(4): Tx.copy(B_sbuf[i, :], A_sbuf) @T.prim_func def expected(): T.func_attr({"global_symbol": "copy"}) - identity = T.alloc_buffer((128, 128), scope="trn.sbuf") - acc_psum = T.alloc_buffer((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) + identity = T.alloc_tensor((128, 128), scope="trn.sbuf") + acc_psum = T.alloc_tensor((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for rhs_f_loop in T.serial(128, annotations={"nki_dim": "F"}): T.nki.identity(identity[p_loop, rhs_f_loop], 128) - A_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") for i in range(4): for b_loop in range(4): for extend_b_loop in range(1): @@ -283,15 +283,15 @@ def test_copy_different_f(): @T.prim_func def copy() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.copy(B_sbuf, A_sbuf) @T.prim_func def expected(): T.func_attr({"global_symbol": "copy"}) - A_sbuf = T.alloc_buffer((128, 256), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 256), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 256), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 256), scope="trn.sbuf") for b_loop in T.serial(0, 64): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim": "P"}): @@ -319,17 +319,17 @@ def test_copy_different_shape(): @T.prim_func def copy() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) B_sbuf_view = B_sbuf.view(512, 4) Tx.copy(B_sbuf_view, A_sbuf[:, 0:4]) @T.prim_func def expected(): T.func_attr({"global_symbol": "copy"}) - A_sbuf = T.alloc_buffer((128, 256), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 16), scope="trn.sbuf") - B_sbuf_view = T.decl_buffer((128, 16), data=B_sbuf.data, scope="trn.sbuf", layout=None) + A_sbuf = T.alloc_tensor((128, 256), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 16), scope="trn.sbuf") + B_sbuf_view = T.decl_tensor((128, 16), data=B_sbuf.data, scope="trn.sbuf", layout=None) for b_loop in T.serial(0, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim": "P"}): @@ -352,18 +352,18 @@ def test_copy_irregular_shape(): dst_layout = TileLayout(S[(128, 512) : (1 @ P, 1 @ F)]) @T.prim_func - def copy(A: T.Buffer(src_shape, "float32", layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, "float32", layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(4): Tx.copy(A[:, i * 512 : i * 512 + 512], A_sbuf) @T.prim_func - def expected(A: T.Buffer((128, 10000), layout=None)): + def expected(A: T.Tensor((128, 10000), layout=None)): T.func_attr({"global_symbol": "copy"}) - A_1 = T.decl_buffer((1280000,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") + A_1 = T.decl_tensor((1280000,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") for i, b_loop in T.grid(4, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim": "P"}): @@ -384,19 +384,19 @@ def test_copy_different_shape_dim(): # fmt: off @T.prim_func - def copy(A: T.Buffer(src_shape, 'float32', layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, 'float32', layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(32): Tx.copy(A_sbuf, A[i, :, :]) @T.prim_func - def expected(A: T.Buffer((32, 128, 512), layout=None)): + def expected(A: T.Tensor((32, 128, 512), layout=None)): T.func_attr({"global_symbol": "copy"}) - A_1 = T.decl_buffer((2097152,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") + A_1 = T.decl_tensor((2097152,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") for i, b_loop in T.grid(32, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -416,18 +416,18 @@ def test_copy_with_offset(): dst_layout = TileLayout(S[(4, 128, 512) : (512 @ F, 1 @ P, 1 @ F)]) @T.prim_func - def copy(A: T.Buffer(src_shape, "float32", layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, "float32", layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(2): Tx.copy(A_sbuf[i * 256 : i * 256 + 256, :], A) @T.prim_func - def expected(A: T.Buffer((256, 512), layout=None)): + def expected(A: T.Tensor((256, 512), layout=None)): T.func_attr({"global_symbol": "copy"}) - A_1 = T.decl_buffer((131072,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") + A_1 = T.decl_tensor((131072,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") for i, b_loop in T.grid(2, 2): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim": "P"}): @@ -450,18 +450,18 @@ def test_large_dma_copy(): dst_layout = TileLayout(S[(4, 128, 4096) : (4096 @ F, 1 @ P, 1 @ F)]) @T.prim_func - def copy(A: T.Buffer(src_shape, "float32", layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, "float32", layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(4): Tx.copy(A_sbuf[i * 128 : i * 128 + 128, :], A[i * 128 : i * 128 + 128, :]) @T.prim_func - def expected(A: T.Buffer((512, 4096), layout=None)): + def expected(A: T.Tensor((512, 4096), layout=None)): T.func_attr({"global_symbol": "copy"}) - A_1 = T.decl_buffer((2097152,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") + A_1 = T.decl_tensor((2097152,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") for i, b_loop in T.grid(4, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim": "P"}): @@ -486,16 +486,16 @@ def test_copy_with_inst_size_limit(): @T.prim_func def copy(A_ptr: T.handle) -> None: T.device_entry() - B_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + B_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(4): Tx.copy(A_sbuf[i * 128 : i * 128 + 128, :], B_sbuf[i * 128 : i * 128 + 128, :]) @T.prim_func def expected(A_ptr: T.handle): T.func_attr({"global_symbol": "copy"}) - B_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - A_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") for i, b_loop in T.grid(4, 8): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim": "P"}): @@ -519,18 +519,18 @@ def test_copy_with_complex_index(): # fmt: off @T.prim_func - def copy(A: T.Buffer(A_shape, 'float32', layout=A_layout), ) -> None: + def copy(A: T.Tensor(A_shape, 'float32', layout=A_layout), ) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(A_sbuf_shape, "float32", scope="trn.sbuf", layout=A_sbuf_layout) + A_sbuf = T.alloc_tensor(A_sbuf_shape, "float32", scope="trn.sbuf", layout=A_sbuf_layout) Tx.copy(A_sbuf[1, 0:2048, 0:1024], A[2048: 4096, 3072:4096]) @T.prim_func - def expected(A: T.Buffer((4096, 4096), layout=None)): + def expected(A: T.Tensor((4096, 4096), layout=None)): T.func_attr({"global_symbol": "copy"}) - A_1 = T.decl_buffer((16777216,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 32768), scope="trn.sbuf") + A_1 = T.decl_tensor((16777216,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 32768), scope="trn.sbuf") for b_loop in T.serial(0, 8): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -551,18 +551,18 @@ def test_copy_with_complex_index_2(): # fmt: off @T.prim_func - def copy(A: T.Buffer(A_shape, 'float32', layout=A_layout), ) -> None: + def copy(A: T.Tensor(A_shape, 'float32', layout=A_layout), ) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(A_sbuf_shape, "float32", scope="trn.sbuf", layout=A_sbuf_layout) + A_sbuf = T.alloc_tensor(A_sbuf_shape, "float32", scope="trn.sbuf", layout=A_sbuf_layout) Tx.copy(A_sbuf[2048: 4096, 3072:4096], A[1, 0:2048, 0:1024]) @T.prim_func - def expected(A: T.Buffer((2, 2048, 1024), layout=None)): + def expected(A: T.Tensor((2, 2048, 1024), layout=None)): T.func_attr({"global_symbol": "copy"}) - A_1 = T.decl_buffer((4194304,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 131072), scope="trn.sbuf") + A_1 = T.decl_tensor((4194304,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 131072), scope="trn.sbuf") for b_loop in T.serial(0, 8): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -586,10 +586,10 @@ def test_copy_transpose_with_workspace(): @T.prim_func def copy() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) - identity = T.alloc_buffer((128, 128), "float32", scope="trn.sbuf") - acc_psum = T.alloc_buffer((1, 128, 512), "float32", scope="trn.psum", allocated_addr=(0, 0)) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + identity = T.alloc_tensor((128, 128), "float32", scope="trn.sbuf") + acc_psum = T.alloc_tensor((1, 128, 512), "float32", scope="trn.psum", allocated_addr=(0, 0)) with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): for rhs_f_loop in T.serial(0, 128, annotations={"nki_dim":"F"}): @@ -599,10 +599,10 @@ def copy() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "copy"}) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - identity = T.alloc_buffer((128, 128), scope="trn.sbuf") - acc_psum = T.alloc_buffer((1, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + identity = T.alloc_tensor((128, 128), scope="trn.sbuf") + acc_psum = T.alloc_tensor((1, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for rhs_f_loop in T.serial(128, annotations={"nki_dim": "F"}): @@ -633,20 +633,20 @@ def test_copy_with_guard(): # fmt: off @T.prim_func - def copy(A: T.Buffer(src_shape, 'float32', layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, 'float32', layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for j in range(4): for i in range(4): Tx.copy(A_sbuf[i * 128 : i * 128 + 128, 0:128*j], A[i * 128 : i * 128 + 128, 0:128*j]) # noqa: E501 @T.prim_func - def expected(A: T.Buffer((512, 512), layout=None)): + def expected(A: T.Tensor((512, 512), layout=None)): T.func_attr({"global_symbol": "copy"}) - A_1 = T.decl_buffer((262144,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") + A_1 = T.decl_tensor((262144,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") for j, i, b_loop in T.grid(4, 4, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -669,20 +669,20 @@ def test_copy_with_guard_2(): # fmt: off @T.prim_func - def copy(A: T.Buffer(src_shape, 'float32', layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, 'float32', layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for j in range(4): for i in range(4): Tx.copy(A_sbuf[0:128*j, 0:128*i], A[0:128*j, 0:128*i]) @T.prim_func - def expected(A: T.Buffer((512, 512), layout=None)): + def expected(A: T.Tensor((512, 512), layout=None)): T.func_attr({"global_symbol": "copy"}) - A_1 = T.decl_buffer((262144,), data=A.data, layout=None) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") + A_1 = T.decl_tensor((262144,), data=A.data, layout=None) + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") for j, i, b_loop in T.grid(4, 4, 3): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -707,8 +707,8 @@ def test_copy_transpose_with_guard(): @T.prim_func def copy() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(4): for j in range(4): Tx.copy(B_sbuf[i * 128 : i * 128 + 128, 0:128*j], A_sbuf[i * 128 : i * 128 + 128, 0:128*j]) # noqa: E501 @@ -716,14 +716,14 @@ def copy() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "copy"}) - identity = T.alloc_buffer((128, 128), scope="trn.sbuf") - acc_psum = T.alloc_buffer((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) + identity = T.alloc_tensor((128, 128), scope="trn.sbuf") + acc_psum = T.alloc_tensor((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for rhs_f_loop in T.serial(128, annotations={"nki_dim": "F"}): T.nki.identity(identity[p_loop, rhs_f_loop], 128) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") for i, j, b_loop in T.grid(4, 4, 3): for extend_b_loop in range(1): T.attr(0, "tensorized_nki_instruction", 1) @@ -756,15 +756,15 @@ def test_copy_with_specified_max_inst_size(): @T.prim_func def copy(A_ptr: T.handle) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.copy(A_sbuf, B_sbuf, max_inst_size=128) @T.prim_func def expected(A_ptr: T.handle): T.func_attr({"global_symbol": "copy"}) - A_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf", layout=None) - B_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf", layout=None) + A_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf", layout=None) + B_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf", layout=None) for b_loop in T.serial(0, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(128, annotations={"nki_dim": "P"}): @@ -782,21 +782,21 @@ def test_copy_transpose_with_extended_f(): @T.prim_func def copy(A_ptr: T.handle) -> None: T.device_entry() - A_sbuf = T.alloc_buffer((128, 2048), "float32", scope="trn.sbuf", layout="PF") - B_sbuf = T.alloc_buffer((128, 2048), "float32", scope="trn.sbuf", layout="FP") + A_sbuf = T.alloc_tensor((128, 2048), "float32", scope="trn.sbuf", layout="PF") + B_sbuf = T.alloc_tensor((128, 2048), "float32", scope="trn.sbuf", layout="FP") Tx.copy(B_sbuf, A_sbuf) @T.prim_func def expected(A_ptr: T.handle): T.func_attr({"global_symbol": "copy"}) - identity = T.alloc_buffer((128, 128), scope="trn.sbuf") - acc_psum = T.alloc_buffer((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) + identity = T.alloc_tensor((128, 128), scope="trn.sbuf") + acc_psum = T.alloc_tensor((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for rhs_f_loop in T.serial(128, annotations={"nki_dim": "F"}): T.nki.identity(identity[p_loop, rhs_f_loop], 128) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") for b_loop in range(4): for extend_b_loop in range(4): T.attr(0, "tensorized_nki_instruction", 1) diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_gemm_trn.py b/tests/python/tirx/operator/tile_primitive/trn/test_gemm_trn.py index 563e764ee9a9..e39e37198de5 100644 --- a/tests/python/tirx/operator/tile_primitive/trn/test_gemm_trn.py +++ b/tests/python/tirx/operator/tile_primitive/trn/test_gemm_trn.py @@ -58,17 +58,17 @@ def test_simple_gemm(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((128, 128), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((128, 128), "float32", scope="trn.sbuf", layout=B_layout) - C_psum = T.alloc_buffer((128, 128), "float32", scope="trn.psum", layout=C_layout) + A_sbuf = T.alloc_tensor((128, 128), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((128, 128), "float32", scope="trn.sbuf", layout=B_layout) + C_psum = T.alloc_tensor((128, 128), "float32", scope="trn.psum", layout=C_layout) Tx.gemm(C_psum, A_sbuf, B_sbuf, C_psum) @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - A_sbuf = T.alloc_buffer((128, 128), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 128), scope="trn.sbuf") - C_psum = T.alloc_buffer((1, 128, 128), scope="trn.psum") + A_sbuf = T.alloc_tensor((128, 128), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 128), scope="trn.sbuf") + C_psum = T.alloc_tensor((1, 128, 128), scope="trn.psum") for lhs_b_loop, rhs_b_loop, reduction_b_loop in T.grid(1, 1, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -92,17 +92,17 @@ def test_larger_gemm(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((256, 512), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((512, 256), "float32", scope="trn.sbuf", layout=B_layout) - C_psum = T.alloc_buffer((256, 256), "float32", scope="trn.psum", layout=C_layout) + A_sbuf = T.alloc_tensor((256, 512), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((512, 256), "float32", scope="trn.sbuf", layout=B_layout) + C_psum = T.alloc_tensor((256, 256), "float32", scope="trn.psum", layout=C_layout) Tx.gemm(C_psum, A_sbuf, B_sbuf, C_psum) @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - A_sbuf = T.alloc_buffer((128, 1024), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 1024), scope="trn.sbuf") - C_psum = T.alloc_buffer((1, 128, 512), scope="trn.psum") + A_sbuf = T.alloc_tensor((128, 1024), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 1024), scope="trn.sbuf") + C_psum = T.alloc_tensor((1, 128, 512), scope="trn.psum") for lhs_b_loop, rhs_b_loop, reduction_b_loop in T.grid(2, 1, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -126,9 +126,9 @@ def test_gemm_in_a_loop(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) - C_psum = T.alloc_buffer((512, 256), "float32", scope="trn.psum", layout=C_layout) + A_sbuf = T.alloc_tensor((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) + C_psum = T.alloc_tensor((512, 256), "float32", scope="trn.psum", layout=C_layout) for i in range(2): for k in range(2): Tx.gemm( @@ -141,9 +141,9 @@ def gemm() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - C_psum = T.alloc_buffer((2, 128, 512), scope="trn.psum") + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + C_psum = T.alloc_tensor((2, 128, 512), scope="trn.psum") for i, k, lhs_b_loop, rhs_b_loop, reduction_b_loop in T.grid(2, 2, 2, 1, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -167,9 +167,9 @@ def test_gemm_with_stride(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((512, 512, 2), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((512, 2, 256), "float32", scope="trn.sbuf", layout=B_layout) - C_psum = T.alloc_buffer((512, 256), "float32", scope="trn.psum", layout=C_layout) + A_sbuf = T.alloc_tensor((512, 512, 2), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((512, 2, 256), "float32", scope="trn.sbuf", layout=B_layout) + C_psum = T.alloc_tensor((512, 256), "float32", scope="trn.psum", layout=C_layout) for i in range(2): for k in range(2): Tx.gemm( @@ -182,9 +182,9 @@ def gemm() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 4095), scope="trn.sbuf") - C_psum = T.alloc_buffer((2, 128, 512), scope="trn.psum") + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 4095), scope="trn.sbuf") + C_psum = T.alloc_tensor((2, 128, 512), scope="trn.psum") for i, k, lhs_b_loop, rhs_b_loop, reduction_b_loop in T.grid(2, 2, 2, 1, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -209,9 +209,9 @@ def test_gemm_swap_lhs_rhs(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) - C_psum = T.alloc_buffer((512, 256), "float32", scope="trn.psum", layout=C_layout) + A_sbuf = T.alloc_tensor((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) + C_psum = T.alloc_tensor((512, 256), "float32", scope="trn.psum", layout=C_layout) for i in range(2): for k in range(2): Tx.gemm( @@ -224,9 +224,9 @@ def gemm() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - C_psum = T.alloc_buffer((2, 128, 512), scope="trn.psum") + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + C_psum = T.alloc_tensor((2, 128, 512), scope="trn.psum") for i, k, lhs_b_loop, rhs_b_loop, reduction_b_loop in T.grid(2, 2, 2, 2, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -250,9 +250,9 @@ def test_gemm_with_sbuf_output(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) - C_sbuf = T.alloc_buffer((512, 256), "float32", scope="trn.sbuf", layout=C_layout) + A_sbuf = T.alloc_tensor((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) + C_sbuf = T.alloc_tensor((512, 256), "float32", scope="trn.sbuf", layout=C_layout) for i in range(2): for k in range(2): Tx.gemm( @@ -264,10 +264,10 @@ def gemm() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - buffer = T.alloc_buffer((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 1024), scope="trn.sbuf") + buffer = T.alloc_tensor((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 1024), scope="trn.sbuf") for i, k, lhs_b_loop, rhs_b_loop in T.grid(2, 2, 2, 2): for reduction_b_loop in range(4): T.attr(0, "tensorized_nki_instruction", 1) @@ -298,9 +298,9 @@ def test_gemm_different_shape(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((2, 512, 1024), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) - C_psum = T.alloc_buffer((512, 256), "float32", scope="trn.psum", layout=C_layout) + A_sbuf = T.alloc_tensor((2, 512, 1024), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) + C_psum = T.alloc_tensor((512, 256), "float32", scope="trn.psum", layout=C_layout) for i in range(2): for k in range(2): Tx.gemm( @@ -313,9 +313,9 @@ def gemm() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - A_sbuf = T.alloc_buffer((128, 8192), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - C_psum = T.alloc_buffer((2, 128, 512), scope="trn.psum") + A_sbuf = T.alloc_tensor((128, 8192), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + C_psum = T.alloc_tensor((2, 128, 512), scope="trn.psum") for i, k, lhs_b_loop, rhs_b_loop, reduction_b_loop in T.grid(2, 2, 2, 2, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -339,17 +339,17 @@ def test_gemm_too_large_f_size(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((256, 128), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((128, 1024), "float32", scope="trn.sbuf", layout=B_layout) - C_psum = T.alloc_buffer((256, 1024), "float32", scope="trn.psum", layout=C_layout) + A_sbuf = T.alloc_tensor((256, 128), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((128, 1024), "float32", scope="trn.sbuf", layout=B_layout) + C_psum = T.alloc_tensor((256, 1024), "float32", scope="trn.psum", layout=C_layout) Tx.gemm(C_psum, A_sbuf, B_sbuf, C_psum) @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - A_sbuf = T.alloc_buffer((128, 256), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 1024), scope="trn.sbuf") - C_psum = T.alloc_buffer((4, 128, 512), scope="trn.psum") + A_sbuf = T.alloc_tensor((128, 256), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 1024), scope="trn.sbuf") + C_psum = T.alloc_tensor((4, 128, 512), scope="trn.psum") for lhs_b_loop, rhs_b_loop, reduction_b_loop in T.grid(2, 2, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -373,10 +373,10 @@ def test_gemm_sbuf_output_with_workspace(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) - C_sbuf = T.alloc_buffer((512, 256), "float32", scope="trn.sbuf", layout=C_layout) - C_psum = T.alloc_buffer((1, 128, 512), "float32", scope="trn.psum", allocated_addr=(0, 0)) + A_sbuf = T.alloc_tensor((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) + C_sbuf = T.alloc_tensor((512, 256), "float32", scope="trn.sbuf", layout=C_layout) + C_psum = T.alloc_tensor((1, 128, 512), "float32", scope="trn.psum", allocated_addr=(0, 0)) for i in range(2): for k in range(2): Tx.gemm( @@ -389,10 +389,10 @@ def gemm() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 1024), scope="trn.sbuf") - C_psum = T.alloc_buffer((1, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 1024), scope="trn.sbuf") + C_psum = T.alloc_tensor((1, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) for i, k, lhs_b_loop, rhs_b_loop in T.grid(2, 2, 2, 2): for reduction_b_loop in range(4): T.attr(0, "tensorized_nki_instruction", 1) @@ -422,9 +422,9 @@ def test_gemm_pf_mismatch_fail(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((256, 1024), "float32", scope="trn.sbuf", layout=B_layout) - C_psum = T.alloc_buffer((512, 256), "float32", scope="trn.psum", layout=C_layout) + A_sbuf = T.alloc_tensor((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((256, 1024), "float32", scope="trn.sbuf", layout=B_layout) + C_psum = T.alloc_tensor((512, 256), "float32", scope="trn.psum", layout=C_layout) for i in range(2): for k in range(2): Tx.gemm( @@ -450,9 +450,9 @@ def test_gemm_transpose_AB(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((1024, 512), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((256, 1024), "float32", scope="trn.sbuf", layout=B_layout) - C_psum = T.alloc_buffer((512, 256), "float32", scope="trn.psum", layout=C_layout) + A_sbuf = T.alloc_tensor((1024, 512), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((256, 1024), "float32", scope="trn.sbuf", layout=B_layout) + C_psum = T.alloc_tensor((512, 256), "float32", scope="trn.psum", layout=C_layout) for i in range(2): for k in range(2): Tx.gemm( @@ -467,9 +467,9 @@ def gemm() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - C_psum = T.alloc_buffer((2, 128, 512), scope="trn.psum") + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + C_psum = T.alloc_tensor((2, 128, 512), scope="trn.psum") for i, k, lhs_b_loop, rhs_b_loop, reduction_b_loop in T.grid(2, 2, 2, 1, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -494,9 +494,9 @@ def test_gemm_guard(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) - C_sbuf = T.alloc_buffer((512, 256), "float32", scope="trn.sbuf", layout=C_layout) + A_sbuf = T.alloc_tensor((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) + C_sbuf = T.alloc_tensor((512, 256), "float32", scope="trn.sbuf", layout=C_layout) for i in range(2): for j in range(2): for k in range(2): @@ -509,10 +509,10 @@ def gemm() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - acc_psum = T.alloc_buffer((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 1024), scope="trn.sbuf") + acc_psum = T.alloc_tensor((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 1024), scope="trn.sbuf") for i, j, k, lhs_b_loop, rhs_b_loop in T.grid(2, 2, 2, 2, 2): for reduction_b_loop in range(8): T.attr(0, "tensorized_nki_instruction", 1) @@ -545,9 +545,9 @@ def test_gemm_guard2(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) - C_psum = T.alloc_buffer((512, 256), "float32", scope="trn.psum", layout=C_layout) + A_sbuf = T.alloc_tensor((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) + C_psum = T.alloc_tensor((512, 256), "float32", scope="trn.psum", layout=C_layout) for j in range(4): for i in range(2): for k in range(2): @@ -560,9 +560,9 @@ def gemm() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "gemm"}) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - C_psum = T.alloc_buffer((2, 128, 512), scope="trn.psum") + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + C_psum = T.alloc_tensor((2, 128, 512), scope="trn.psum") for j, i, k, lhs_b_loop, rhs_b_loop, reduction_b_loop in T.grid(4, 2, 2, 2, 1, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_private_alloc_trn.py b/tests/python/tirx/operator/tile_primitive/trn/test_private_alloc_trn.py index d2af21015dbd..cf8f86145d2d 100644 --- a/tests/python/tirx/operator/tile_primitive/trn/test_private_alloc_trn.py +++ b/tests/python/tirx/operator/tile_primitive/trn/test_private_alloc_trn.py @@ -36,23 +36,23 @@ def test_copy_transpose(): @T.prim_func def copy() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.copy(B_sbuf, A_sbuf) @T.prim_func def expected(): T.func_attr({"global_symbol": "copy"}) T.device_entry() - identity = T.alloc_buffer((128, 128), scope="trn.sbuf") - acc_psum = T.alloc_buffer((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) + identity = T.alloc_tensor((128, 128), scope="trn.sbuf") + acc_psum = T.alloc_tensor((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for rhs_f_loop in T.serial(128, annotations={"nki_dim": "F"}): T.nki.identity(identity[p_loop, rhs_f_loop], 128) - A_sbuf = T.alloc_buffer((512, 512), scope="trn.sbuf", + A_sbuf = T.alloc_tensor((512, 512), scope="trn.sbuf", layout=T.TileLayout(T.S[(128, 2048) : (1 @ P, 1@F)])) - B_sbuf = T.alloc_buffer((512, 512), scope="trn.sbuf", + B_sbuf = T.alloc_tensor((512, 512), scope="trn.sbuf", layout=T.TileLayout(T.S[(2048, 128) : (1@F, 1@P)])) Tx.copy(B_sbuf[0:512, 0:512], A_sbuf[0:512, 0:512], workspace={"acc_psum": acc_psum, "identity": identity}) # noqa: E501 @@ -71,10 +71,10 @@ def test_normal_copy(): # fmt: off @T.prim_func - def copy(A: T.Buffer(src_shape, 'float32', layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, 'float32', layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.copy(A_sbuf, A) # fmt: on with target: @@ -95,22 +95,22 @@ def test_unary_with_bias_scale(): @T.prim_func def unary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.exp(C_sbuf, A_sbuf, bias=bias, scale=scale) @T.prim_func def expected(): T.func_attr({"global_symbol": "unary"}) T.device_entry() - const_bias = T.alloc_buffer((128, 512), scope="trn.sbuf") + const_bias = T.alloc_tensor((128, 512), scope="trn.sbuf") with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for f_loop in T.serial(512, annotations={"nki_dim": "F"}): T.nki.memset(const_bias[p_loop, f_loop], T.float32(1.0)) - A_sbuf = T.alloc_buffer((512, 1024), scope="trn.sbuf", + A_sbuf = T.alloc_tensor((512, 1024), scope="trn.sbuf", layout=T.TileLayout(T.S[(128, 4096) : (1@P, 1@F)])) - C_sbuf = T.alloc_buffer((512, 1024), scope="trn.sbuf", + C_sbuf = T.alloc_tensor((512, 1024), scope="trn.sbuf", layout=T.TileLayout(T.S[(128, 4096) : (1@P, 1@F)])) Tx.exp(C_sbuf[0:512, 0:1024], A_sbuf[0:512, 0:1024], T.float32(1.0), T.float32(2.0), workspace={"const_bias": const_bias}) # noqa: E501 # fmt: on @@ -130,18 +130,18 @@ def test_reduction_two_stage(): @T.prim_func def reduction(): T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.sum(B_sbuf, A_sbuf, axes=(1, 3)) @T.prim_func def expected(): T.func_attr({"global_symbol": "reduction"}) T.device_entry() - partial_reduce = T.alloc_buffer((128, 32), scope="trn.sbuf") - A_sbuf = T.alloc_buffer((128, 32, 4, 32), scope="trn.sbuf", + partial_reduce = T.alloc_tensor((128, 32), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 32, 4, 32), scope="trn.sbuf", layout=T.TileLayout(T.S[(128, 32 * 32 * 4) : (1@P, 1@F)])) - B_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf", + B_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf", layout=T.TileLayout(T.S[(128, 4) : (1@P, 1@F)])) Tx.sum(B_sbuf[0:128, 0:4], A_sbuf[0:128, 0:32, 0:4, 0:32], [1, 3], False, workspace={"partial_reduce": partial_reduce}) # noqa: E501 @@ -162,9 +162,9 @@ def test_gemm(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) - C_sbuf = T.alloc_buffer((512, 256), "float32", scope="trn.sbuf", layout=C_layout) + A_sbuf = T.alloc_tensor((512, 1024), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((1024, 256), "float32", scope="trn.sbuf", layout=B_layout) + C_sbuf = T.alloc_tensor((512, 256), "float32", scope="trn.sbuf", layout=C_layout) for i in range(2): for k in range(2): Tx.gemm( @@ -177,12 +177,12 @@ def gemm() -> None: def expected(): T.func_attr({"global_symbol": "gemm"}) T.device_entry() - acc_psum = T.alloc_buffer((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) - A_sbuf = T.alloc_buffer((512, 1024), scope="trn.sbuf", + acc_psum = T.alloc_tensor((8, 128, 512), scope="trn.psum", allocated_addr=[0, 0]) + A_sbuf = T.alloc_tensor((512, 1024), scope="trn.sbuf", layout=T.TileLayout(T.S[(4, 128, 8, 128) : (1024@F, 1@F, 1@F, 1@P)])) # noqa: E501 - B_sbuf = T.alloc_buffer((1024, 256), scope="trn.sbuf", + B_sbuf = T.alloc_tensor((1024, 256), scope="trn.sbuf", layout=T.TileLayout(T.S[(8, 128, 2, 128) : (256@F, 1@P, 128@F, 1@F)])) # noqa: E501 - C_sbuf = T.alloc_buffer((512, 256), scope="trn.sbuf", + C_sbuf = T.alloc_tensor((512, 256), scope="trn.sbuf", layout=T.TileLayout(T.S[(4, 128, 2, 128) : (256@F, 1@F, 128@F, 1@P)])) # noqa: E501 for i, k in T.grid(2, 2): Tx.gemm(C_sbuf[256 * i:256 * i + 256, 0:256], A_sbuf[256 * i:256 * i + 256, 512 * k:512 * k + 512], B_sbuf[512 * k:512 * k + 512, 0:256], C_sbuf[256 * i:256 * i + 256, 0:256], False, False, T.float32(1.0), T.float32(0.0), workspace={"acc_psum": acc_psum}) # noqa: E501 @@ -205,21 +205,21 @@ def test_binary_reduce_two_stage(): @T.prim_func def tensor_scalar_reduce() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) - B_sbuf = T.alloc_buffer(dst1_shape, "float32", scope="trn.sbuf", layout=dst1_layout) - C_sbuf = T.alloc_buffer(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 + A_sbuf = T.alloc_tensor(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) + B_sbuf = T.alloc_tensor(dst1_shape, "float32", scope="trn.sbuf", layout=dst1_layout) + C_sbuf = T.alloc_tensor(reduce_dst_shape, "float32", scope="trn.sbuf", layout=reduce_dst_layout) # noqa: E501 Tx.binary_reduce(B_sbuf, C_sbuf, A_sbuf, 1.0, "add", "sum", reduce_axes=(1, 2)) @T.prim_func def expected(): T.func_attr({"global_symbol": "tensor_scalar_reduce"}) T.device_entry() - partial_reduce = T.alloc_buffer((128, 4), scope="trn.sbuf") - A_sbuf = T.alloc_buffer((512, 1024, 4), scope="trn.sbuf", + partial_reduce = T.alloc_tensor((128, 4), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((512, 1024, 4), scope="trn.sbuf", layout=T.TileLayout(T.S[(128, 4096, 4) : (1 @ P, 1 @ F, 4096 @ F)])) - B_sbuf = T.alloc_buffer((512, 1024, 4), scope="trn.sbuf", + B_sbuf = T.alloc_tensor((512, 1024, 4), scope="trn.sbuf", layout=T.TileLayout(T.S[(128, 4096, 4) : (1 @ P, 1 @ F, 4096 @ F)])) - C_sbuf = T.alloc_buffer((512,), scope="trn.sbuf", + C_sbuf = T.alloc_tensor((512,), scope="trn.sbuf", layout=T.TileLayout(T.S[(128, 4) : (1 @ P, 1 @ F)])) Tx.binary_reduce(B_sbuf[0:512, 0:1024, 0:4], C_sbuf[0:512], A_sbuf[0:512, 0:1024, 0:4], T.float32(1.0), "add", "sum", [1, 2], workspace={"partial_reduce": partial_reduce}) # noqa: E501 # fmt: on @@ -241,9 +241,9 @@ def test_activation_reduce_two_stage(): @T.prim_func def activation_reduce(): T.device_entry() - A = T.alloc_buffer(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) - B = T.alloc_buffer(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) - C = T.alloc_buffer(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) + A = T.alloc_tensor(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) + B = T.alloc_tensor(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) + C = T.alloc_tensor(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) for i in range(2): Tx.unary_reduce(B, C, A[i*16:i*16+16], "sqrt", "sum", reduce_axes=(0,1)) @@ -251,17 +251,17 @@ def activation_reduce(): def expected(): T.func_attr({"global_symbol": "activation_reduce"}) T.device_entry() - partial_reduce = T.alloc_buffer((128, 8), scope="trn.sbuf") - const_bias = T.alloc_buffer((128, 1024), scope="trn.sbuf") + partial_reduce = T.alloc_tensor((128, 8), scope="trn.sbuf") + const_bias = T.alloc_tensor((128, 1024), scope="trn.sbuf") with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for f_loop in T.serial(1024, annotations={"nki_dim": "F"}): T.nki.memset(const_bias[p_loop, f_loop], T.float32(0.0)) - A = T.alloc_buffer((32, 512, 128), scope="trn.sbuf", + A = T.alloc_tensor((32, 512, 128), scope="trn.sbuf", layout=T.TileLayout(T.S[(16 * 1024, 128) : (1@F, 1@P)])) - B = T.alloc_buffer((16, 512, 128), scope="trn.sbuf", + B = T.alloc_tensor((16, 512, 128), scope="trn.sbuf", layout=T.TileLayout(T.S[(2, 4, 1024, 128) : (1024@F, 2048@F, 1@F, 1@P)])) - C = T.alloc_buffer((1, 128), scope="trn.sbuf", + C = T.alloc_tensor((1, 128), scope="trn.sbuf", layout=T.TileLayout(T.S[(1, 128) : (1@F, 1@P)])) for i in range(2): Tx.unary_reduce(B[0:16, 0:512, 0:128], C[0, 0:128], A[i * 16:i * 16 + 16, 0:512, 0:128], "sqrt", "sum", None, None, [0, 1], workspace={"const_bias": const_bias, "partial_reduce": partial_reduce}) # noqa: E501 @@ -284,10 +284,10 @@ def test_partial_workspace_specify(): @T.prim_func def activation_reduce(): T.device_entry() - partial_reduce = T.alloc_buffer((128, 16), scope="trn.sbuf") - A = T.alloc_buffer(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) - B = T.alloc_buffer(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) - C = T.alloc_buffer(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) + partial_reduce = T.alloc_tensor((128, 16), scope="trn.sbuf") + A = T.alloc_tensor(A_shape, dtype="float32", scope="trn.sbuf", layout=A_layout) + B = T.alloc_tensor(B_shape, dtype="float32", scope="trn.sbuf", layout=B_layout) + C = T.alloc_tensor(C_shape, dtype="float32", scope="trn.sbuf", layout=C_layout) for i in range(2): Tx.unary_reduce(B, C, A[i*16:i*16+16], "sqrt", "sum", reduce_axes=(0,1), workspace={"partial_reduce": partial_reduce}) # noqa: E501 @@ -295,17 +295,17 @@ def activation_reduce(): def expected(): T.func_attr({"global_symbol": "activation_reduce"}) T.device_entry() - const_bias = T.alloc_buffer((128, 1024), scope="trn.sbuf") + const_bias = T.alloc_tensor((128, 1024), scope="trn.sbuf") with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for f_loop in T.serial(1024, annotations={"nki_dim": "F"}): T.nki.memset(const_bias[p_loop, f_loop], T.float32(0.0)) - partial_reduce = T.alloc_buffer((128, 16), scope="trn.sbuf") - A = T.alloc_buffer((32, 512, 128), scope="trn.sbuf", + partial_reduce = T.alloc_tensor((128, 16), scope="trn.sbuf") + A = T.alloc_tensor((32, 512, 128), scope="trn.sbuf", layout=T.TileLayout(T.S[(16 * 1024, 128) : (1@F, 1@P)])) - B = T.alloc_buffer((16, 512, 128), scope="trn.sbuf", + B = T.alloc_tensor((16, 512, 128), scope="trn.sbuf", layout=T.TileLayout(T.S[(2, 4, 1024, 128) : (1024@F, 2048@F, 1@F, 1@P)])) - C = T.alloc_buffer((1, 128), scope="trn.sbuf", + C = T.alloc_tensor((1, 128), scope="trn.sbuf", layout=T.TileLayout(T.S[(1, 128) : (1@F, 1@P)])) for i in range(2): Tx.unary_reduce(B[0:16, 0:512, 0:128], C[0, 0:128], A[i * 16:i * 16 + 16, 0:512, 0:128], "sqrt", "sum", None, None, [0, 1], workspace={"const_bias": const_bias, "partial_reduce": partial_reduce}) # noqa: E501 @@ -327,8 +327,8 @@ def test_workspace_reuse(): @T.prim_func def unary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.exp(C_sbuf, A_sbuf, bias=0.0, scale=scale, max_inst_size=1024) Tx.exp(C_sbuf, C_sbuf) @@ -336,14 +336,14 @@ def unary() -> None: def expected(): T.func_attr({"global_symbol": "unary"}) T.device_entry() - const_bias = T.alloc_buffer((128, 1024), scope="trn.sbuf") + const_bias = T.alloc_tensor((128, 1024), scope="trn.sbuf") with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for f_loop in T.serial(1024, annotations={"nki_dim": "F"}): T.nki.memset(const_bias[p_loop, f_loop], T.float32(0.0)) - A_sbuf = T.alloc_buffer((512, 1024), scope="trn.sbuf", + A_sbuf = T.alloc_tensor((512, 1024), scope="trn.sbuf", layout=T.TileLayout(T.S[(128, 4096) : (1 @ P, 1 @ F)])) - C_sbuf = T.alloc_buffer((512, 1024), scope="trn.sbuf", + C_sbuf = T.alloc_tensor((512, 1024), scope="trn.sbuf", layout=T.TileLayout(T.S[(128, 4096) : (1 @ P, 1 @ F)])) Tx.exp(C_sbuf[0:512, 0:1024], A_sbuf[0:512, 0:1024], T.float32(0.0), T.float32(2.0), workspace={"const_bias": const_bias}, max_inst_size=1024) # noqa: E501 Tx.exp(C_sbuf[0:512, 0:1024], C_sbuf[0:512, 0:1024], None, None, workspace={"const_bias": const_bias}) # noqa: E501 @@ -366,9 +366,9 @@ def test_no_rewrite_with_existing_workspace(): @T.prim_func def reduction(): T.device_entry() - intermediate_buffer = T.alloc_buffer((128, 64), scope="trn.sbuf") - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + intermediate_buffer = T.alloc_tensor((128, 64), scope="trn.sbuf") + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.sum(B_sbuf, A_sbuf, axes=(1, 3), workspace={"partial_reduce": intermediate_buffer}) # fmt: on with target: @@ -387,9 +387,9 @@ def test_no_rewrite_with_psum_output(): @T.prim_func def gemm() -> None: T.device_entry() - A_sbuf = T.alloc_buffer((128, 128), "float32", scope="trn.sbuf", layout=A_layout) - B_sbuf = T.alloc_buffer((128, 128), "float32", scope="trn.sbuf", layout=B_layout) - C_psum = T.alloc_buffer((128, 128), "float32", scope="trn.psum", layout=C_layout) + A_sbuf = T.alloc_tensor((128, 128), "float32", scope="trn.sbuf", layout=A_layout) + B_sbuf = T.alloc_tensor((128, 128), "float32", scope="trn.sbuf", layout=B_layout) + C_psum = T.alloc_tensor((128, 128), "float32", scope="trn.psum", layout=C_layout) Tx.gemm(C_psum, A_sbuf, B_sbuf, C_psum) # fmt: on with target: diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_reduction_trn.py b/tests/python/tirx/operator/tile_primitive/trn/test_reduction_trn.py index eb80931cd78b..13cf4e74302d 100644 --- a/tests/python/tirx/operator/tile_primitive/trn/test_reduction_trn.py +++ b/tests/python/tirx/operator/tile_primitive/trn/test_reduction_trn.py @@ -67,15 +67,15 @@ def test_simple_reduction(op_type): @T.prim_func def reduction() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) tx_func(B_sbuf, A_sbuf, axes=-1) @T.prim_func def expected(): T.func_attr({"global_symbol": "reduction"}) - A_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 1), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 1), scope="trn.sbuf") for b_loop in range(1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -99,15 +99,15 @@ def test_reduction_with_multiple_axes(): @T.prim_func def reduction(): T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.sum(B_sbuf, A_sbuf, axes=(1, 2), max_inst_size=2048) @T.prim_func def expected(): T.func_attr({"global_symbol": "reduction"}) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 1), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 1), scope="trn.sbuf") for b_loop in range(1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -131,16 +131,16 @@ def test_reduction_in_loop(): @T.prim_func def reduction(): T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(4): Tx.sum(B_sbuf[:, i], A_sbuf[:, :, i], axes=-2) @T.prim_func def expected(): T.func_attr({"global_symbol": "reduction"}) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") for i, b_loop in T.grid(4, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -163,16 +163,16 @@ def test_reduction_two_stage(): @T.prim_func def reduction(): T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.sum(B_sbuf, A_sbuf, axes=(1, 3)) @T.prim_func def expected(): T.func_attr({"global_symbol": "reduction"}) - intermediate_buffer = T.alloc_buffer((128, 32), scope="trn.sbuf") - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") + intermediate_buffer = T.alloc_tensor((128, 32), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") for b_loop in range(4): for reduction_b_loop in range(32): T.attr(0, "tensorized_nki_instruction", 1) @@ -202,8 +202,8 @@ def test_reduction_with_guard(): @T.prim_func def reduction() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(4): for j in range(4): Tx.sum(B_sbuf[0: (i+1) * 128, 0], A_sbuf[0: (i+1) * 128, 0: (j+1) * 256], max_inst_size=512) # noqa: E501 @@ -211,9 +211,9 @@ def reduction() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "reduction"}) - intermediate_buffer = T.alloc_buffer((128, 2), scope="trn.sbuf") - A_sbuf = T.alloc_buffer((128, 8192), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") + intermediate_buffer = T.alloc_tensor((128, 2), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 8192), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") for i, j in T.grid(4, 4): for b_loop in range(4): for reduction_b_loop in range(2): @@ -249,17 +249,17 @@ def test_reduction_two_stage_workspace(): @T.prim_func def reduction(): T.device_entry() - intermediate_buffer = T.alloc_buffer((128, 64), scope="trn.sbuf") - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + intermediate_buffer = T.alloc_tensor((128, 64), scope="trn.sbuf") + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.sum(B_sbuf, A_sbuf, axes=(1, 3), workspace={"partial_reduce": intermediate_buffer}) @T.prim_func def expected(): T.func_attr({"global_symbol": "reduction"}) - intermediate_buffer = T.alloc_buffer((128, 64), scope="trn.sbuf") - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") + intermediate_buffer = T.alloc_tensor((128, 64), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") for b_loop in range(4): for reduction_b_loop in range(32): T.attr(0, "tensorized_nki_instruction", 1) diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_select_trn.py b/tests/python/tirx/operator/tile_primitive/trn/test_select_trn.py index 1d2f068fab0b..19d8bc6d70bb 100644 --- a/tests/python/tirx/operator/tile_primitive/trn/test_select_trn.py +++ b/tests/python/tirx/operator/tile_primitive/trn/test_select_trn.py @@ -58,15 +58,15 @@ def test_select(): @T.prim_func def select() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.select(B_sbuf, A_sbuf, 0.0, lambda i, j: i < j) @T.prim_func def expected(): T.func_attr({"global_symbol": "select"}) - A_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") for b_loop in T.serial(0, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -91,16 +91,16 @@ def test_select_in_loop(): @T.prim_func def select() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(2): Tx.select(B_sbuf, A_sbuf[i*16, :, :], 0.0, lambda a, b: (i+1)* a < b) @T.prim_func def expected(): T.func_attr({"global_symbol": "select"}) - A_sbuf = T.alloc_buffer((128, 16384), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 16384), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") for i, b_loop in T.grid(2, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -125,15 +125,15 @@ def test_select_expr_affine(): @T.prim_func def select() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.select(B_sbuf, A_sbuf, 0.0, lambda i, j: i < j) @T.prim_func def expected(): T.func_attr({"global_symbol": "select"}) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") for b_loop in T.serial(0, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -157,8 +157,8 @@ def test_select_with_guard(): @T.prim_func def select() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(4): for j in range(4): Tx.select(B_sbuf[0: (i+1) * 128, 0: (j+1) * 128], A_sbuf[0: (i+1) * 128, 0: (j+1) * 128], 0.0, lambda a, b: a < b) # noqa: E501 @@ -166,8 +166,8 @@ def select() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "select"}) - A_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") for i, j, b_loop in T.grid(4, 4, 4): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py b/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py index 0773ea8132ba..c61c2e07caec 100644 --- a/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py +++ b/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py @@ -63,8 +63,8 @@ def test_simple_unary(op_type): @T.prim_func def unary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) if T.constexpr(op_type == "memset"): tx_func(B_sbuf, T.float32(0.0)) else: @@ -73,8 +73,8 @@ def unary() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "unary"}) - A_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 512), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 512), scope="trn.sbuf") for b_loop in T.serial(0, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -105,8 +105,8 @@ def test_unary_in_a_loop(op_type): @T.prim_func def unary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) A_sbuf_view = A_sbuf.view(128, 8, 512) B_sbuf_view = B_sbuf.view(128, 4, 512) for i in range(4): @@ -118,10 +118,10 @@ def unary() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "unary"}) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 2048), scope="trn.sbuf") - A_sbuf_view = T.decl_buffer((128, 4096), data=A_sbuf.data, scope="trn.sbuf", layout=None) - B_sbuf_view = T.decl_buffer((128, 2048), data=B_sbuf.data, scope="trn.sbuf", layout=None) + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 2048), scope="trn.sbuf") + A_sbuf_view = T.decl_tensor((128, 4096), data=A_sbuf.data, scope="trn.sbuf", layout=None) + B_sbuf_view = T.decl_tensor((128, 2048), data=B_sbuf.data, scope="trn.sbuf", layout=None) for i, b_loop in T.grid(4, 1): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -145,13 +145,13 @@ def test_unary_complex1(): @T.prim_func def unary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.memset(A_sbuf, T.float32(0.0)) @T.prim_func def expected(): T.func_attr({"global_symbol": "unary"}) - A_sbuf = T.alloc_buffer((128, 8192), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 8192), scope="trn.sbuf") for b_loop in T.serial(0, 16): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -179,17 +179,17 @@ def test_unary_with_bias_scale(op_type): @T.prim_func def unary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(bias_shape, "float32", scope="trn.sbuf", layout=bias_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(bias_shape, "float32", scope="trn.sbuf", layout=bias_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) tx_func(C_sbuf, A_sbuf, bias=B_sbuf, scale=scale) @T.prim_func def expected(): T.func_attr({"global_symbol": "unary"}) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") for b_loop in T.serial(0, 8): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): @@ -216,20 +216,20 @@ def test_unary_with_bias_scale_2(op_type): @T.prim_func def unary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) tx_func(C_sbuf, A_sbuf, bias=bias, scale=scale) @T.prim_func def expected(): T.func_attr({"global_symbol": "unary"}) - const_bias = T.alloc_buffer((128, 512), scope="trn.sbuf") + const_bias = T.alloc_tensor((128, 512), scope="trn.sbuf") with T.attr(0, "tensorized_nki_instruction", 1): for p_loop in T.serial(128, annotations={"nki_dim": "P"}): for f_loop in T.serial(512, annotations={"nki_dim": "F"}): T.nki.memset(const_bias[p_loop, f_loop], T.float32(1.0)) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") for b_loop in T.serial(0, 8): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(128, annotations={"nki_dim": "P"}): @@ -256,9 +256,9 @@ def test_unary_with_guard(): @T.prim_func def unary() -> None: T.device_entry() - A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) - B_sbuf = T.alloc_buffer(bias_shape, "float32", scope="trn.sbuf", layout=bias_layout) - C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(src_shape, "float32", scope="trn.sbuf", layout=src_layout) + B_sbuf = T.alloc_tensor(bias_shape, "float32", scope="trn.sbuf", layout=bias_layout) + C_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) for i in range(4): for j in range(4): Tx.sqrt(C_sbuf[0: (i+1) * 128, 0: (j+1)*256], A_sbuf[0: (i+1) * 128, 0: (j+1)*256], bias=B_sbuf[0: (i+1) * 128, 0], scale=scale) # noqa: E501 @@ -266,9 +266,9 @@ def unary() -> None: @T.prim_func def expected(): T.func_attr({"global_symbol": "unary"}) - A_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") - B_sbuf = T.alloc_buffer((128, 4), scope="trn.sbuf") - C_sbuf = T.alloc_buffer((128, 4096), scope="trn.sbuf") + A_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") + B_sbuf = T.alloc_tensor((128, 4), scope="trn.sbuf") + C_sbuf = T.alloc_tensor((128, 4096), scope="trn.sbuf") for i, j, b_loop in T.grid(4, 4, 8): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): diff --git a/tests/python/tirx/script/test_tirx_script_basic_usage.py b/tests/python/tirx/script/test_tirx_script_basic_usage.py index fccdf5630652..68ef2354d39f 100644 --- a/tests/python/tirx/script/test_tirx_script_basic_usage.py +++ b/tests/python/tirx/script/test_tirx_script_basic_usage.py @@ -56,16 +56,16 @@ def main(): def test_tir_buffer_annotation(): - buffer_0 = T.Buffer((128, 128), "float32") + buffer_0 = T.Tensor((128, 128), "float32") assert ( - isinstance(buffer_0, tirx.BufferType) + isinstance(buffer_0, tirx.TensorType) and list(buffer_0.shape) == [128, 128] and buffer_0.dtype == ir.PrimType("float32") ) - buffer_1 = T.Buffer((64, 64, 64), "int32") + buffer_1 = T.Tensor((64, 64, 64), "int32") assert ( - isinstance(buffer_1, tirx.BufferType) + isinstance(buffer_1, tirx.TensorType) and list(buffer_1.shape) == [64, 64, 64] and buffer_1.dtype == ir.PrimType("int32") ) @@ -260,7 +260,7 @@ def test_pointer_expression_assignment_uses_bind(): @T.prim_func def func() -> None: T.device_entry() - buf = T.alloc_buffer((4,), "uint32", scope="shared") + buf = T.alloc_tensor((4,), "uint32", scope="shared") ptr = buf.ptr_to([1]) T.evaluate(T.reinterpret("uint64", ptr)) # fmt: on @@ -290,7 +290,7 @@ def test_pointer_expression_rebinding_creates_distinct_native_bindings(): @T.prim_func def func() -> None: T.device_entry() - buf = T.alloc_buffer((4,), "uint32", scope="shared") + buf = T.alloc_tensor((4,), "uint32", scope="shared") ptr = buf.ptr_to([0]) ptr = buf.ptr_to([1]) T.evaluate(T.reinterpret("uint64", ptr)) @@ -318,9 +318,9 @@ def test_pointer_expression_assignment_can_shadow_extra_var(): @T.prim_func def func() -> None: T.device_entry() - buf = T.alloc_buffer((4,), "uint32", scope="shared") + buf = T.alloc_tensor((4,), "uint32", scope="shared") ptr = buf.ptr_to([1]) - view = T.decl_buffer((3,), "uint32", data=ptr, scope="shared") + view = T.decl_tensor((3,), "uint32", data=ptr, scope="shared") view[0] = T.uint32(0) """ func = tvm.script.from_source(source, extra_vars={"T": T, "ptr": object()}) @@ -341,7 +341,7 @@ def test_roundtrip_unary_inplace(): # fmt: off @T.prim_func - def test(A: T.Buffer((128,), "float32", scope="global")) -> None: + def test(A: T.Tensor((128,), "float32", scope="global")) -> None: T.device_entry() cta_id = T.cta_id([1]) warp_id = T.warp_id([1]) @@ -369,8 +369,8 @@ def test_roundtrip_unary_different_dst_src(): # fmt: off @T.prim_func def test( - A: T.Buffer((128,), "float32", scope="global"), - B: T.Buffer((128,), "float32", scope="global"), + A: T.Tensor((128,), "float32", scope="global"), + B: T.Tensor((128,), "float32", scope="global"), ) -> None: T.device_entry() cta_id = T.cta_id([1]) @@ -572,7 +572,7 @@ class Rejected: def test_roundtrip_break_for(): # fmt: off @T.prim_func - def test(A: T.Buffer((10,), 'int32')): + def test(A: T.Tensor((10,), 'int32')): T.device_entry() for i in T.serial(10): @@ -588,10 +588,10 @@ def test(A: T.Buffer((10,), 'int32')): def test_roundtrip_break_while(): # fmt: off @T.prim_func - def test(A: T.Buffer((10,), 'int32')): + def test(A: T.Tensor((10,), 'int32')): T.device_entry() - i = T.alloc_buffer((1,), "int32", scope="local") + i = T.alloc_tensor((1,), "int32", scope="local") i[0] = 0 while i[0] < 10: A[i[0]] = i[0] * 2 @@ -607,10 +607,10 @@ def test(A: T.Buffer((10,), 'int32')): def test_roundtrip_break_nested(): # fmt: off @T.prim_func - def test(A: T.Buffer((9,), 'int32')): + def test(A: T.Tensor((9,), 'int32')): T.device_entry() - idx = T.alloc_buffer((1,), "int32", scope="local") + idx = T.alloc_tensor((1,), "int32", scope="local") idx[0] = 0 for i in T.serial(3): for j in T.serial(3): @@ -627,7 +627,7 @@ def test(A: T.Buffer((9,), 'int32')): def test_roundtrip_continue_for(): # fmt: off @T.prim_func - def test(A: T.Buffer((10,), 'int32')): + def test(A: T.Tensor((10,), 'int32')): T.device_entry() for i in T.serial(10): @@ -643,10 +643,10 @@ def test(A: T.Buffer((10,), 'int32')): def test_roundtrip_continue_while(): # fmt: off @T.prim_func - def test(A: T.Buffer((10,), 'int32')): + def test(A: T.Tensor((10,), 'int32')): T.device_entry() - i = T.alloc_buffer((1,), "int32", scope="local") + i = T.alloc_tensor((1,), "int32", scope="local") i[0] = 0 while i[0] < 10: if (i[0] % 2) == 1: @@ -663,10 +663,10 @@ def test(A: T.Buffer((10,), 'int32')): def test_roundtrip_continue_nested(): # fmt: off @T.prim_func - def test(A: T.Buffer((9,), 'int32')): + def test(A: T.Tensor((9,), 'int32')): T.device_entry() - idx = T.alloc_buffer((1,), dtype="int32", scope="local") + idx = T.alloc_tensor((1,), dtype="int32", scope="local") idx[0] = 0 for i in T.serial(3): for j in T.serial(3): @@ -683,7 +683,7 @@ def test(A: T.Buffer((9,), 'int32')): def test_roundtrip_break_and_continue(): # fmt: off @T.prim_func - def test(A: T.Buffer((10,), 'int32')): + def test(A: T.Tensor((10,), 'int32')): T.device_entry() for i in T.serial(10): @@ -701,7 +701,7 @@ def test(A: T.Buffer((10,), 'int32')): def test_roundtrip_unreachable_after_break(): # fmt: off @T.prim_func - def test(A: T.Buffer((5,), 'int32')): + def test(A: T.Tensor((5,), 'int32')): T.device_entry() for i in T.serial(5): @@ -720,7 +720,7 @@ def test_roundtrip_serial_unroll_false(): # fmt: off @T.prim_func - def test(A: T.Buffer((128,), 'float32', scope='global')) -> None: + def test(A: T.Tensor((128,), 'float32', scope='global')) -> None: T.device_entry() cta_id = T.cta_id([1]) @@ -742,7 +742,7 @@ def test_roundtrip_serial_unroll_true(): # fmt: off @T.prim_func - def test(A: T.Buffer((128,), 'float32', scope='global')) -> None: + def test(A: T.Tensor((128,), 'float32', scope='global')) -> None: T.device_entry() cta_id = T.cta_id([1]) @@ -764,7 +764,7 @@ def test_roundtrip_serial_unroll_count(): # fmt: off @T.prim_func - def test(A: T.Buffer((128,), 'float32', scope='global')) -> None: + def test(A: T.Tensor((128,), 'float32', scope='global')) -> None: T.device_entry() cta_id = T.cta_id([1]) @@ -786,7 +786,7 @@ def test_roundtrip_serial_unroll_false_with_other_annotations(): # fmt: off @T.prim_func - def test(A: T.Buffer((128,), 'float32', scope='global')) -> None: + def test(A: T.Tensor((128,), 'float32', scope='global')) -> None: T.device_entry() cta_id = T.cta_id([1]) @@ -805,7 +805,7 @@ def test(A: T.Buffer((128,), 'float32', scope='global')) -> None: def test_loop_var_dtype_uint32(): # fmt: off @T.prim_func - def func(A: T.Buffer((128,), 'float32')): + def func(A: T.Tensor((128,), 'float32')): for i in T.serial(128, dtype="uint32"): A[i] = T.float32(1) @@ -827,7 +827,7 @@ def _assert_roundtrip(func): def test_loop_var_dtype_uint32_with_step(): # fmt: off @T.prim_func - def func(A: T.Buffer((128,), 'float32')): + def func(A: T.Tensor((128,), 'float32')): for i in T.serial(4, 128, step=2, dtype="uint32"): A[i] = T.float32(1) @@ -845,7 +845,7 @@ def func(A: T.Buffer((128,), 'float32')): def test_loop_var_dtype_uint32_all_for_kinds(for_kind): # fmt: off @T.prim_func - def func(A: T.Buffer((4,), 'float32')): + def func(A: T.Tensor((4,), 'float32')): for i in getattr(T, for_kind)(4, dtype="uint32"): A[i] = T.float32(1) @@ -858,7 +858,7 @@ def func(A: T.Buffer((4,), 'float32')): def test_grid_loop_var_dtype_uint32(): # fmt: off @T.prim_func - def func(A: T.Buffer((8, 16), 'float32')): + def func(A: T.Tensor((8, 16), 'float32')): for i, j in T.grid(8, 16, dtype="uint32"): A[i, j] = T.float32(1) @@ -873,7 +873,7 @@ def func(A: T.Buffer((8, 16), 'float32')): def test_loop_var_dtype_defaults_to_int32(): # fmt: off @T.prim_func - def func(A: T.Buffer((128,), 'float32')): + def func(A: T.Tensor((128,), 'float32')): for i in range(128): A[i] = T.float32(1) @@ -888,7 +888,7 @@ def test_loop_var_dtype_inferred_from_unsigned_extent(): # fmt: off @T.prim_func - def func(A: T.Buffer((128,), 'float32'), n: T.uint32): + def func(A: T.Tensor((128,), 'float32'), n: T.uint32): for i in range(n): A[i] = T.float32(1) @@ -903,7 +903,7 @@ def test_loop_var_dtype_casts_mismatched_bound(): # fmt: off @T.prim_func - def func(A: T.Buffer((128,), 'float32'), n: T.int32): + def func(A: T.Tensor((128,), 'float32'), n: T.int32): for i in T.serial(n, dtype="uint32"): A[i] = T.float32(1) diff --git a/tests/python/tirx/script/test_tirx_script_buffers.py b/tests/python/tirx/script/test_tirx_script_buffers.py index 1692b06e6b4a..6640eceebd1b 100644 --- a/tests/python/tirx/script/test_tirx_script_buffers.py +++ b/tests/python/tirx/script/test_tirx_script_buffers.py @@ -51,18 +51,18 @@ def get_layout5(): # fmt: off @T.prim_func - def test(_: T.Buffer((64,), 'float32', scope='global')) -> None: + def test(_: T.Tensor((64,), 'float32', scope='global')) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) warp_id = T.warp_id([1]) lane_id = T.lane_id([32]) - C = T.alloc_buffer([128, 128], dtype="float16", scope="shared", layout=get_layout3()) - D = T.alloc_buffer([128, 32], dtype="float16", scope="shared", layout=get_layout4()) - A_warp = T.alloc_buffer([64, 64], dtype="float16", scope="shared", layout=get_layout1()) - B_warp = T.alloc_buffer([64, 64], dtype="float16", scope="shared", layout=get_layout2()) + C = T.alloc_tensor([128, 128], dtype="float16", scope="shared", layout=get_layout3()) + D = T.alloc_tensor([128, 32], dtype="float16", scope="shared", layout=get_layout4()) + A_warp = T.alloc_tensor([64, 64], dtype="float16", scope="shared", layout=get_layout1()) + B_warp = T.alloc_tensor([64, 64], dtype="float16", scope="shared", layout=get_layout2()) - E = T.alloc_buffer([64, 256], dtype="float16", scope="shared", layout=get_layout5()) + E = T.alloc_tensor([64, 256], dtype="float16", scope="shared", layout=get_layout5()) T.evaluate(A_warp[0, 0] + B_warp[0, 0] + C[0, 0] + D[0, 0] + E[0, 0]) # fmt: on @@ -97,10 +97,10 @@ def get_full(): @T.prim_func def test() -> None: T.device_entry() - A = T.alloc_buffer([8], dtype="float16", scope="shared", layout=get_shard_replica()) - B = T.alloc_buffer([8], dtype="float16", scope="shared", layout=get_shard_offset_single()) - C = T.alloc_buffer([8], dtype="float16", scope="shared", layout=get_shard_offset_multi()) - D = T.alloc_buffer([32], dtype="float16", scope="shared", layout=get_full()) + A = T.alloc_tensor([8], dtype="float16", scope="shared", layout=get_shard_replica()) + B = T.alloc_tensor([8], dtype="float16", scope="shared", layout=get_shard_offset_single()) + C = T.alloc_tensor([8], dtype="float16", scope="shared", layout=get_shard_offset_multi()) + D = T.alloc_tensor([32], dtype="float16", scope="shared", layout=get_full()) T.evaluate(A[0] + B[0] + C[0] + D[0]) # fmt: on @@ -114,7 +114,7 @@ def test_roundtrip_buffer_view_get1(): @T.prim_func def test() -> None: T.device_entry() - A = T.alloc_buffer([2], dtype="float16", scope="local") + A = T.alloc_tensor([2], dtype="float16", scope="local") A_layout = T.TileLayout(T.S[(1, 2) : (2, 1)]) A_warp_layout = A_layout.tile(L_LANE, (8, 4), (1, 2)) A_warp = A.view(8, 8, layout=A_warp_layout) @@ -133,14 +133,14 @@ def test() -> None: def test_roundtrip_buffer_view_get2(): # fmt: off @T.prim_func - def test(out: T.Buffer(2, 'float32', scope='global')) -> None: + def test(out: T.Tensor(2, 'float32', scope='global')) -> None: T.device_entry() bx, by, bz = T.cta_id([32, 32, 1]) tx, ty, tz = T.thread_id([16, 8, 1]) warp_id = T.warp_id([4]) lane_id = T.lane_id([32]) - A = T.alloc_buffer([2,], dtype="float16", scope="local") + A = T.alloc_tensor([2,], dtype="float16", scope="local") A_layout = T.TileLayout(T.S[(1, 2) : (2, 1)]) B_layout = A_layout.tile(L_LANE, (8, 4), (1, 2)) B = A.view(8, 8, layout=B_layout) @@ -157,7 +157,7 @@ def test_roundtrip_buffer_view_get3(): @T.prim_func def test() -> None: T.device_entry() - A = T.alloc_buffer([8, 8], dtype="float32", scope="local") + A = T.alloc_tensor([8, 8], dtype="float32", scope="local") A_f16 = A.view("float16") A_f64 = A.view("float64") A_f16[0, 0] = T.float16(0) @@ -175,7 +175,7 @@ def test_roundtrip_allocated_addr(): @T.prim_func def test(): T.device_entry() - A = T.alloc_buffer([10], "float32", scope="trn.sbuf", allocated_addr=1024) + A = T.alloc_tensor([10], "float32", scope="trn.sbuf", allocated_addr=1024) for i in T.serial(2): Tx.memset(A[i*5:i*5+5], T.float32(0.0)) @@ -188,7 +188,7 @@ def test(): def test_roundtrip_implicit_buffer_region(): # fmt: off @T.prim_func - def test(A: T.Buffer((10, 10, 10), 'float32', layout=T.TileLayout(T.S[10, 10, 10]))): + def test(A: T.Tensor((10, 10, 10), 'float32', layout=T.TileLayout(T.S[10, 10, 10]))): T.device_entry() Tx.memset(A[0], T.float32(0.0)) @@ -205,7 +205,7 @@ def test_roundtrip_alloc_under_any_scope(): def test(): T.device_entry() for i in T.serial(10): - A = T.alloc_buffer([100], "float32", scope="trn.sbuf", allocated_addr=1024) + A = T.alloc_tensor([100], "float32", scope="trn.sbuf", allocated_addr=1024) Tx.memset(A[i*10:i*10+10], T.float32(0.0)) # fmt: on @@ -247,13 +247,13 @@ def test(): # scalar buffer (alloc) C = T.shared_scalar("float16") D: T.float16 - pool = T.alloc_buffer([10], "uint8", scope="shared.dyn") + pool = T.alloc_tensor([10], "uint8", scope="shared.dyn") # scalar buffer (decl) E = T.decl_scalar("float16", pool.data, "shared.dyn", 0) # normal 1-dim buffer with shape (1,) F = T.alloc_local((1,), "float16") Ta: T.float16 - inner_pool = T.decl_buffer(shape=[10], data=pool.data, dtype="uint8", scope="shared.dyn") + inner_pool = T.decl_tensor(shape=[10], data=pool.data, dtype="uint8", scope="shared.dyn") test = Test(Ta, inner_pool) # noqa: F821 test.init() A[0] = C @@ -284,7 +284,7 @@ def test(): def test_alloc_apis_reject_name_argument(): with pytest.raises(TypeError): - T.alloc_buffer((1,), "int32", name="buf") + T.alloc_tensor((1,), "int32", name="buf") with pytest.raises(TypeError): T.local_scalar("int32", name="idx") @@ -294,25 +294,25 @@ def test_buffer(): # fmt: off @T.prim_func(private=True) def test( - A: T.Buffer((10, 11), "float32", layout=None), - B: T.Buffer((10, 11), "float32", scope="global"), - C: T.Buffer((10, 11), "float32", layout="default"), - D: T.Buffer((10, 11), "float32", layout=T.TileLayout(T.S[(10, 11) : (1, 10)])), - _E: T.Buffer([10, 11], 'float16', layout=None), - _F: T.Buffer([10, 11], 'float16', scope='global'), - _G: T.Buffer([10, 11], 'float16', layout='default'), - _H: T.Buffer([10, 11], 'float16', layout=T.TileLayout(T.S[(10, 11):(1, 10)])), + A: T.Tensor((10, 11), "float32", layout=None), + B: T.Tensor((10, 11), "float32", scope="global"), + C: T.Tensor((10, 11), "float32", layout="default"), + D: T.Tensor((10, 11), "float32", layout=T.TileLayout(T.S[(10, 11) : (1, 10)])), + _E: T.Tensor([10, 11], 'float16', layout=None), + _F: T.Tensor([10, 11], 'float16', scope='global'), + _G: T.Tensor([10, 11], 'float16', layout='default'), + _H: T.Tensor([10, 11], 'float16', layout=T.TileLayout(T.S[(10, 11):(1, 10)])), ): - _A0 = T.decl_buffer((10, 11), "float32", data=A.data, layout=None) - _B0 = T.decl_buffer((10, 11), "float32", data=B.data, scope="global") - _C0 = T.decl_buffer((10, 11), "float32", data=C.data, layout="default") - _D0 = T.decl_buffer((10, 11), "float32", data=D.data, layout=T.TileLayout(T.S[(10, 11) : (1, 10)])) # noqa: E501 - _A1 = T.alloc_buffer((10, 11), "float32", layout=None) - _B1 = T.alloc_buffer((10, 11), "float32", scope="global") - _C1 = T.alloc_buffer((10, 11), "float32", layout="default") - _D1 = T.alloc_buffer((10, 11), "float32", layout=T.TileLayout(T.S[(10, 11) : (1, 10)])) + _A0 = T.decl_tensor((10, 11), "float32", data=A.data, layout=None) + _B0 = T.decl_tensor((10, 11), "float32", data=B.data, scope="global") + _C0 = T.decl_tensor((10, 11), "float32", data=C.data, layout="default") + _D0 = T.decl_tensor((10, 11), "float32", data=D.data, layout=T.TileLayout(T.S[(10, 11) : (1, 10)])) # noqa: E501 + _A1 = T.alloc_tensor((10, 11), "float32", layout=None) + _B1 = T.alloc_tensor((10, 11), "float32", scope="global") + _C1 = T.alloc_tensor((10, 11), "float32", layout="default") + _D1 = T.alloc_tensor((10, 11), "float32", layout=T.TileLayout(T.S[(10, 11) : (1, 10)])) pass # fmt: on @@ -324,7 +324,7 @@ def test( @pytest.mark.skipif(sys.version_info < (3, 12), reason="PEP 695 requires Python 3.12") def test_buffer_shape_repeated_var_prints_out_of_line(): n = tvm.tirx.Var("n", "int32") - buffer = tvm.tirx.decl_buffer((n + n,), name="A") + buffer = tvm.tirx.decl_tensor((n + n,), name="A") func = tvm.tirx.PrimFunc([buffer], tvm.tirx.Evaluate(0)) code = func.script() @@ -373,14 +373,14 @@ def test_scalar_allocbuffer_annotation_sugar(): # fmt: off @T.prim_func def test(): - x = T.alloc_buffer((1,), "int32", scope="local") + x = T.alloc_tensor((1,), "int32", scope="local") x[0] = T.int32(0) T.evaluate(x[0]) # fmt: on code = test.script() assert "x: T.int32 = 0" in code - assert "x = T.alloc_buffer" not in code + assert "x = T.alloc_tensor" not in code assert from_source(code).script() == code assert_structural_equal(test, from_source(code)) @@ -390,7 +390,7 @@ def test_roundtrip_buffer_permute(): @T.prim_func def test() -> None: T.device_entry() - A = T.alloc_buffer([8, 4], dtype="float16", scope="local", + A = T.alloc_tensor([8, 4], dtype="float16", scope="local", layout=T.TileLayout(T.S[(8, 4) : (4, 1)])) B = A.permute(1, 0) B[0, 0] = T.float16(0) @@ -405,7 +405,7 @@ def test_roundtrip_buffer_local_auto(): @T.prim_func def test() -> None: T.device_entry() - A = T.alloc_buffer([2], dtype="float16", scope="local") + A = T.alloc_tensor([2], dtype="float16", scope="local") A_layout = T.TileLayout(T.S[(1, 2) : (2, 1)]) B = A.view(8, 8, layout=A_layout.tile(L_LANE, (8, 4), (1, 2))) B_local = B.local() @@ -423,7 +423,7 @@ def test_buffer_local_ir(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([2], dtype="float16", scope="local") + A = T.alloc_tensor([2], dtype="float16", scope="local") A_layout = T.TileLayout(T.S[(1, 2) : (2, 1)]) B = A.view(8, 8, layout=A_layout.tile(L_LANE, (8, 4), (1, 2))) B_local = B.local() @@ -444,7 +444,7 @@ def func() -> None: # Round-trip code = func.script() assert ( - 'v_1 = T.decl_buffer((2,), "float16", data=v.data, ' + 'v_1 = T.decl_tensor((2,), "float16", data=v.data, ' 'scope="local", layout="default")' ) in code assert from_source(code).script() == code @@ -465,7 +465,7 @@ def _collect_buffers(func): buffers = [] def visit(node): - if _is_buffer_binding(node, "tirx.alloc_buffer", "tirx.decl_buffer"): + if _is_buffer_binding(node, "tirx.alloc_tensor", "tirx.decl_tensor"): buffers.append(node.var) tvm_ffi.structural_walk(func.body, visit) @@ -477,7 +477,7 @@ def _buffer_source(func, buffer): sources = [] def visit(node): - if _is_buffer_binding(node, "tirx.decl_buffer") and node.var.same_as(buffer): + if _is_buffer_binding(node, "tirx.decl_tensor") and node.var.same_as(buffer): sources.append(node.value.args[0]) tvm_ffi.structural_walk(func.body, visit) @@ -493,7 +493,7 @@ def test_buffer_local_physical_order(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([32], dtype="float32", scope="local") + A = T.alloc_tensor([32], dtype="float32", scope="local") B = A.view(64, 64, layout=tcgen05_atom_layout("16x256b", (64, 64), "float32")) B_flat = B.local() B_2d = B.local(4, 8) @@ -520,11 +520,11 @@ def func() -> None: code = func.script() assert ( - 'v_1 = T.decl_buffer((32,), "float32", data=v.data, ' + 'v_1 = T.decl_tensor((32,), "float32", data=v.data, ' 'scope="local", layout="default")' ) in code assert ( - 'v_2 = T.decl_buffer((4, 8), "float32", data=v.data, ' + 'v_2 = T.decl_tensor((4, 8), "float32", data=v.data, ' 'scope="local", layout="default")' ) in code assert from_source(code).script() == code @@ -539,7 +539,7 @@ def test_buffer_local_layout_overrides_roundtrip(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([32], dtype="float32", scope="local") + A = T.alloc_tensor([32], dtype="float32", scope="local") B = A.view(64, 64, layout=tcgen05_atom_layout("16x256b", (64, 64), "float32")) B_storage = B.local(layout=B.layout.storage()) # An explicit layout is an escape hatch and may describe a smaller @@ -558,10 +558,10 @@ def func() -> None: storage_line = next(line for line in code.splitlines() if "v_1 =" in line) custom_line = next(line for line in code.splitlines() if "v_2 =" in line) assert ( - 'T.decl_buffer((32,), "float32", data=v.data, scope="local", layout=' + 'T.decl_tensor((32,), "float32", data=v.data, scope="local", layout=' ) in storage_line assert ( - 'T.decl_buffer((2, 4), "float32", data=v.data, scope="local", ' + 'T.decl_tensor((2, 4), "float32", data=v.data, scope="local", ' 'layout=' ) in custom_line assert_structural_equal(func, from_source(code)) @@ -575,7 +575,7 @@ def test_buffer_local_explicit_layout_without_parent_layout(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer((4,), dtype="float32", scope="local", layout=None) + A = T.alloc_tensor((4,), dtype="float32", scope="local", layout=None) B = A.local(4, layout=T.TileLayout(T.S[4])) B[0] = T.float32(1) # fmt: on @@ -596,7 +596,7 @@ def test_buffer_local_compose_layout_printer_roundtrip(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer( + A = T.alloc_tensor( (8, 8), dtype="float32", scope="local", @@ -611,7 +611,7 @@ def func() -> None: assert b_buf.ty.layout.is_trivial() code = func.script() local_line = next(line for line in code.splitlines() if "v =" in line) - assert 'T.decl_buffer((64,), "float32", data=A.data, scope="local", layout=' in local_line + assert 'T.decl_tensor((64,), "float32", data=A.data, scope="local", layout=' in local_line parsed = from_source(code) assert_structural_equal(func, parsed) assert parsed.script() == code @@ -625,7 +625,7 @@ def test_buffer_local_inference_without_parent_layout_has_clear_diagnostic(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer((4,), dtype="float32", scope="local", layout=None) + A = T.alloc_tensor((4,), dtype="float32", scope="local", layout=None) B = A.local(layout=T.TileLayout(T.S[4])) B[0] = T.float32(1) @@ -637,7 +637,7 @@ def test_buffer_local_physical_span_includes_gaps_and_offset(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([6], dtype="float32", scope="local") + A = T.alloc_tensor([6], dtype="float32", scope="local") B = A.view(32, 2, layout=T.TileLayout(T.S[(32, 2) : (1 @ laneid, 2)] + 3)) B_flat = B.local() B_2d = B.local(2, 3) @@ -664,7 +664,7 @@ def func() -> None: code = func.script() storage_line = next(line for line in code.splitlines() if "v_3 =" in line) - assert 'T.decl_buffer((2,), "float32", data=v.data, scope="local", layout=' in storage_line + assert 'T.decl_tensor((2,), "float32", data=v.data, scope="local", layout=' in storage_line assert_structural_equal(func, from_source(code)) assert from_source(code).script() == code @@ -677,7 +677,7 @@ def test_buffer_local_printer_is_stable_with_multiple_aliases(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([32], dtype="float32", scope="local") + A = T.alloc_tensor([32], dtype="float32", scope="local") B = A.view(64, 64, layout=tcgen05_atom_layout("16x256b", (64, 64), "float32")) B_flat = B.local() B_2d = B.local(4, 8) @@ -687,16 +687,16 @@ def func() -> None: expected = func.script() assert ( - 'v_1 = T.decl_buffer((32,), "float32", data=v.data, ' + 'v_1 = T.decl_tensor((32,), "float32", data=v.data, ' 'scope="local", layout="default")' ) in expected assert ( - 'v_2 = T.decl_buffer((4, 8), "float32", data=v.data, ' + 'v_2 = T.decl_tensor((4, 8), "float32", data=v.data, ' 'scope="local", layout="default")' ) in expected storage_line = next(line for line in expected.splitlines() if "v_3 =" in line) assert ( - 'T.decl_buffer((32,), "float32", data=v.data, scope="local", layout=' + 'T.decl_tensor((32,), "float32", data=v.data, scope="local", layout=' ) in storage_line for _ in range(20): parsed = from_source(expected) @@ -711,14 +711,14 @@ def test_buffer_local_printer_preserves_inherited_metadata(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer( + A = T.alloc_tensor( [32, 2], dtype="float32", elem_offset=8, scope="local", layout=T.TileLayout(T.S[(32, 2) : (1 @ laneid, 2)]), ) - B_align = T.decl_buffer( + B_align = T.decl_tensor( (2,), dtype="float32", data=A.data, @@ -726,7 +726,7 @@ def func() -> None: scope="local", align=128, ) - B_factor = T.decl_buffer( + B_factor = T.decl_tensor( (2,), dtype="float32", data=A.data, @@ -740,8 +740,8 @@ def func() -> None: code = func.script() align_line = next(line for line in code.splitlines() if "B_align =" in line) factor_line = next(line for line in code.splitlines() if "B_factor =" in line) - assert "T.decl_buffer" in align_line and "align=128" in align_line - assert "T.decl_buffer" in factor_line and "offset_factor=8" in factor_line + assert "T.decl_tensor" in align_line and "align=128" in align_line + assert "T.decl_tensor" in factor_line and "offset_factor=8" in factor_line assert ".local(" not in align_line assert ".local(" not in factor_line parsed = from_source(code) @@ -757,7 +757,7 @@ def test_buffer_local_rejects_shape_that_does_not_match_physical_span(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([6], dtype="float32", scope="local") + A = T.alloc_tensor([6], dtype="float32", scope="local") B = A.view(32, 2, layout=T.TileLayout(T.S[(32, 2) : (1 @ laneid, 2)] + 3)) B_local = B.local(2) B_local[0] = T.float32(0) @@ -770,7 +770,7 @@ def test_buffer_permute_ir(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([8, 4], dtype="float16", scope="local", + A = T.alloc_tensor([8, 4], dtype="float16", scope="local", layout=T.TileLayout(T.S[(8, 4) : (4, 1)])) B = A.permute(1, 0) B[0, 0] = T.float16(0) @@ -794,7 +794,7 @@ def test_buffer_rearrange_allows_arbitrary_axis_names(): @T.prim_func def ordinary_axis() -> None: T.device_entry() - A = T.alloc_buffer( + A = T.alloc_tensor( (8, 4), "float16", scope="local", @@ -806,7 +806,7 @@ def ordinary_axis() -> None: @T.prim_func def buf_axis() -> None: T.device_entry() - A = T.alloc_buffer( + A = T.alloc_tensor( (8, 4), "float16", scope="local", @@ -818,7 +818,7 @@ def buf_axis() -> None: @T.prim_func def self_axis() -> None: T.device_entry() - A = T.alloc_buffer( + A = T.alloc_tensor( (8, 4), "float16", scope="local", @@ -830,7 +830,7 @@ def self_axis() -> None: @T.prim_func def pattern_axis() -> None: T.device_entry() - A = T.alloc_buffer( + A = T.alloc_tensor( (8, 4), "float16", scope="local", @@ -842,7 +842,7 @@ def pattern_axis() -> None: @T.prim_func def keyword_pattern() -> None: T.device_entry() - A = T.alloc_buffer( + A = T.alloc_tensor( (8, 4), "float16", scope="local", @@ -867,7 +867,7 @@ def test_buffer_permute_compose_layout_ir(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer( + A = T.alloc_tensor( [4, 4, 4, 64], dtype="bfloat16", scope="shared.dyn", layout=T.ComposeLayout(3, 3, 3, T.TileLayout(T.S[(4, 4, 4, 64) : (1024, 256, 64, 1)])), ) @@ -900,7 +900,7 @@ def test_buffer_sub_multi_iter_dim_ir(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([8, 16], dtype="float16", scope="local", + A = T.alloc_tensor([8, 16], dtype="float16", scope="local", layout=T.TileLayout(T.S[(2, 4, 16) : (1024, 64, 1)])) B = A.sub[5] B[0] = T.float16(0) @@ -917,7 +917,7 @@ def func() -> None: def test_buffer_sub_multi_iter_misaligned_rejected(): - buf = tvm.tirx.decl_buffer( + buf = tvm.tirx.decl_tensor( (8, 16), "float16", layout=tvm.tirx.layout.TileLayout(T.S[(2, 4, 16) : (1024, 64, 1)]) ) # sub[2:6] narrows the multi-iter dim 0 at a misaligned offset. @@ -934,7 +934,7 @@ def test_buffer_sub_ir(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([4, 8, 16], dtype="float16", scope="local", + A = T.alloc_tensor([4, 8, 16], dtype="float16", scope="local", layout=T.TileLayout(T.S[(4, 8, 16) : (256, 16, 1)])) B = A.sub[1, 2:6] B[0, 0] = T.float16(0) @@ -961,10 +961,10 @@ def func() -> None: def test_buffer_view_surgery_static_bounds_rejected(): """Statically-known out-of-range sub arguments must be rejected loudly (review finding: OOB offsets were silent).""" - buf = tvm.tirx.decl_buffer( + buf = tvm.tirx.decl_tensor( (10,), "float16", layout=tvm.tirx.layout.TileLayout(T.S[(10,) : (1,)]) ) - grid = tvm.tirx.decl_buffer( + grid = tvm.tirx.decl_tensor( (4, 8), "float16", layout=tvm.tirx.layout.TileLayout(T.S[(4, 8) : (8, 1)]) ) # int index: static bounds @@ -1013,7 +1013,7 @@ def addr(buf, base, *coords): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([4, 1024], dtype="bfloat16", scope="shared.dyn", layout=compose) + A = T.alloc_tensor([4, 1024], dtype="bfloat16", scope="shared.dyn", layout=compose) B = A.sub[1] # offset 1024 = 2 * period: folds into elem_offset B[0] = T.bfloat16(0) C = A.sub[:, 512:1024] # offset 512 = period: folds into elem_offset @@ -1039,7 +1039,7 @@ def func() -> None: @T.prim_func def func2() -> None: T.device_entry() - A = T.alloc_buffer([2, 16, 8], dtype="float16", scope="shared.dyn", layout=compose2) + A = T.alloc_tensor([2, 16, 8], dtype="float16", scope="shared.dyn", layout=compose2) B = A.sub[:, 1] # offset 8 B[0, 0] = T.float16(0) C = A.sub[:, :, 1] # offset 1 @@ -1086,7 +1086,7 @@ def func2() -> None: @T.prim_func def func3() -> None: T.device_entry() - A = T.alloc_buffer([64], dtype="bfloat16", scope="shared.dyn", layout=compose3) + A = T.alloc_tensor([64], dtype="bfloat16", scope="shared.dyn", layout=compose3) B = A.sub[8:16] B[0] = T.bfloat16(0) # fmt: on @@ -1105,7 +1105,7 @@ def test_buffer_tile_ir(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([3, 64, 512], dtype="float16", scope="shared", + A = T.alloc_tensor([3, 64, 512], dtype="float16", scope="shared", layout=T.TileLayout(T.S[(3, 64, 512) : (64 * 512, 512, 1)])) for w in T.serial(4): B = A.tile((1, (-1, 4, 4)))[:, w, :] @@ -1124,7 +1124,7 @@ def func() -> None: @T.prim_func def func_multi() -> None: T.device_entry() - A = T.alloc_buffer([64, 128], dtype="float16", scope="shared", + A = T.alloc_tensor([64, 128], dtype="float16", scope="shared", layout=T.TileLayout(T.S[(64, 128) : (128, 1)])) for wx in T.serial(4): for wy in T.serial(2): @@ -1137,7 +1137,7 @@ def func_multi() -> None: @T.prim_func def func_multipick() -> None: T.device_entry() - A = T.alloc_buffer([128, 16], dtype="float16", scope="shared", + A = T.alloc_tensor([128, 16], dtype="float16", scope="shared", layout=T.TileLayout(T.S[(128, 16) : (16, 1)])) for a in T.serial(2): for b in T.serial(4): @@ -1165,7 +1165,7 @@ def func_multipick() -> None: def test_buffer_tile_rejected(): - buf = tvm.tirx.decl_buffer( + buf = tvm.tirx.decl_tensor( (3, 64, 512), "float16", layout=tvm.tirx.layout.TileLayout(T.S[(3, 64, 512) : (64 * 512, 512, 1)]), @@ -1194,10 +1194,10 @@ def test_buffer_chunk_ir(): the hand-written a*k:(a+1)*k slice — no reshape, no extra dim.""" compose = T.ComposeLayout(3, 3, 3, T.TileLayout(T.S[(4, 512) : (512, 1)])) - A = tvm.tirx.decl_buffer( + A = tvm.tirx.decl_tensor( (4, 8, 16), "float16", layout=tvm.tirx.layout.TileLayout(T.S[(4, 8, 16) : (128, 16, 1)]) ) - C = tvm.tirx.decl_buffer((4, 512), "bfloat16", layout=compose) + C = tvm.tirx.decl_tensor((4, 512), "bfloat16", layout=compose) # chunk((None, None, 2))[:, :, 1] narrows dim 2 (extent 16) to chunk 1 of 2 # → [8:16] (k = 16 // 2 = 8); rank preserved, dims 0/1 pass through as ':'. @@ -1242,7 +1242,7 @@ def test_buffer_view_dtype_ir(): @T.prim_func def func() -> None: T.device_entry() - A = T.alloc_buffer([8, 8], dtype="float16", scope="local") + A = T.alloc_tensor([8, 8], dtype="float16", scope="local") B = A.view("float32") B[0, 0] = T.float32(0) # fmt: on @@ -1262,9 +1262,9 @@ def func() -> None: def test_buffer_slice_region(): - """Verify A[slice] returns BufferRegion (not DeclBuffer).""" + """Verify A[slice] returns BufferRegion (not DeclTensor).""" - buf = tvm.tirx.decl_buffer((128, 64), "float16") + buf = tvm.tirx.decl_tensor((128, 64), "float16") br = buf[32:64, 0:32] assert isinstance(br, TensorRegion) assert br.source.same_as(buf) @@ -1310,9 +1310,9 @@ def add(a: T.float32, b: T.float32) -> T.float32: @T.prim_func def main( - A: T.Buffer((16,), "float32"), - B: T.Buffer((16,), "float32"), - C: T.Buffer((16,), "float32"), + A: T.Tensor((16,), "float32"), + B: T.Tensor((16,), "float32"), + C: T.Tensor((16,), "float32"), ): for i in range(16): C[i] = Module.add(A[i], B[i]) @@ -1329,17 +1329,17 @@ def test_buffer_sub_tmem_offset_uses_physical_columns(): @T.prim_func def func() -> None: T.device_entry() - Q = T.decl_buffer( + Q = T.decl_tensor( (2, 64, 288), "bfloat16", scope="tmem", allocated_addr=256, layout=T.TileLayout(T.S[(2, 64, 288) : (64 @ TLane, 1 @ TLane, 1 @ TCol)]), ) Q_tail = Q.sub[:, :, 256:288] - F8 = T.decl_buffer( + F8 = T.decl_tensor( (64, 128), "float8_e4m3fn", scope="tmem", allocated_addr=32, layout=T.TileLayout(T.S[(64, 128) : (1 @ TLane, 1 @ TCol)]), ) F8_tail = F8.sub[:, 64:96] - F32 = T.decl_buffer( + F32 = T.decl_tensor( (64, 128), "float32", scope="tmem", allocated_addr=64, layout=T.TileLayout(T.S[(64, 128) : (1 @ TLane, 1 @ TCol)]), ) @@ -1369,7 +1369,7 @@ def build(): @T.prim_func def func() -> None: T.device_entry() - A = T.decl_buffer( + A = T.decl_tensor( (64, 16), "bfloat16", scope="tmem", allocated_addr=0, layout=buf_layout, ) _ = A.sub[:, 1:3] @@ -1382,7 +1382,7 @@ def func() -> None: def test_roundtrip_tmem_decl_buffer(): - """DeclBuffer with tmem scope: data kwarg must be suppressed, allocated_addr + """DeclTensor with tmem scope: data kwarg must be suppressed, allocated_addr must print as Expr (not Array), and scalar buffer index must not get a .source suffix.""" @@ -1392,8 +1392,8 @@ def func(): with T.launch_thread("blockIdx.x", 1): T.launch_thread("threadIdx.x", 128) addr = T.alloc_shared((1,), "uint32", layout=None) - addr_alias = T.decl_buffer((1,), "uint32", data=addr.data, scope="shared") - buf = T.decl_buffer((64,), scope="tmem", layout=None, allocated_addr=addr_alias[0]) + addr_alias = T.decl_tensor((1,), "uint32", data=addr.data, scope="shared") + buf = T.decl_tensor((64,), scope="tmem", layout=None, allocated_addr=addr_alias[0]) # fmt: on code = func.script() @@ -1402,7 +1402,7 @@ def func(): decls = [] tvm_ffi.structural_walk( func.body, - lambda node: decls.append(node) if _is_buffer_binding(node, "tirx.decl_buffer") else None, + lambda node: decls.append(node) if _is_buffer_binding(node, "tirx.decl_tensor") else None, ) # The shared alias has an explicit definition before the tensor-memory use. assert len(decls) == 2 diff --git a/tests/python/tirx/script/test_tirx_script_dynamic_shape.py b/tests/python/tirx/script/test_tirx_script_dynamic_shape.py index 74b0c055f04f..bf37bb0856d4 100644 --- a/tests/python/tirx/script/test_tirx_script_dynamic_shape.py +++ b/tests/python/tirx/script/test_tirx_script_dynamic_shape.py @@ -33,10 +33,10 @@ def test_tir_bound_prim_param_reused_in_dependent_annotations(): @T.prim_func def func( n: T.int32, - direct: T.Buffer((n,), "float32"), - repeated: T.Buffer((n,), "float32"), - compound: T.Buffer((n + 1,), "float32"), - ) -> T.Buffer((n,), "float32"): + direct: T.Tensor((n,), "float32"), + repeated: T.Tensor((n,), "float32"), + compound: T.Tensor((n + 1,), "float32"), + ) -> T.Tensor((n,), "float32"): return repeated n, direct, repeated, compound = func.params @@ -50,7 +50,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.Tensor((n + 1,), "float32")): T.evaluate(n) n, A = Module["main"].params @@ -62,13 +62,13 @@ def test_tir_external_symbol_adopted_by_later_prim_param(dtype): n = T.dynamic("n", dtype) @T.prim_func - def func(A: T.Buffer((n,), "float32"), n: n): + def func(A: T.Tensor((n,), "float32"), n: n): T.evaluate(n) @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), "float32"), n: n): + def main(A: T.Tensor((n,), "float32"), n: n): T.evaluate(n) for function in [func, Module["main"]]: @@ -81,7 +81,7 @@ def test_tir_external_dynamic_symbol_preserves_dtype(): n = T.dynamic("n", "int64") @T.prim_func - def func(A: T.Buffer((n,), "float32")): + def func(A: T.Tensor((n,), "float32")): T.evaluate(n) n = func.params[0].ty.shape[0] @@ -93,19 +93,19 @@ def test_tir_undeclared_shape_symbol_is_undefined(): with pytest.raises(NameError): @T.prim_func - def main(A: T.Buffer((n, n), "float32")): # noqa: F821 + def main(A: T.Tensor((n, n), "float32")): # noqa: F821 T.evaluate(0) def test_tir_direct_later_prim_param_reuses_shape_symbol(): @T.prim_func - def func(A: T.Buffer((n,), "float32"), n: T.int32): + def func(A: T.Tensor((n,), "float32"), n: T.int32): T.evaluate(n) @I.ir_module class Module: @T.prim_func - def main(A: T.Buffer((n,), "float32"), n: T.int32): + def main(A: T.Tensor((n,), "float32"), n: T.int32): T.evaluate(n) for function in [func, Module["main"]]: @@ -119,8 +119,8 @@ def test_tir_return_annotation_does_not_define_symbolic_var(): with pytest.raises(NameError): @T.prim_func - def main() -> T.Buffer((n,), "float32"): # noqa: F821 - A = T.alloc_buffer((n,), "float32") # noqa: F821 + def main() -> T.Tensor((n,), "float32"): # noqa: F821 + A = T.alloc_tensor((n,), "float32") # noqa: F821 return A @@ -130,18 +130,18 @@ def test_type_vars_roundtrip(): UNUSED = I.dynamic("UNUSED") @T.prim_func(private=True) - def func(A: T.Buffer((M, M * 2), "float32")): + def func(A: T.Tensor((M, M * 2), "float32")): A[0, 0] = T.float32(1) script = func.script() assert script.startswith("from __future__ import annotations\n\n") assert "def main[M](" in script - assert 'T.Buffer((M, M * T.int64(2)), "float32", layout="default")' in script + assert 'T.Tensor((M, M * T.int64(2)), "float32", layout="default")' in script assert "M = T.int64()" not in script typed = tvm.script.from_source( """ @T.prim_func(private=True) -def func[M: int](A: T.Buffer((M, M * 2), "float32")): +def func[M: int](A: T.Tensor((M, M * 2), "float32")): A[0, 0] = T.float32(1) """, extra_vars={"I": tvm.script.ir, "T": tvm.script.tirx}, @@ -152,7 +152,7 @@ def func[M: int](A: T.Buffer((M, M * 2), "float32")): repeated = func.script() assert "from __future__ import annotations" in repeated assert "def main[M](" in repeated - assert 'T.Buffer((M, M * T.int64(2)), "float32", layout="default")' in repeated + assert 'T.Tensor((M, M * T.int64(2)), "float32", layout="default")' in repeated assert "UNUSED" not in script assert "M = T.int64()" not in repeated assert len(func.params) == 1 @@ -172,7 +172,7 @@ def test_dynamic_int32_roundtrip(): n = I.dynamic("n", "int32") @T.prim_func(private=True) - def func(A: T.Buffer((n,), "float32")): + def func(A: T.Tensor((n,), "float32")): A[0] = T.float32(1) source = func.script() @@ -219,7 +219,7 @@ def test_captured_shape_requires_concrete_symbols(): # Native shape construction preserves concrete symbols and rejects strings. def build(shape): @T.prim_func - def main(x: T.Buffer(shape, "float32")): + def main(x: T.Tensor(shape, "float32")): T.evaluate(0) return main @@ -243,7 +243,7 @@ def test_dynamic_symbols_are_fresh_and_scope_independent(): @I.ir_module class Module: @T.prim_func - def first(x: T.Buffer((n,), "float32")): + def first(x: T.Tensor((n,), "float32")): T.evaluate(n) assert Module["first"].params[0].ty.shape[0].same_as(n) diff --git a/tests/python/tirx/script/test_tirx_script_gpu.py b/tests/python/tirx/script/test_tirx_script_gpu.py index 874c25552915..4a582aba15b2 100644 --- a/tests/python/tirx/script/test_tirx_script_gpu.py +++ b/tests/python/tirx/script/test_tirx_script_gpu.py @@ -31,13 +31,13 @@ def test_roundtrip_scopeid1(): # fmt: off @T.prim_func - def test(A: T.Buffer((64,), 'float32', scope='global')) -> None: + def test(A: T.Tensor((64,), 'float32', scope='global')) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) warp_id = T.warp_id([1]) lane_id = T.lane_id([32]) - A_local = T.alloc_buffer([1], dtype="float16", scope="local") + A_local = T.alloc_tensor([1], dtype="float16", scope="local") for i in T.serial(2): A_local[0] = A[lane_id * 2 + i] # fmt: on @@ -54,7 +54,7 @@ def from_source(code): def test_roundtrip_scopeid2(): # fmt: off @T.prim_func - def test(_: T.Buffer((64,), 'float32', scope='global')) -> None: + def test(_: T.Tensor((64,), 'float32', scope='global')) -> None: T.device_entry() bx, by, bz = T.cta_id([8, 10, 12]) @@ -79,7 +79,7 @@ def test_roundtrip_scopeid_deferred(): # fmt: off @T.prim_func(private=True) - def test(_: T.Buffer((64,), 'float32', scope='global')) -> None: + def test(_: T.Tensor((64,), 'float32', scope='global')) -> None: T.device_entry() bx = T.cta_id() # deferred kernel→cta @@ -100,7 +100,7 @@ def test(_: T.Buffer((64,), 'float32', scope='global')) -> None: def test_exec_scope_filter_guard_roundtrip(): @T.prim_func(private=True) - def test(A: T.Buffer((1,), "float32", scope="global")) -> None: + def test(A: T.Tensor((1,), "float32", scope="global")) -> None: T.device_entry() T.cta_id([1]) tx = T.thread_id([128]) @@ -115,13 +115,13 @@ def test(A: T.Buffer((1,), "float32", scope="global")) -> None: def test_roundtrip_op1(): # fmt: off @T.prim_func - def test(A: T.Buffer((64,), 'float32', scope='global')) -> None: + def test(A: T.Tensor((64,), 'float32', scope='global')) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) warp_id = T.warp_id([1]) lane_id = T.lane_id([32]) - A_smem = T.alloc_buffer([64], dtype="float32", scope="shared") + A_smem = T.alloc_tensor([64], dtype="float32", scope="shared") Tx.cta.copy(A_smem, A) for i in range(10): @@ -139,19 +139,19 @@ def test_roundtrip_op2(): # fmt: off @T.prim_func def test( - A: T.Buffer((128, 128), "float16", scope="global"), - B: T.Buffer((128, 64), "float16", scope="global"), - C: T.Buffer((128, 64), "float32", scope="global"), + A: T.Tensor((128, 128), "float16", scope="global"), + B: T.Tensor((128, 64), "float16", scope="global"), + C: T.Tensor((128, 64), "float32", scope="global"), ) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) warp_id = T.warp_id([4]) lane_id = T.lane_id([32]) - A_smem = T.alloc_buffer([128, 32], dtype="float16", scope="shared") - B_smem = T.alloc_buffer([32, 64], dtype="float16", scope="shared") + A_smem = T.alloc_tensor([128, 32], dtype="float16", scope="shared") + B_smem = T.alloc_tensor([32, 64], dtype="float16", scope="shared") - C_local = T.alloc_buffer([128, 64], dtype="float32", scope="local") + C_local = T.alloc_tensor([128, 64], dtype="float32", scope="local") for k in range(4): Tx.cta.copy(A_smem, A[:, k * 32 : k * 32 + 32]) Tx.cta.copy(B_smem, B[k * 32 : k * 32 + 32, 0:64]) @@ -171,19 +171,19 @@ def test_roundtrip_op3(): @T.prim_func def test( - A: T.Buffer((128, K), "float16", scope="global"), - B: T.Buffer((K, 64), "float16", scope="global"), - C: T.Buffer((128, 64), "float32", scope="global"), + A: T.Tensor((128, K), "float16", scope="global"), + B: T.Tensor((K, 64), "float16", scope="global"), + C: T.Tensor((128, 64), "float32", scope="global"), ) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) warp_id = T.warp_id([4]) lane_id = T.lane_id([32]) - A_smem = T.alloc_buffer([NUM_STAGES, 128, 32], dtype="float16", scope="shared") - B_smem = T.alloc_buffer([NUM_STAGES, 32, 64], dtype="float16", scope="shared") + A_smem = T.alloc_tensor([NUM_STAGES, 128, 32], dtype="float16", scope="shared") + B_smem = T.alloc_tensor([NUM_STAGES, 32, 64], dtype="float16", scope="shared") - C_local = T.alloc_buffer([128, 64], dtype="float32", scope="local") + C_local = T.alloc_tensor([128, 64], dtype="float32", scope="local") for i in range(NUM_STAGES - 1): Tx.cta.copy(A_smem[i, :, :], A[:, i * 32 : i * 32 + 32]) Tx.cta.copy(B_smem[i, :, :], B[i * 32 : i * 32 + 32, :]) @@ -207,7 +207,7 @@ def test( def test_roundtrip_tensormap(): # fmt: off @T.prim_func - def func1(A: T.Buffer([128], "float32")): + def func1(A: T.Tensor([128], "float32")): T.func_attr({"global_symbol": "func"}) A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) @@ -237,11 +237,11 @@ def test_roundtrip_op_call_workspace(): # fmt: off @T.prim_func def test( - A: T.Buffer([10], "float32", scope="global"), B: T.Buffer([10], "float32", scope="global") + A: T.Tensor([10], "float32", scope="global"), B: T.Tensor([10], "float32", scope="global") ): T.device_entry() - smem = T.alloc_buffer([10], "float32", scope="shared") + smem = T.alloc_tensor([10], "float32", scope="shared") Tx.add(B, A, T.float32(1), workspace={"smem": smem}) # fmt: on code = test.script() @@ -253,7 +253,7 @@ def test_roundtrip_op_call_config(): # fmt: off @T.prim_func def test( - A: T.Buffer([10], "float32", scope="global"), B: T.Buffer([10], "float32", scope="global") + A: T.Tensor([10], "float32", scope="global"), B: T.Tensor([10], "float32", scope="global") ): T.device_entry() @@ -269,8 +269,8 @@ def test_predicate(): @T.prim_func def test(): T.device_entry() - A = T.alloc_buffer([10, 10], "float32") - B = T.alloc_buffer([10, 10], "float32") + A = T.alloc_tensor([10, 10], "float32") + B = T.alloc_tensor([10, 10], "float32") Tx.select(B, A, 1.0, lambda i, j: i < j) # fmt: on code = test.script() @@ -281,7 +281,7 @@ def test(): def test_kwargs_op_call(): # fmt: off @T.prim_func(private=True) - def test(A: T.Buffer((10, 10), "float32"), B: T.Buffer((10, 10), "float32")): + def test(A: T.Tensor((10, 10), "float32"), B: T.Tensor((10, 10), "float32")): T.device_entry() kwargs = T.meta_var({"dispatch": "tma_auto", "cta_group": 2}) Tx.copy_async(A[:, :], B[:, :], **kwargs) @@ -299,9 +299,9 @@ def test_workspace_default_none(): ``if workspace is None: workspace = {}`` guard.""" from tvm.tirx import BufferRegion - A_buf = tvm.tirx.decl_buffer((128, 128), "float16", name="A") - B_buf = tvm.tirx.decl_buffer((128, 128), "float16", name="B") - C_buf = tvm.tirx.decl_buffer((128,), "float16", name="C") + A_buf = tvm.tirx.decl_tensor((128, 128), "float16", name="A") + B_buf = tvm.tirx.decl_tensor((128, 128), "float16", name="B") + C_buf = tvm.tirx.decl_tensor((128,), "float16", name="C") A = BufferRegion(A_buf, [tvm.ir.Range(0, 128), tvm.ir.Range(0, 128)]) B = BufferRegion(B_buf, [tvm.ir.Range(0, 128), tvm.ir.Range(0, 128)]) C = BufferRegion(C_buf, [tvm.ir.Range(0, 128)]) @@ -333,7 +333,7 @@ def test_roundtrip_persistent_decorator(): # fmt: off @T.prim_func(persistent=True) - def test(A: T.Buffer((128,), 'float32', scope='global')) -> None: + def test(A: T.Tensor((128,), 'float32', scope='global')) -> None: T.device_entry() cta_id = T.cta_id([1]) @@ -354,7 +354,7 @@ def test_roundtrip_persistent_not_present(): # fmt: off @T.prim_func - def test(A: T.Buffer((128,), 'float32', scope='global')) -> None: + def test(A: T.Tensor((128,), 'float32', scope='global')) -> None: T.device_entry() cta_id = T.cta_id([1]) @@ -373,7 +373,7 @@ def test_warp_role(): # fmt: off @T.prim_func - def test(A: T.Buffer((128,), 'float32', scope='global')) -> None: + def test(A: T.Tensor((128,), 'float32', scope='global')) -> None: T.device_entry() cta_id = T.cta_id([1]) @@ -408,7 +408,7 @@ def test_warpgroup_role(): # fmt: off @T.prim_func - def test(A: T.Buffer((128,), 'float32', scope='global')) -> None: + def test(A: T.Tensor((128,), 'float32', scope='global')) -> None: T.device_entry() cta_id = T.cta_id([1]) @@ -453,12 +453,12 @@ def test_roundtrip_cp_async_bulk_tensor_g2s_cluster(): # fmt: off @T.prim_func(check_well_formed=False) - def func(_: T.Buffer((16, 16), 'float32')): + def func(_: T.Tensor((16, 16), 'float32')): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) with T.launch_thread("blockIdx.x", 1): T.launch_thread("threadIdx.x", 128) - A_smem = T.alloc_buffer((16, 16), "float32", scope="shared") + A_smem = T.alloc_tensor((16, 16), "float32", scope="shared") T.ptx["cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes"]( A_smem.data, T.address_of(A_map), 0, 0, T.uint32(0) ) @@ -474,12 +474,12 @@ def test_roundtrip_cp_async_bulk_tensor_s2g(): # fmt: off @T.prim_func(check_well_formed=False) - def func(_: T.Buffer((16, 16), 'float32')): + def func(_: T.Tensor((16, 16), 'float32')): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) with T.launch_thread("blockIdx.x", 1): T.launch_thread("threadIdx.x", 128) - A_smem = T.alloc_buffer((16, 16), "float32", scope="shared") + A_smem = T.alloc_tensor((16, 16), "float32", scope="shared") T.ptx["cp.async.bulk.tensor.2d.global.shared::cta.tile.bulk_group"]( T.address_of(A_map), 0, 0, A_smem.data ) @@ -495,7 +495,7 @@ def test_roundtrip_cp_async_bulk_tensor_prefetch(): # fmt: off @T.prim_func(check_well_formed=False) - def func(_: T.Buffer((16, 16), 'float32')): + def func(_: T.Tensor((16, 16), 'float32')): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) with T.launch_thread("blockIdx.x", 1): @@ -515,12 +515,12 @@ def test_roundtrip_cp_async_bulk_tensor_s2g_reduce(): # fmt: off @T.prim_func(check_well_formed=False) - def func(_: T.Buffer((16, 16), 'float32')): + def func(_: T.Tensor((16, 16), 'float32')): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) with T.launch_thread("blockIdx.x", 1): T.launch_thread("threadIdx.x", 128) - A_smem = T.alloc_buffer((16, 16), "float32", scope="shared") + A_smem = T.alloc_tensor((16, 16), "float32", scope="shared") T.ptx["cp.reduce.async.bulk.tensor.2d.global.shared::cta.add.tile.bulk_group"]( T.address_of(A_map), 0, 0, A_smem.data ) @@ -534,7 +534,7 @@ def func(_: T.Buffer((16, 16), 'float32')): def test_scope_id_dtype_uint32(): # fmt: off @T.prim_func - def func(A: T.Buffer((128,), 'float32')): + def func(A: T.Tensor((128,), 'float32')): T.device_entry() bx = T.cta_id([1]) @@ -570,7 +570,7 @@ def _assert_roundtrip(func): def test_scope_id_dtype_uint32_lane_and_warp(): # fmt: off @T.prim_func - def func(A: T.Buffer((32,), 'float32')): + def func(A: T.Tensor((32,), 'float32')): T.device_entry() _ = T.cta_id([1]) @@ -588,7 +588,7 @@ def func(A: T.Buffer((32,), 'float32')): def test_scope_id_dtype_uint32_with_preferred(): # fmt: off @T.prim_func - def func(A: T.Buffer((4,), 'float32')): + def func(A: T.Tensor((4,), 'float32')): T.device_entry() _ = T.cluster_id([2]) @@ -608,7 +608,7 @@ def test_scope_id_dtype_uint32_deferred_extent(): # fmt: off @T.prim_func - def func(A: T.Buffer((32,), 'float32')): + def func(A: T.Tensor((32,), 'float32')): T.device_entry() _ = T.cta_id([1]) @@ -636,7 +636,7 @@ def test_scope_id_dtype_rejects_unsupported(dtype): with pytest.raises(Exception, match='must be "int32" or "uint32"'): @T.prim_func - def func(A: T.Buffer((128,), 'float32')): + def func(A: T.Tensor((128,), 'float32')): T.device_entry() _ = T.cta_id([1]) diff --git a/tests/python/tirx/script/test_tirx_script_ir_builder.py b/tests/python/tirx/script/test_tirx_script_ir_builder.py index fc6e80de6cdb..63ade9a7ea84 100644 --- a/tests/python/tirx/script/test_tirx_script_ir_builder.py +++ b/tests/python/tirx/script/test_tirx_script_ir_builder.py @@ -269,13 +269,13 @@ def test_ir_builder_tir_inline(): ], ) def test_concrete_buffer_indices(shape, index, expected): - buffer = tirx.decl_buffer(shape, "float32") + buffer = tirx.decl_tensor(shape, "float32") assert T.buffer_indices(buffer, index) == expected def test_symbolic_buffer_indices(): m, n, k = [tirx.Var(name, "int32") for name in ("m", "n", "k")] - buffer = tirx.decl_buffer((m, n), "float32") + buffer = tirx.decl_tensor((m, n), "float32") actual = T.buffer_indices(buffer, k) for index, expected in zip(actual, [k // n, k % n]): tvm.ir.assert_structural_equal(index, expected) @@ -283,7 +283,7 @@ def test_symbolic_buffer_indices(): @pytest.mark.parametrize("layout", [None, TileLayout(S[(2, 3) : (1, 2)])]) def test_flat_buffer_store_preserves_identity_and_emits_once(layout): - buffer = tirx.decl_buffer((2, 3), "float32", strides=(5, 1), elem_offset=2, layout=layout) + buffer = tirx.decl_tensor((2, 3), "float32", strides=(5, 1), elem_offset=2, layout=layout) indices = T.buffer_indices(buffer, 4) load = buffer[indices] assert load.source.same_as(buffer) @@ -412,7 +412,7 @@ def test_gpu_imperative_buffers(kind): elif kind == "scan": from tvm.topi.gpu.scan import exclusive_scan_ir - body = exclusive_scan_ir(tirx.decl_buffer((2, 4)), tirx.decl_buffer((2, 4))) + body = exclusive_scan_ir(tirx.decl_tensor((2, 4)), tirx.decl_tensor((2, 4))) assert isinstance(body, tirx.Stmt) return elif kind == "scatter_nd": @@ -452,7 +452,7 @@ def test_concrete_mutable_scalar(declare): with IRBuilder() as ib: with T.prim_func(): if declare: - owner = T.alloc_buffer((4,), "int32", scope="local") + owner = T.alloc_tensor((4,), "int32", scope="local") scalar = T.decl_scalar("int32", owner.data, "local", elem_offset=2) else: scalar = T.alloc_scalar("int32", "local") @@ -471,7 +471,7 @@ def test_concrete_mutable_scalar(declare): assert all(store.buffer.same_as(scalar.source) for store in stores) if declare: declaration = next( - stmt for stmt in ib.get().body.seq if _is_buffer_binding(stmt, "tirx.decl_buffer") + stmt for stmt in ib.get().body.seq if _is_buffer_binding(stmt, "tirx.decl_tensor") ) assert declaration.var.same_as(scalar.source) tvm.ir.assert_structural_equal(declaration.value.args[0], owner.data) diff --git a/tests/python/tirx/script/test_tirx_script_jit_specialization.py b/tests/python/tirx/script/test_tirx_script_jit_specialization.py index 9e2cd65a7ad8..9ff38d01568a 100644 --- a/tests/python/tirx/script/test_tirx_script_jit_specialization.py +++ b/tests/python/tirx/script/test_tirx_script_jit_specialization.py @@ -24,11 +24,11 @@ def test_jit_buffer_annotation(): @T.jit(private=True) - def kernel(output: T.Buffer((5,), "int32")): + def kernel(output: T.Tensor((5,), "int32")): output[0] = 7 @T.prim_func(private=True) - def expected(output: T.Buffer((5,), "int32")): + def expected(output: T.Tensor((5,), "int32")): output[0] = 7 assert_structural_equal(kernel.specialize(), expected, map_free_vars=True) @@ -36,7 +36,7 @@ def expected(output: T.Buffer((5,), "int32")): def test_jit_optional_buffer(): @T.jit(private=True) - def kernel(value: T.Optional(T.Buffer((5,), "int32"))): + def kernel(value: T.Optional(T.Tensor((5,), "int32"))): if T.constexpr(value is not None): value[0] = 3 else: @@ -47,7 +47,7 @@ def absent(): T.evaluate(0) @T.prim_func(private=True) - def present(value: T.Buffer((5,), "int32")): + def present(value: T.Tensor((5,), "int32")): value[0] = 3 assert_structural_equal(kernel.specialize(value=None), absent, map_free_vars=True) diff --git a/tests/python/tirx/script/test_tirx_script_meta_programming.py b/tests/python/tirx/script/test_tirx_script_meta_programming.py index f362a2089900..c44cd2de3b67 100644 --- a/tests/python/tirx/script/test_tirx_script_meta_programming.py +++ b/tests/python/tirx/script/test_tirx_script_meta_programming.py @@ -33,7 +33,7 @@ def test_meta_class_constructor_rejects_unowned_resource(): @T.meta_class class Bad: def __init__(self): - tmp = T.alloc_buffer((1,), "int32", scope="local") + tmp = T.alloc_tensor((1,), "int32", scope="local") with pytest.raises(ValueError): @@ -50,14 +50,14 @@ def test_meta_class_multiple_instances_preserve_owned_resources(): class Holder: def __init__(self, external): self.external = external - self.buf = T.alloc_buffer((2,), "int32", scope="local") + self.buf = T.alloc_tensor((2,), "int32", scope="local") self.scalar = T.local_scalar("int32") instances.append(self) @T.prim_func(private=True) def test(): T.device_entry() - external = T.alloc_buffer((2,), "int32", scope="local") + external = T.alloc_tensor((2,), "int32", scope="local") first = Holder(external) second = Holder(external) T.evaluate( @@ -255,17 +255,17 @@ def test_prim_func_closure_shape(): def f(M=16): @T.prim_func - def func(A: T.Buffer((M,), "float32")): + def func(A: T.Tensor((M,), "float32")): T.evaluate(0) return func @T.prim_func - def expected_16(A: T.Buffer((16,), "float32")): + def expected_16(A: T.Tensor((16,), "float32")): T.evaluate(0) @T.prim_func - def expected_32(A: T.Buffer((32,), "float32")): + def expected_32(A: T.Tensor((32,), "float32")): T.evaluate(0) tvm.ir.assert_structural_equal(_normalize(f(16)), _normalize(expected_16)) @@ -282,17 +282,17 @@ def test_prim_func_closure_dtype(): def f(dtype="float32"): @T.prim_func - def func(A: T.Buffer((16,), dtype)): + def func(A: T.Tensor((16,), dtype)): T.evaluate(0) return func @T.prim_func - def expected_f32(A: T.Buffer((16,), "float32")): + def expected_f32(A: T.Tensor((16,), "float32")): T.evaluate(0) @T.prim_func - def expected_f16(A: T.Buffer((16,), "float16")): + def expected_f16(A: T.Tensor((16,), "float16")): T.evaluate(0) tvm.ir.assert_structural_equal(_normalize(f("float32")), _normalize(expected_f32)) @@ -311,7 +311,7 @@ def test_prim_func_nested_closure(): def outer(M=16): def middle(N=8): @T.prim_func - def func(A: T.Buffer((M, N), "float32")): + def func(A: T.Tensor((M, N), "float32")): T.evaluate(0) return func @@ -319,11 +319,11 @@ def func(A: T.Buffer((M, N), "float32")): return middle() @T.prim_func - def expected_16_8(A: T.Buffer((16, 8), "float32")): + def expected_16_8(A: T.Tensor((16, 8), "float32")): T.evaluate(0) @T.prim_func - def expected_32_8(A: T.Buffer((32, 8), "float32")): + def expected_32_8(A: T.Tensor((32, 8), "float32")): T.evaluate(0) tvm.ir.assert_structural_equal(_normalize(outer(16)), _normalize(expected_16_8)) @@ -337,17 +337,17 @@ def f(M=16): @I.ir_module class Mod: @T.prim_func - def main(A: T.Buffer((M,), "float32")): + def main(A: T.Tensor((M,), "float32")): T.evaluate(0) return Mod @T.prim_func - def expected_16(A: T.Buffer((16,), "float32")): + def expected_16(A: T.Tensor((16,), "float32")): T.evaluate(0) @T.prim_func - def expected_32(A: T.Buffer((32,), "float32")): + def expected_32(A: T.Tensor((32,), "float32")): T.evaluate(0) tvm.ir.assert_structural_equal(_normalize(f(16)["main"]), _normalize(expected_16)) @@ -359,17 +359,17 @@ def test_mixed_closure_usage(): def f(M=16): @T.prim_func - def func(A: T.Buffer((M,), "float32")): + def func(A: T.Tensor((M,), "float32")): T.evaluate(M) return func @T.prim_func - def expected_16(A: T.Buffer((16,), "float32")): + def expected_16(A: T.Tensor((16,), "float32")): T.evaluate(16) @T.prim_func - def expected_32(A: T.Buffer((32,), "float32")): + def expected_32(A: T.Tensor((32,), "float32")): T.evaluate(32) tvm.ir.assert_structural_equal(_normalize(f(16)), _normalize(expected_16)) diff --git a/tests/python/tirx/script/test_tirx_script_printer.py b/tests/python/tirx/script/test_tirx_script_printer.py index b55e9623230f..e352d4368ee1 100644 --- a/tests/python/tirx/script/test_tirx_script_printer.py +++ b/tests/python/tirx/script/test_tirx_script_printer.py @@ -100,10 +100,10 @@ def test_config_reserves_dialect_prefixes_before_variable_definition(prefixes): def test_buffer(): - a = tirx.decl_buffer((128, 128), "float16", name="A") + a = tirx.decl_tensor((128, 128), "float16", name="A") _assert_print( a, - """A = T.Var("A", T.Buffer((128, 128), "float16", layout="default")) + """A = T.Var("A", T.Tensor((128, 128), "float16", layout="default")) A""", ) @@ -113,7 +113,7 @@ def _assert_print(obj, expected): def test_buffer_region(): - src = tirx.decl_buffer((128, 128), "float32", name="src") + src = tirx.decl_tensor((128, 128), "float32", name="src") obj = tirx.BufferRegion( src, [ @@ -124,33 +124,33 @@ def test_buffer_region(): _assert_print( obj, """ -src = T.Var("src", T.Buffer((128, 128), "float32", layout="default")) +src = T.Var("src", T.Tensor((128, 128), "float32", layout="default")) src[64:128, 64:128] """, ) def test_buffer_load(): - a = tirx.decl_buffer((128, 128), "float16", name="A") + a = tirx.decl_tensor((128, 128), "float16", name="A") obj = tirx.BufferLoad(a, [128, 128]) _assert_print( obj, """ -A = T.Var("A", T.Buffer((128, 128), "float16", layout="default")) +A = T.Var("A", T.Tensor((128, 128), "float16", layout="default")) A[128, 128] """, ) def test_buffer_store(): - a = tirx.decl_buffer((128, 128), "float16", name="A") + a = tirx.decl_tensor((128, 128), "float16", name="A") with IRBuilder() as ib: TB.buffer_store(a, a[128, 128] + 1, [128, 128]) obj = ib.get() _assert_print( obj, """ -A = T.Var("A", T.Buffer((128, 128), "float16", layout="default")) +A = T.Var("A", T.Tensor((128, 128), "float16", layout="default")) A[128, 128] = A[128, 128] + T.float16(1.0) """, ) @@ -218,55 +218,56 @@ def test_while(): ) -def test_allocate(): +@pytest.mark.parametrize("declare", [TB.alloc_tensor, TB.decl_tensor]) +def test_allocate(declare): with IRBuilder() as ib: with TB.prim_func(): TB.func_name_("test") - buf = TB.alloc_buffer([128, 128], "float32") + buf = declare([128, 128], "float32") TB.evaluate(1) obj = ib.get() _assert_print( obj.body, """ -v = T.alloc_buffer((128, 128), "float32", layout="default") +v = T.alloc_tensor((128, 128), "float32", layout="default") T.evaluate(1) """, ) def test_allocate_with_decl_buffer_sugar(): - # AllocBuffer and DeclBuffer are flat siblings + # AllocTensor and DeclTensor are flat siblings with IRBuilder() as ib: with TB.prim_func(): TB.func_name_("test") - buf = TB.alloc_buffer([128, 128], "float32") - buf2 = TB.decl_buffer([128, 128], "float32", data=buf.data) + buf = TB.alloc_tensor([128, 128], "float32") + buf2 = TB.decl_tensor([128, 128], "float32", data=buf.data) TB.evaluate(1) obj = ib.get() _assert_print( obj.body, """ -v = T.alloc_buffer((128, 128), "float32", layout="default") -v_1 = T.decl_buffer((128, 128), "float32", data=v.data, layout="default") +v = T.alloc_tensor((128, 128), "float32", layout="default") +v_1 = T.decl_tensor((128, 128), "float32", data=v.data, layout="default") T.evaluate(1) """, ) def test_allocate_with_decl_buffer_sugar_multi_usage(): - # AllocBuffer and DeclBuffer are flat siblings + # AllocTensor and DeclTensor are flat siblings with IRBuilder() as ib: with TB.prim_func(): TB.func_name_("test") - buf = TB.alloc_buffer([128, 128], "float32") - buf2 = TB.decl_buffer([128, 128], "float32", data=buf.data) + buf = TB.alloc_tensor([128, 128], "float32") + buf2 = TB.decl_tensor([128, 128], "float32", data=buf.data) TB.evaluate(buf.data) obj = ib.get() _assert_print( obj.body, """ -v = T.alloc_buffer((128, 128), "float32", layout="default") -v_1 = T.decl_buffer((128, 128), "float32", data=v.data, layout="default") +v = T.alloc_tensor((128, 128), "float32", layout="default") +v_1 = T.decl_tensor((128, 128), "float32", data=v.data, layout="default") T.evaluate(v.data) """, ) @@ -276,26 +277,26 @@ def test_allocate_with_decl_buffer_no_sugar_mismatch(): with IRBuilder() as ib: with TB.prim_func(): TB.func_name_("test") - buf = TB.alloc_buffer([128, 128], "float32") - buf2 = TB.decl_buffer([256, 256], "float32", data=buf.data) + buf = TB.alloc_tensor([128, 128], "float32") + buf2 = TB.decl_tensor([256, 256], "float32", data=buf.data) TB.evaluate(buf.data) obj = ib.get() _assert_print( obj.body, """ -v = T.alloc_buffer((128, 128), "float32", layout="default") -v_1 = T.decl_buffer((256, 256), "float32", data=v.data, layout="default") +v = T.alloc_tensor((128, 128), "float32", layout="default") +v_1 = T.decl_tensor((256, 256), "float32", data=v.data, layout="default") T.evaluate(v.data) """, ) def test_decl_buffer(): - # DeclBuffer is flat: we need a frame to hold multiple stmts + # DeclTensor is flat: we need a frame to hold multiple stmts with IRBuilder() as ib: with TB.prim_func(): TB.func_name_("test") - buf = TB.decl_buffer((10, 10), data=TB.ptr("float32")) + buf = TB.decl_tensor((10, 10), data=TB.ptr("float32")) TB.evaluate(1) obj = ib.get() # Print only the body (skip PrimFunc wrapper) @@ -303,7 +304,7 @@ def test_decl_buffer(): obj.body, """ v_1: T.handle("float32", "global") = T.handle("float32", "global") -v = T.decl_buffer((10, 10), "float32", data=v_1, layout="default") +v = T.decl_tensor((10, 10), "float32", data=v_1, layout="default") T.evaluate(1) """, ) @@ -668,7 +669,7 @@ def test_print_kwargs_schedule_op_full_code(): # fmt: off @T.prim_func def test(): - A = T.alloc_buffer((16,), "float32") + A = T.alloc_tensor((16,), "float32") Tx.memset(A[0:16], T.float32(1.25), dispatch="v10", bar=7, foo=42) # fmt: on @@ -677,7 +678,7 @@ def test(): "\n" "@T.prim_func\n" "def test():\n" - ' A = T.alloc_buffer((16,), "float32", layout="default")\n' + ' A = T.alloc_tensor((16,), "float32", layout="default")\n' ' T.tile.memset(A[0:16], T.float32(1.25), dispatch="v10", bar=7, foo=42)' ) code = test.script() @@ -707,7 +708,7 @@ def _make_minimal_tirx_prim_func(): source = ( "# from tvm.script import tirx as T\n\n" "@T.prim_func()\n" - 'def f(A: T.Buffer((1,), "float32")):\n' + 'def f(A: T.Tensor((1,), "float32")):\n' " A[0] = T.float32(1)" ) return from_source(source) @@ -919,100 +920,100 @@ def test_printer_nvshmem_more(): def test_printer_nki_namespace(): - A = tir.decl_buffer([1], dtype="float16", name="A") - B = tir.decl_buffer([1], dtype="float16", name="B") + A = tir.decl_tensor([1], dtype="float16", name="A") + B = tir.decl_tensor([1], dtype="float16", name="B") a0 = A[0] b0 = B[0] _assert_namespace_print( trn_op.nki_load(a0, b0), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' - 'B = T.Var("B", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' + 'B = T.Var("B", T.Tensor((1,), "float16", layout="default"))\n' "T.nki.load(A[0], B[0])", ) _assert_namespace_print( trn_op.nki_store(a0, b0), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' - 'B = T.Var("B", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' + 'B = T.Var("B", T.Tensor((1,), "float16", layout="default"))\n' "T.nki.store(A[0], B[0])", ) _assert_namespace_print( trn_op.nki_tensor_copy(a0, b0), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' - 'B = T.Var("B", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' + 'B = T.Var("B", T.Tensor((1,), "float16", layout="default"))\n' "T.nki.tensor_copy(A[0], B[0])", ) _assert_namespace_print( trn_op.nki_matmul(a0, a0, b0), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' - 'B = T.Var("B", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' + 'B = T.Var("B", T.Tensor((1,), "float16", layout="default"))\n' "T.nki.matmul(A[0], A[0], B[0], T.bool(True))", ) _assert_namespace_print( trn_op.nki_activation(a0, b0, "relu", 0.0, 1.0), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' - 'B = T.Var("B", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' + 'B = T.Var("B", T.Tensor((1,), "float16", layout="default"))\n' 'T.nki.activation(A[0], B[0], "relu", ' "T.float32(0.0), T.float32(1.0))", ) _assert_namespace_print( trn_op.nki_memset(a0, 0), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\nT.nki.memset(A[0], 0)', + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\nT.nki.memset(A[0], 0)', ) _assert_namespace_print( trn_op.nki_identity(a0, 1), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\nT.nki.identity(A[0], 1)', + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\nT.nki.identity(A[0], 1)', ) _assert_namespace_print( trn_op.nki_reciprocal(a0, b0), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' - 'B = T.Var("B", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' + 'B = T.Var("B", T.Tensor((1,), "float16", layout="default"))\n' "T.nki.reciprocal(A[0], B[0])", ) _assert_namespace_print( trn_op.nki_tensorreduce(a0, b0, "sum", False, 0), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' - 'B = T.Var("B", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' + 'B = T.Var("B", T.Tensor((1,), "float16", layout="default"))\n' 'T.nki.tensorreduce(A[0], B[0], "sum", ' "T.bool(False), 0)", ) _assert_namespace_print( trn_op.nki_tensortensor(a0, a0, b0, "add"), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' - 'B = T.Var("B", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' + 'B = T.Var("B", T.Tensor((1,), "float16", layout="default"))\n' 'T.nki.tensortensor(A[0], A[0], B[0], "add")', ) _assert_namespace_print( trn_op.nki_tensorscalar(a0, a0, 1.0, "mul", False), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' "T.nki.tensorscalar(A[0], A[0], T.float32(1.0), " '"mul", T.bool(False))', ) _assert_namespace_print( trn_op.nki_tensorscalar_reduce(a0, a0, 1.0, "mul", "sum", False), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' 'T.nki.tensorscalar_reduce(A[0], A[0], T.float32(1.0), "mul", "sum", T.bool(False), T.bool(False))', # noqa: E501 ) _assert_namespace_print( trn_op.nki_scalar_tensor_tensor(a0, a0, 1.0, a0, "add", "add"), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' 'T.nki.scalar_tensor_tensor(A[0], A[0], T.float32(1.0), A[0], "add", "add", T.bool(False), T.bool(False))', # noqa: E501 ) _assert_namespace_print( trn_op.nki_scalar_tensor_scalar(a0, a0, 1.0, 1.0, "add", "add"), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' 'T.nki.scalar_tensor_scalar(A[0], A[0], T.float32(1.0), T.float32(1.0), "add", "add", T.bool(False), T.bool(False))', # noqa: E501 ) _assert_namespace_print( trn_op.nki_activation_reduce(a0, a0, b0, "relu", "sum", 0.0, 1.0), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' - 'B = T.Var("B", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' + 'B = T.Var("B", T.Tensor((1,), "float16", layout="default"))\n' "T.nki.activation_reduce(A[0], A[0], B[0], " '"relu", "sum", T.float32(0.0), ' "T.float32(1.0))", ) _assert_namespace_print( trn_op.nki_affine_select(a0, a0, a0, 1.0), - 'A = T.Var("A", T.Buffer((1,), "float16", layout="default"))\n' + 'A = T.Var("A", T.Tensor((1,), "float16", layout="default"))\n' "T.nki.affine_select(A[0], A[0], A[0], T.float32(1.0))", ) diff --git a/tests/python/tirx/script/test_tirx_script_printer_highlight.py b/tests/python/tirx/script/test_tirx_script_printer_highlight.py index 13948ce705ae..e2008e3a6ce6 100644 --- a/tests/python/tirx/script/test_tirx_script_printer_highlight.py +++ b/tests/python/tirx/script/test_tirx_script_printer_highlight.py @@ -27,9 +27,9 @@ def test_highlight_script(): class Module: @T.prim_func def main( # type: ignore - A: T.Buffer([16, 128, 128]), - B: T.Buffer([16, 128, 128]), - C: T.Buffer([16, 128, 128]), + A: T.Tensor([16, 128, 128]), + B: T.Tensor([16, 128, 128]), + C: T.Tensor([16, 128, 128]), ) -> None: # pylint: disable=no-self-argument T.func_attr({"global_symbol": "main", "tirx.noalias": True}) for n, i, j in T.grid(16, 128, 128): diff --git a/tests/python/tirx/script/test_tirx_script_printer_structural_equal.py b/tests/python/tirx/script/test_tirx_script_printer_structural_equal.py index 9a178a5a6aeb..abe2c378378a 100644 --- a/tests/python/tirx/script/test_tirx_script_printer_structural_equal.py +++ b/tests/python/tirx/script/test_tirx_script_printer_structural_equal.py @@ -46,11 +46,11 @@ def test_prim_type_hidden_path_exact_message(): def test_prim_func_buffer_param(): @T.prim_func - def func1(A: T.Buffer((128, 128)), B: T.Buffer((128, 128))): + def func1(A: T.Tensor((128, 128)), B: T.Tensor((128, 128))): pass @T.prim_func - def func2(A: T.Buffer((128, 128)), B: T.Buffer((128, 256))): + def func2(A: T.Tensor((128, 128)), B: T.Tensor((128, 256))): pass func1 = func1.with_attr("global_symbol", "main") @@ -125,11 +125,11 @@ def func(): def test_allocate(): @T.prim_func def func1(): - a = T.alloc_buffer((128, 128), dtype="float32") + a = T.alloc_tensor((128, 128), dtype="float32") @T.prim_func def func2(): - a = T.alloc_buffer((256, 128), dtype="float32") + a = T.alloc_tensor((256, 128), dtype="float32") func1 = func1.with_attr("global_symbol", "main") func2 = func2.with_attr("global_symbol", "main") diff --git a/tests/python/tirx/script/test_tirx_script_source_locations.py b/tests/python/tirx/script/test_tirx_script_source_locations.py index 642676e4aba8..6e87386f3b84 100644 --- a/tests/python/tirx/script/test_tirx_script_source_locations.py +++ b/tests/python/tirx/script/test_tirx_script_source_locations.py @@ -45,7 +45,7 @@ def test_parser_attaches_span_to_direct_call(): @_capture_source(sources) def direct_call(): T.device_entry() - barriers = T.alloc_buffer((1,), "uint64", scope="shared") + barriers = T.alloc_tensor((1,), "uint64", scope="shared") T.cuda.mbarrier_wait( T.address_of(barriers[0]), 0, @@ -98,8 +98,8 @@ def test_parser_attaches_span_to_nested_tensor_load(): @T.prim_func @_capture_source(sources) def nested_load(): - source_buffer = T.alloc_buffer((1,), "int32") - output = T.alloc_buffer((1,), "int32") + source_buffer = T.alloc_tensor((1,), "int32") + output = T.alloc_tensor((1,), "int32") output[0] = source_buffer[0] + 1 source = sources[0] @@ -130,7 +130,7 @@ def wait_impl(barrier): @_capture_source(sources) def inline_call(): T.device_entry() - barriers = T.alloc_buffer((1,), "uint64", scope="shared") + barriers = T.alloc_tensor((1,), "uint64", scope="shared") wait(T.address_of(barriers[0])) caller_source = sources[0] @@ -156,7 +156,7 @@ def test_parser_attaches_span_to_tile_primitive_call(): @T.prim_func @_capture_source(sources) def tile_call(): - A = T.alloc_buffer((16,), "float32") + A = T.alloc_tensor((16,), "float32") Tx.memset(A[0:16], T.float32(0)) source = sources[0] @@ -186,7 +186,7 @@ def test_statement_receipts_keep_emitted_nodes_and_spans(): with I.IRBuilder(): with T.function_(private=True) as frame: T.func_name_("receipts") - output = T.arg_("output", T.Buffer((1,), "int32")) + output = T.arg_("output", T.Tensor((1,), "int32")) stored = T.setitem_(value=3, target=output, key=0, span=location) holder = SimpleNamespace(value=output) updated = T.setattr_(holder, "value", 4, span=location) @@ -267,7 +267,7 @@ def observe(value): monkeypatch.setattr(ir.Var, "view", view) @T.prim_func - def main(A: T.Buffer((4, 4), "float32")): + def main(A: T.Tensor((4, 4), "float32")): renamed = A.view(mark()) alias = renamed observe(renamed) @@ -323,7 +323,7 @@ def observe(*items): observed.extend(items) @T.prim_func - def main(A: T.Buffer((4,), "float32")): + def main(A: T.Tensor((4,), "float32")): initialize(A) renamed_layout = make(0) renamed_holder = make(1) diff --git a/tests/python/tirx/test_buffer_data_reinfer_type.py b/tests/python/tirx/test_buffer_data_reinfer_type.py index 45439ec59db8..2c6b70147340 100644 --- a/tests/python/tirx/test_buffer_data_reinfer_type.py +++ b/tests/python/tirx/test_buffer_data_reinfer_type.py @@ -22,8 +22,8 @@ def test_buffer_data_reinfer_type_from_rewritten_argument(): - global_buffer = tirx.decl_buffer((8,), "float32", name="global_buffer", scope="global") - local_buffer = tirx.decl_buffer((8,), "float16", name="local_buffer", scope="local") + global_buffer = tirx.decl_tensor((8,), "float32", name="global_buffer", scope="global") + local_buffer = tirx.decl_tensor((8,), "float16", name="local_buffer", scope="local") stale_type = tirx.buffer_data_pointer_type(global_buffer) expected_type = tirx.buffer_data_pointer_type(local_buffer) diff --git a/tests/python/tirx/test_buffer_print.py b/tests/python/tirx/test_buffer_print.py index 691823e8f691..867fbf14b560 100644 --- a/tests/python/tirx/test_buffer_print.py +++ b/tests/python/tirx/test_buffer_print.py @@ -199,7 +199,7 @@ def test_vector_add_1D(dtype, dtype_str): @T.prim_func def add_func( - A: T.Buffer((M,), dtype_str), B: T.Buffer((M,), dtype_str), C: T.Buffer((M,), dtype_str) + A: T.Tensor((M,), dtype_str), B: T.Tensor((M,), dtype_str), C: T.Tensor((M,), dtype_str) ) -> None: for i in T.thread_binding(M, thread="threadIdx.x"): C[i] = A[i] + B[i] @@ -219,9 +219,9 @@ def test_vector_add_2D(dtype, dtype_str): @T.prim_func def add_func( - A: T.Buffer((M, N), dtype_str), - B: T.Buffer((M, N), dtype_str), - C: T.Buffer((M, N), dtype_str), + A: T.Tensor((M, N), dtype_str), + B: T.Tensor((M, N), dtype_str), + C: T.Tensor((M, N), dtype_str), ) -> None: for i in T.thread_binding(M, thread="threadIdx.x"): for j in T.thread_binding(N, thread="threadIdx.y"): @@ -242,9 +242,9 @@ def test_vector_add_3D(dtype, dtype_str): @T.prim_func def add_func( - A: T.Buffer((M, N, K), dtype_str), - B: T.Buffer((M, N, K), dtype_str), - C: T.Buffer((M, N, K), dtype_str), + A: T.Tensor((M, N, K), dtype_str), + B: T.Tensor((M, N, K), dtype_str), + C: T.Tensor((M, N, K), dtype_str), ) -> None: for i in T.thread_binding(M, thread="threadIdx.x"): for j in T.thread_binding(N, thread="threadIdx.y"): @@ -266,7 +266,7 @@ def test_const_scalar(dtype, dtype_str): @T.prim_func def add_func( - A: T.Buffer((M,), dtype_str), B: T.Buffer((M,), dtype_str), C: T.Buffer((M,), dtype_str) + A: T.Tensor((M,), dtype_str), B: T.Tensor((M,), dtype_str), C: T.Tensor((M,), dtype_str) ) -> None: Ten: T.let = T.IntImm(dtype_str, 10) @@ -288,7 +288,7 @@ def test_string(dtype, dtype_str, test_string): @T.prim_func def add_func( - A: T.Buffer((M,), dtype_str), B: T.Buffer((M,), dtype_str), C: T.Buffer((M,), dtype_str) + A: T.Tensor((M,), dtype_str), B: T.Tensor((M,), dtype_str), C: T.Tensor((M,), dtype_str) ) -> None: string_var = tvm.ir.StringImm(test_string) diff --git a/tests/python/tirx/test_compilation_pipeline.py b/tests/python/tirx/test_compilation_pipeline.py index b4a5b33d26d3..ef45fe22e437 100644 --- a/tests/python/tirx/test_compilation_pipeline.py +++ b/tests/python/tirx/test_compilation_pipeline.py @@ -29,8 +29,8 @@ @pytest.mark.parametrize("pipeline", [None, "default", "tirx"]) def test_default_pipeline_allocations(pipeline): @T.prim_func - def add_one(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): - temp = T.alloc_buffer((16,), "float32") + def add_one(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): + temp = T.alloc_tensor((16,), "float32") for i in range(16): temp[i] = A[i] + T.float32(1) for i in range(16): @@ -46,7 +46,7 @@ def add_one(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): def test_unify_thread_binding(): @T.prim_func - def before(A: T.Buffer((32,), "int32")): + def before(A: T.Tensor((32,), "int32")): for bx in T.thread_binding(1, thread="blockIdx.x"): for tx in T.thread_binding(32, thread="threadIdx.x"): A[tx] = tx @@ -54,7 +54,7 @@ def before(A: T.Buffer((32,), "int32")): A[tx] = A[tx] + 1 @T.prim_func - def expected(A: T.Buffer((32,), "int32")): + def expected(A: T.Tensor((32,), "int32")): for bx in T.thread_binding(1, thread="blockIdx.x"): for tx in T.thread_binding(32, thread="threadIdx.x"): A[tx] = tx diff --git a/tests/python/tirx/test_control_flow.py b/tests/python/tirx/test_control_flow.py index c449f5d3b329..6e3e00dc33fa 100644 --- a/tests/python/tirx/test_control_flow.py +++ b/tests/python/tirx/test_control_flow.py @@ -43,7 +43,7 @@ def run_and_check(): def test_break_continue1(): # fmt: off @T.prim_func - def func(A: T.Buffer((10,), 'int32')): + def func(A: T.Tensor((10,), 'int32')): T.device_entry() cta_id = T.cta_id([1]) @@ -65,12 +65,12 @@ def func(A: T.Buffer((10,), 'int32')): def test_break_continue2(): # fmt: off @T.prim_func - def func(A: T.Buffer((9,), 'int32')): + def func(A: T.Tensor((9,), 'int32')): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([32]) - idx = T.alloc_buffer((1,), "int32", scope="local") + idx = T.alloc_tensor((1,), "int32", scope="local") idx[0] = 0 for i in T.serial(3): if i == 0: @@ -92,12 +92,12 @@ def func(A: T.Buffer((9,), 'int32')): def test_break_continue3(): # fmt: off @T.prim_func - def func(A: T.Buffer((10,), 'int32')): + def func(A: T.Tensor((10,), 'int32')): T.device_entry() cta_id = T.cta_id([1]) tid = T.thread_id([32]) - i = T.alloc_buffer((1,), "int32", scope="local") + i = T.alloc_tensor((1,), "int32", scope="local") i[0] = 0 while i[0] < 10: if (i[0] % 2) == 1: diff --git a/tests/python/tirx/test_hint.py b/tests/python/tirx/test_hint.py index 25840ef0e8e1..46b4322e9cc6 100644 --- a/tests/python/tirx/test_hint.py +++ b/tests/python/tirx/test_hint.py @@ -29,11 +29,11 @@ def from_source(code): def test_hint_keyword_arg_on_tx_op(): """Tx.op(..., hint="msg") stores hint in TilePrimitiveCall.config.""" - from tvm.tirx.buffer import decl_buffer + from tvm.tirx.buffer import decl_tensor from tvm.tirx.tile_primitive import TilePrimitiveCall - A = decl_buffer((64, 64), "float32", scope="global") - A_sm = decl_buffer((64, 64), "float32", scope="shared") + A = decl_tensor((64, 64), "float32", scope="global") + A_sm = decl_tensor((64, 64), "float32", scope="shared") op_call = TilePrimitiveCall( A[0:64, 0:64], @@ -52,7 +52,7 @@ def test_hint_keyword_arg_on_tx_op_roundtrip(): @T.prim_func def func( - A: T.Buffer([10], "float32", scope="global"), B: T.Buffer([10], "float32", scope="global") + A: T.Tensor([10], "float32", scope="global"), B: T.Tensor([10], "float32", scope="global") ): Tx.add(B, A, T.float32(1), hint="use_fast_math") diff --git a/tests/python/tirx/test_inline.py b/tests/python/tirx/test_inline.py index 4e3c8f8aa2ab..380ac8a66217 100644 --- a/tests/python/tirx/test_inline.py +++ b/tests/python/tirx/test_inline.py @@ -27,7 +27,7 @@ def test_local_shadows_enclosing(): """A local parameter in the inline shadows a variable from the enclosing scope.""" @T.prim_func(private=True) - def func(A: T.Buffer((128,), "int32")) -> None: + def func(A: T.Tensor((128,), "int32")) -> None: T.int32(10) @T.inline @@ -38,7 +38,7 @@ def write(x): write(T.int32(20)) @T.prim_func(private=True) - def expected(A: T.Buffer((128,), "int32")) -> None: + def expected(A: T.Tensor((128,), "int32")) -> None: T.int32(10) A[0] = T.int32(20) @@ -54,11 +54,11 @@ def write_val(A): A[0] = val @T.prim_func(private=True) - def func(A: T.Buffer((128,), "int32")) -> None: + def func(A: T.Tensor((128,), "int32")) -> None: write_val(A) @T.prim_func(private=True) - def expected(A: T.Buffer((128,), "int32")) -> None: + def expected(A: T.Tensor((128,), "int32")) -> None: A[0] = 64 assert_structural_equal(func, expected) @@ -77,11 +77,11 @@ def add_two(A): add_one(A) @T.prim_func(private=True) - def func(A: T.Buffer((128,), "int32")) -> None: + def func(A: T.Tensor((128,), "int32")) -> None: add_two(A) @T.prim_func(private=True) - def expected(A: T.Buffer((128,), "int32")) -> None: + def expected(A: T.Tensor((128,), "int32")) -> None: A[0] = A[0] + 1 A[0] = A[0] + 1 @@ -96,11 +96,11 @@ def write_const(A): A[0] = MODULE_CONST @T.prim_func(private=True) - def func(A: T.Buffer((128,), "int32")) -> None: + def func(A: T.Tensor((128,), "int32")) -> None: write_const(A) @T.prim_func(private=True) - def expected(A: T.Buffer((128,), "int32")) -> None: + def expected(A: T.Tensor((128,), "int32")) -> None: A[0] = 42 assert_structural_equal(func, expected) @@ -110,7 +110,7 @@ def test_shadowing_in_inner_scope(): """An inline defined inside a for-loop captures the loop variable.""" @T.prim_func(private=True) - def func(A: T.Buffer((10,), "int32")) -> None: + def func(A: T.Tensor((10,), "int32")) -> None: for i in T.serial(10): @T.inline @@ -120,7 +120,7 @@ def write_i(A): write_i(A) @T.prim_func(private=True) - def expected(A: T.Buffer((10,), "int32")) -> None: + def expected(A: T.Tensor((10,), "int32")) -> None: for i in range(10): A[i] = i @@ -138,12 +138,12 @@ def static_capture(A, B): B[()] = A[x_value] @T.prim_func(private=True) - def func(A: T.Buffer((1024,), "int32"), B: T.Buffer((), "int32")) -> None: + def func(A: T.Tensor((1024,), "int32"), B: T.Tensor((), "int32")) -> None: for x_value in T.serial(10): static_capture(A, B) @T.prim_func(private=True) - def expected(A: T.Buffer((1024,), "int32"), B: T.Buffer((), "int32")) -> None: + def expected(A: T.Tensor((1024,), "int32"), B: T.Tensor((), "int32")) -> None: for x_value in range(10): B[()] = A[128] @@ -162,11 +162,11 @@ def inc(A): A[0] = A[0] + 1 @T.prim_func(private=True) - def func(A: T.Buffer((128,), "int32")) -> None: + def func(A: T.Tensor((128,), "int32")) -> None: apply_fn(inc, A) @T.prim_func(private=True) - def expected(A: T.Buffer((128,), "int32")) -> None: + def expected(A: T.Tensor((128,), "int32")) -> None: A[0] = A[0] + 1 assert_structural_equal(func, expected) @@ -184,12 +184,12 @@ def write_b(A): A[1] = 2 @T.prim_func(private=True) - def func(A: T.Buffer((128,), "int32")) -> None: + def func(A: T.Tensor((128,), "int32")) -> None: write_a(A) write_b(A) @T.prim_func(private=True) - def expected(A: T.Buffer((128,), "int32")) -> None: + def expected(A: T.Tensor((128,), "int32")) -> None: A[0] = 1 A[1] = 2 @@ -230,7 +230,7 @@ def test_late_binding(): """Variable defined after inline but before call (inside prim_func).""" @T.prim_func(private=True) - def func(A: T.Buffer((128,), "int32")) -> None: + def func(A: T.Tensor((128,), "int32")) -> None: @T.inline def write(A): A[0] = val @@ -239,7 +239,7 @@ def write(A): write(A) @T.prim_func(private=True) - def expected(A: T.Buffer((128,), "int32")) -> None: + def expected(A: T.Tensor((128,), "int32")) -> None: val = T.int32(99) A[0] = val diff --git a/tests/python/tirx/test_jit.py b/tests/python/tirx/test_jit.py index b36a784b3088..530651ffd0eb 100644 --- a/tests/python/tirx/test_jit.py +++ b/tests/python/tirx/test_jit.py @@ -33,9 +33,9 @@ def test_int_constexpr_specializes_loop_bound(): @T.jit(private=True) def add( - A: T.Buffer((N,), "int32"), - B: T.Buffer((N,), "int32"), - C: T.Buffer((N,), "int32"), + A: T.Tensor((N,), "int32"), + B: T.Tensor((N,), "int32"), + C: T.Tensor((N,), "int32"), *, N: T.constexpr, ): @@ -44,9 +44,9 @@ def add( @T.prim_func(private=True) def expected( - A: T.Buffer((128,), "int32"), - B: T.Buffer((128,), "int32"), - C: T.Buffer((128,), "int32"), + A: T.Tensor((128,), "int32"), + B: T.Tensor((128,), "int32"), + C: T.Tensor((128,), "int32"), ): for i in range(128): C[i] = A[i] + B[i] @@ -57,9 +57,9 @@ def expected( def test_constexpr_in_2d_buffer_shape(): @T.jit(private=True) def matadd( - A: T.Buffer((M, K), "int32"), - B: T.Buffer((M, K), "int32"), - C: T.Buffer((M, K), "int32"), + A: T.Tensor((M, K), "int32"), + B: T.Tensor((M, K), "int32"), + C: T.Tensor((M, K), "int32"), *, M: T.constexpr, K: T.constexpr, @@ -70,9 +70,9 @@ def matadd( @T.prim_func(private=True) def expected( - A: T.Buffer((4, 8), "int32"), - B: T.Buffer((4, 8), "int32"), - C: T.Buffer((4, 8), "int32"), + A: T.Tensor((4, 8), "int32"), + B: T.Tensor((4, 8), "int32"), + C: T.Tensor((4, 8), "int32"), ): for m in range(4): for k in range(8): @@ -84,8 +84,8 @@ def expected( def test_constexpr_in_body_expression(): @T.jit(private=True) def scaled_copy( - A: T.Buffer((N,), "int32"), - B: T.Buffer((N,), "int32"), + A: T.Tensor((N,), "int32"), + B: T.Tensor((N,), "int32"), *, N: T.constexpr, SCALE: T.constexpr, @@ -95,8 +95,8 @@ def scaled_copy( @T.prim_func(private=True) def expected( - A: T.Buffer((16,), "int32"), - B: T.Buffer((16,), "int32"), + A: T.Tensor((16,), "int32"), + B: T.Tensor((16,), "int32"), ): for i in range(16): B[i] = A[i] * 3 @@ -107,7 +107,7 @@ def expected( def test_specialize_cache_returns_same_instance(): @T.jit(private=True) def k( - A: T.Buffer((N,), "int32"), + A: T.Tensor((N,), "int32"), *, N: T.constexpr, ): @@ -122,7 +122,7 @@ def k( def test_specialize_different_args_produce_different_funcs(): @T.jit(private=True) def k( - A: T.Buffer((N,), "int32"), + A: T.Tensor((N,), "int32"), *, N: T.constexpr, ): @@ -135,7 +135,7 @@ def k( def test_specialize_missing_constexpr_raises(): @T.jit(private=True) def k( - A: T.Buffer((N,), "int32"), + A: T.Tensor((N,), "int32"), *, N: T.constexpr, SCALE: T.constexpr, @@ -150,7 +150,7 @@ def k( def test_specialize_extra_kwarg_raises(): @T.jit(private=True) def k( - A: T.Buffer((N,), "int32"), + A: T.Tensor((N,), "int32"), *, N: T.constexpr, ): @@ -164,7 +164,7 @@ def k( def test_jit_kernel_with_nested_inline_helper(): @T.jit(private=True) def k( - A: T.Buffer((N,), "int32"), + A: T.Tensor((N,), "int32"), *, N: T.constexpr, ): @@ -177,7 +177,7 @@ def double(x): @T.prim_func(private=True) def expected( - A: T.Buffer((4,), "int32"), + A: T.Tensor((4,), "int32"), ): for i in range(4): A[i] = A[i] * 2 @@ -188,7 +188,7 @@ def expected( def test_constexpr_default_value(): @T.jit(private=True) def k( - A: T.Buffer((N,), "int32"), + A: T.Tensor((N,), "int32"), *, N: T.constexpr, SCALE: T.constexpr = 7, @@ -198,7 +198,7 @@ def k( @T.prim_func(private=True) def expected( - A: T.Buffer((8,), "int32"), + A: T.Tensor((8,), "int32"), ): for i in range(8): A[i] = 7 @@ -212,7 +212,7 @@ def expected( def test_specialize_returns_primfunc(): @T.jit(private=True) def k( - A: T.Buffer((N,), "int32"), + A: T.Tensor((N,), "int32"), *, N: T.constexpr, ): @@ -228,9 +228,9 @@ def k( def test_constexpr_specializes_nested_selector_condition(): @T.jit(private=True) def k( - A: T.Buffer((8,), "float16"), - B: T.Buffer((8,), "float16"), - C: T.Buffer((8,), "float16"), + A: T.Tensor((8,), "float16"), + B: T.Tensor((8,), "float16"), + C: T.Tensor((8,), "float16"), flag: T.int32, *, LIMIT: T.constexpr, @@ -253,18 +253,18 @@ def k( def test_optional_param_present_and_absent_ir(): @T.jit(private=True) - def kernel(a: T.Optional(T.Buffer((1,), "int32")), out: T.Buffer((1,), "int32")): + def kernel(a: T.Optional(T.Tensor((1,), "int32")), out: T.Tensor((1,), "int32")): if T.constexpr(a is not None): out[0] = a[0] else: out[0] = -1 @T.prim_func(private=True) - def expected_present(A: T.Buffer((1,), "int32"), out: T.Buffer((1,), "int32")): + def expected_present(A: T.Tensor((1,), "int32"), out: T.Tensor((1,), "int32")): out[0] = A[0] @T.prim_func(private=True) - def expected_absent(out: T.Buffer((1,), "int32")): + def expected_absent(out: T.Tensor((1,), "int32")): out[0] = -1 present = kernel.specialize() @@ -279,7 +279,7 @@ def expected_absent(out: T.Buffer((1,), "int32")): def test_optional_specialization_cache_includes_presence(): @T.jit(private=True) - def kernel(a: T.Optional(T.Buffer((1,), "int32")), out: T.Buffer((1,), "int32")): + def kernel(a: T.Optional(T.Tensor((1,), "int32")), out: T.Tensor((1,), "int32")): if T.constexpr(a is not None): out[0] = a[0] else: @@ -295,11 +295,11 @@ def kernel(a: T.Optional(T.Buffer((1,), "int32")), out: T.Buffer((1,), "int32")) def test_multiple_optional_params_preserve_runtime_order(): @T.jit(private=True) def kernel( - first: T.Buffer((1,), "int32"), - a: T.Optional(T.Buffer((1,), "int32")), + first: T.Tensor((1,), "int32"), + a: T.Optional(T.Tensor((1,), "int32")), scale: T.int32, - b: T.Optional(T.Buffer((1,), "int32")), - out: T.Buffer((1,), "int32"), + b: T.Optional(T.Tensor((1,), "int32")), + out: T.Tensor((1,), "int32"), ): out[0] = first[0] * scale if T.constexpr(a is not None): @@ -329,7 +329,7 @@ def kernel( def test_optional_only_accepts_none_at_specialization_time(): @T.jit(private=True) - def kernel(a: T.Optional(T.Buffer((1,), "int32")), out_h: T.Buffer((1,), "int32")): + def kernel(a: T.Optional(T.Tensor((1,), "int32")), out_h: T.Tensor((1,), "int32")): if T.constexpr(a is not None): T.evaluate(a[0]) T.evaluate(out_h[0]) @@ -363,7 +363,7 @@ def invalid(a: T.Optional(T.handle)): def test_compile_time_if_binding_uses_python_scope(): @T.jit(private=True) - def kernel(a: T.Optional(T.Buffer((1,), "int32")), out_h: T.Buffer((1,), "int32")): + def kernel(a: T.Optional(T.Tensor((1,), "int32")), out_h: T.Tensor((1,), "int32")): if T.constexpr(a is None): selected = out_h else: @@ -383,7 +383,7 @@ def fail_if_evaluated(): raise RuntimeError("dead expression was evaluated") @T.jit(private=True) - def kernel(a: T.Optional(T.handle), out: T.Buffer((1,), "int32")): + def kernel(a: T.Optional(T.handle), out: T.Tensor((1,), "int32")): if T.constexpr(a is None or fail_if_evaluated()): out[0] = 1 if T.constexpr(a is not None and fail_if_evaluated()): @@ -396,7 +396,7 @@ def kernel(a: T.Optional(T.handle), out: T.Buffer((1,), "int32")): def test_runtime_tir_if_cannot_guard_absent_optional_param(): @T.jit(private=True) - def kernel(a: T.Optional(T.Buffer((1,), "int32")), flag: T.int32): + def kernel(a: T.Optional(T.Tensor((1,), "int32")), flag: T.int32): if flag != 0: T.evaluate(a[0]) @@ -413,7 +413,7 @@ def kernel(a: T.Optional(T.Buffer((1,), "int32")), flag: T.int32): ) def test_unguarded_absent_optional_param_reports_source(operation, source_text, error_type): @T.jit(private=True) - def kernel(a: T.Optional(T.Buffer((1,), "int32"))): + def kernel(a: T.Optional(T.Tensor((1,), "int32"))): if T.constexpr(operation == "subscript"): a[10] elif T.constexpr(operation == "attribute"): @@ -428,7 +428,7 @@ def kernel(a: T.Optional(T.Buffer((1,), "int32"))): def test_present_optional_param_still_rejects_ffi_none(): @T.jit - def kernel(A: T.Buffer((1,), "int32")): + def kernel(A: T.Tensor((1,), "int32")): A[0] = 0 executable = tvm.compile(kernel.specialize(), target="llvm", tir_pipeline="tirx") diff --git a/tests/python/tirx/test_op.py b/tests/python/tirx/test_op.py index b9bfead5f9ac..db8068c568d6 100644 --- a/tests/python/tirx/test_op.py +++ b/tests/python/tirx/test_op.py @@ -19,7 +19,7 @@ import pytest from tvm.ir import Op, assert_structural_equal -from tvm.tirx.buffer import decl_buffer +from tvm.tirx.buffer import decl_tensor from tvm.tirx.exec_scope import ExecScope from tvm.tirx.tile_primitive import TilePrimitiveCall @@ -29,28 +29,28 @@ def _test(op: str, *args): def test_copy(): - A = decl_buffer((64, 64), "float32", scope="global") - A_sm = decl_buffer((64, 64), "float32", scope="shared") + A = decl_tensor((64, 64), "float32", scope="global") + A_sm = decl_tensor((64, 64), "float32", scope="shared") _test("copy", A[0:64, 0:64], A_sm[0:64, 0:64]) def test_fill(): - A = decl_buffer((64, 64), "float32", scope="global") + A = decl_tensor((64, 64), "float32", scope="global") _test("fill", A[0:64, 0:64], 1.0) def test_gemm(): - A = decl_buffer((64, 64), "float32", scope="global") - B = decl_buffer((64, 64), "float32", scope="global") - C = decl_buffer((64, 64), "float32", scope="global") - D = decl_buffer((64, 64), "float32", scope="global") + A = decl_tensor((64, 64), "float32", scope="global") + B = decl_tensor((64, 64), "float32", scope="global") + C = decl_tensor((64, 64), "float32", scope="global") + D = decl_tensor((64, 64), "float32", scope="global") _test("gemm", D[:, :], A[:, :], B[:, :], C[:, :], True, False, 1.0, 0.0) def test_tile_primitive_call_pickle_roundtrip(): """TilePrimitiveCall reflection must provide a deserialization creator.""" - A = decl_buffer((64,), "float32", scope="local") - workspace = decl_buffer((16,), "float32", scope="shared") + A = decl_tensor((64,), "float32", scope="local") + workspace = decl_tensor((16,), "float32", scope="shared") call = TilePrimitiveCall( A[:], 1.0, @@ -79,8 +79,8 @@ def test_buffer_replacer_no_shared_default(): r1 = BufferReplacer() r2 = BufferReplacer() - A = decl_buffer((64,), "float32") - B = decl_buffer((64,), "float32") + A = decl_tensor((64,), "float32") + B = decl_tensor((64,), "float32") r1.buffer_map[A] = B # r2 must not see r1's mutation assert len(r2.buffer_map) == 0 @@ -93,7 +93,7 @@ def test_buffer_replacer_replaces_strides_and_elem_offset(): n = Var("n", "int32") m = Var("m", "int32") - A = decl_buffer((64,), "float32", strides=[n], elem_offset=n) + A = decl_tensor((64,), "float32", strides=[n], elem_offset=n) store = BufferStore(A, 1.0, [0]) new = BufferReplacer(var_map={n: m})(store) @@ -105,10 +105,10 @@ def test_gemm_async_partial_scale_factor(): """Regression test for F7: gemm_async must reject partial scale factors.""" from tvm.tirx.script.ir_builder.tirx import gemm_async - A = decl_buffer((64, 64), "float16", scope="shared") - B = decl_buffer((64, 64), "float16", scope="shared") - C = decl_buffer((64, 64), "float16", scope="shared") - SF = decl_buffer((64,), "float16", scope="shared") + A = decl_tensor((64, 64), "float16", scope="shared") + B = decl_tensor((64, 64), "float16", scope="shared") + C = decl_tensor((64, 64), "float16", scope="shared") + SF = decl_tensor((64,), "float16", scope="shared") with pytest.raises(ValueError, match="SFA and SFB must both be provided or both be None"): gemm_async(C[:, :], A[:, :], B[:, :], SFA=SF[:]) diff --git a/tests/python/tirx/test_op_namespace_cleanup.py b/tests/python/tirx/test_op_namespace_cleanup.py index c2337421b5fb..ba6d8d6f60a4 100644 --- a/tests/python/tirx/test_op_namespace_cleanup.py +++ b/tests/python/tirx/test_op_namespace_cleanup.py @@ -130,7 +130,7 @@ def marker(): def test_tile_shorthand_and_scoped_aliases_use_tile_ops(): @T.prim_func(check_well_formed=False) - def tile_aliases(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): + def tile_aliases(A: T.Tensor((16,), "float32"), B: T.Tensor((16,), "float32")): T.tile.copy(A[0:16], B[0:16]) Tx.cast(A[0:16], B[0:16]) T.cta.cast(A[0:16], B[0:16]) @@ -171,7 +171,7 @@ def test_device_intrinsic_namespaces_are_canonical_and_classified(): assert T.metal is builder_op.metal assert T.nki is builder_op.nki - buffer = tvm.tirx.decl_buffer((1,), "float32") + buffer = tvm.tirx.decl_tensor((1,), "float32") calls = [ T.cuda.elect_sync(), T.cuda.thread_fence(), @@ -289,7 +289,7 @@ def register_backend(): def test_device_intrinsic_printer_roundtrips_canonical_namespaces(): @T.prim_func - def device_namespaces(dst: T.handle, A: T.Buffer((1,), "float32")): + def device_namespaces(dst: T.handle, A: T.Tensor((1,), "float32")): T.cuda.cta_sync() T.metal.simd_shuffle(A[0], 0) T.metal.simd_shuffle_up(A[0], 1) diff --git a/tests/python/tirx/test_roundtrip_namespaces.py b/tests/python/tirx/test_roundtrip_namespaces.py index d74bc74da96f..3a1b124c8da9 100644 --- a/tests/python/tirx/test_roundtrip_namespaces.py +++ b/tests/python/tirx/test_roundtrip_namespaces.py @@ -27,7 +27,7 @@ def from_source(code): def test_roundtrip_tir_namespaces_minimal(): # Exercise a selection of namespace ops and ensure round-trip consistency @T.prim_func - def func(A: T.Buffer((2, 2), "float16")) -> None: + def func(A: T.Tensor((2, 2), "float16")) -> None: T.ptx.wgmma.commit_group.sync.aligned() T.cuda.cluster_sync() T.ptx.cp.async_.wait_group(0) diff --git a/tests/python/tirx/test_verifier.py b/tests/python/tirx/test_verifier.py index b3429f9f3a5d..92208925e2b0 100644 --- a/tests/python/tirx/test_verifier.py +++ b/tests/python/tirx/test_verifier.py @@ -171,7 +171,7 @@ def test1(): T.cta_id([32]) T.warp_id([4]) T.lane_id([32]) - A = T.alloc_buffer((2,), layout=T.TileLayout(T.S[2, 1])) + A = T.alloc_tensor((2,), layout=T.TileLayout(T.S[2, 1])) A[0] = 0 # fmt: on @@ -185,7 +185,7 @@ def test2(): T.cta_id([32]) T.warp_id([4]) T.lane_id([32]) - A = T.alloc_buffer( + A = T.alloc_tensor( (512,), scope="shared", layout=T.ComposeLayout(3, 3, 3, T.TileLayout(T.S[(512,)])) ) @@ -197,7 +197,7 @@ def test2(): def test_host(): # fmt: off @T.prim_func(check_well_formed=False) - def test1(A: T.Buffer((16, 16), dtype='float32', align=16)): + def test1(A: T.Tensor((16, 16), dtype='float32', align=16)): A_map: T.let[T.handle("tensormap")] = T.tvm_stack_alloca("tensormap", 1) T.call_packed("runtime.cuTensorMapEncodeTiled", A_map, "float32", 2, A.data, 16, 16, 64, 16, 16, 1, 1, 0, 0, 0, 0) # noqa: E501 @@ -205,9 +205,9 @@ def test1(A: T.Buffer((16, 16), dtype='float32', align=16)): T.device_entry() for blockIdx in T.thread_binding(1, thread="blockIdx.x"): for threadIdx in T.thread_binding(128, thread="threadIdx.x"): - bar = T.alloc_buffer((1,), "uint64", scope="shared", align=8) - phase = T.alloc_buffer((1,), "int32", scope="local") - A_smem = T.alloc_buffer((16, 16), "float32", scope="shared", align=128) + bar = T.alloc_tensor((1,), "uint64", scope="shared", align=8) + phase = T.alloc_tensor((1,), "int32", scope="local") + A_smem = T.alloc_tensor((16, 16), "float32", scope="shared", align=128) phase[0] = 0 if threadIdx == 0: @@ -231,14 +231,14 @@ def test_device_func(): # is dropped. # fmt: off @T.prim_func(check_well_formed=False) - def test1(A: T.Buffer((128,), "float32")): + def test1(A: T.Tensor((128,), "float32")): T.device_entry() T.cta_id([1]) T.thread_id([128]) Tx.cta.fill(A, 0.) @T.prim_func(check_well_formed=False) - def test2(A: T.Buffer((128,), "float32")): + def test2(A: T.Tensor((128,), "float32")): T.device_entry() T.cta_id([128]) T.thread_id([128]) diff --git a/tests/python/tirx/transform/test_transform_flatten_buffer.py b/tests/python/tirx/transform/test_transform_flatten_buffer.py index 3e59bf931b47..9f828d66526e 100644 --- a/tests/python/tirx/transform/test_transform_flatten_buffer.py +++ b/tests/python/tirx/transform/test_transform_flatten_buffer.py @@ -44,7 +44,7 @@ def _collect_defined_buffers(func): defined = set() def visit(node): - if _is_buffer_binding(node, "tirx.alloc_buffer", "tirx.decl_buffer"): + if _is_buffer_binding(node, "tirx.alloc_tensor", "tirx.decl_tensor"): defined.add(node.var) tvm_ffi.structural_walk(func.body, visit) @@ -53,7 +53,7 @@ def visit(node): def _assert_loads_reference_defined_buffers(func): """Every BufferLoad — direct, or embedded in a buffer's type fields — - must reference a buffer defined by an AllocBuffer/DeclBuffer in the + must reference a buffer defined by an AllocTensor/DeclTensor in the function.""" defined = _collect_defined_buffers(func) @@ -76,7 +76,7 @@ def visit(node): stale.append(f"access of {buffer.name}") for index in node.indices: check_expr(index, f"index of {buffer.name}") - if _is_buffer_binding(node, "tirx.alloc_buffer", "tirx.decl_buffer"): + if _is_buffer_binding(node, "tirx.alloc_tensor", "tirx.decl_tensor"): for extent in node.var.shape: check_expr(extent, f"shape of {node.var.name}") if node.var.elem_offset is not None: @@ -100,8 +100,8 @@ def test_flatten_remaps_loads_in_view_shape(): def before(): n = T.alloc_local([1], "int32") n[0] = 8 - data = T.alloc_buffer([64], "float16", scope="shared") - view = T.decl_buffer((n[0],), "float16", data.data, scope="shared") + data = T.alloc_tensor([64], "float16", scope="shared") + view = T.decl_tensor((n[0],), "float16", data.data, scope="shared") view[0] = T.float16(0) _assert_loads_reference_defined_buffers(_flatten(before)) @@ -116,8 +116,8 @@ def test_flatten_remaps_loads_in_folded_elem_offset(): def before(): n = T.alloc_local([1], "int32") n[0] = 4 - base = T.alloc_buffer([128], "uint64", scope="shared") - mbar = T.decl_buffer((1,), "uint64", base.data, elem_offset=n[0], scope="shared") + base = T.alloc_tensor([128], "uint64", scope="shared") + mbar = T.decl_tensor((1,), "uint64", base.data, elem_offset=n[0], scope="shared") mbar[0] = T.uint64(1) after = _flatten(before) @@ -145,13 +145,13 @@ def test_flatten_keeps_identity_of_already_flat_buffers(): @T.prim_func(private=True) def before(): - flat = T.alloc_buffer([32], "float32", scope="shared", layout=None) + flat = T.alloc_tensor([32], "float32", scope="shared", layout=None) flat[0] = T.float32(0) before_allocs = {} def collect_before(node): - if _is_buffer_binding(node, "tirx.alloc_buffer"): + if _is_buffer_binding(node, "tirx.alloc_tensor"): before_allocs[node.var.name] = node.var tvm_ffi.structural_walk(before.body, collect_before) @@ -160,7 +160,7 @@ def collect_before(node): preserved = [] def visit(node): - if _is_buffer_binding(node, "tirx.alloc_buffer") and node.var.name in before_allocs: + if _is_buffer_binding(node, "tirx.alloc_tensor") and node.var.name in before_allocs: preserved.append(node.var.same_as(before_allocs[node.var.name])) tvm_ffi.structural_walk(after.body, visit) diff --git a/tests/python/tirx/transform/test_transform_lower_tirx.py b/tests/python/tirx/transform/test_transform_lower_tirx.py index e9e1c50cdafa..2c18f6b8eb40 100644 --- a/tests/python/tirx/transform/test_transform_lower_tirx.py +++ b/tests/python/tirx/transform/test_transform_lower_tirx.py @@ -69,7 +69,7 @@ def collect(node): def test_lower_tirx_opaque_optional_pragma_annotations(): @T.prim_func(private=True) - def before(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): + def before(A: T.Tensor(8, "float32"), B: T.Tensor(8, "float32")): for i in T.serial(8, annotations={"pragma_unroll": None}): B[i] = A[i] + 1.0 for i in T.serial(8, annotations={"pragma_unroll_explicit": None}): @@ -80,7 +80,7 @@ def before(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): B[i] = A[i] + 4.0 @T.prim_func(private=True) - def after(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): + def after(A: T.Tensor(8, "float32"), B: T.Tensor(8, "float32")): for i in T.serial(8): B[i] = A[i] + 1.0 for i in T.serial(8): @@ -100,12 +100,12 @@ def after(A: T.Buffer(8, "float32"), B: T.Buffer(8, "float32")): def test_lower_view_get(): @T.prim_func(private=True) - def before1(in_buf: T.Buffer(64, "float32"), out: T.Buffer(64, "float32")) -> None: + def before1(in_buf: T.Tensor(64, "float32"), out: T.Tensor(64, "float32")) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) T.warp_id([1]) lane_id = T.lane_id([32]) - A = T.alloc_buffer([2], dtype="float16", scope="local", layout=T.TileLayout(T.S[2:1])) + A = T.alloc_tensor([2], dtype="float16", scope="local", layout=T.TileLayout(T.S[2:1])) B_layout = A.layout.tile(L_LANE, (32,), (2,)) B = A.view(64, layout=B_layout) A_local = B.local(2) @@ -117,9 +117,9 @@ def before1(in_buf: T.Buffer(64, "float32"), out: T.Buffer(64, "float32")) -> No out[lane_id * 2 + i] = T.float32(A_local_1[i]) @T.prim_func(private=True) - def after1(in_buf: T.Buffer((64,), layout=None), out: T.Buffer((64,), layout=None)): - out_1 = T.decl_buffer((64,), data=out.data, layout=None) - in_buf_1 = T.decl_buffer((64,), data=in_buf.data, layout=None) + def after1(in_buf: T.Tensor((64,), layout=None), out: T.Tensor((64,), layout=None)): + out_1 = T.decl_tensor((64,), data=out.data, layout=None) + in_buf_1 = T.decl_tensor((64,), data=in_buf.data, layout=None) blockIdx_x = T.launch_thread("blockIdx.x", 1) threadIdx_x = T.launch_thread("threadIdx.x", 32) blockIdx_y = T.launch_thread("blockIdx.y", 1) @@ -133,19 +133,19 @@ def after1(in_buf: T.Buffer((64,), layout=None), out: T.Buffer((64,), layout=Non v: T.let[T.int32] = warp_id_in_cta lane_id: T.let[T.int32] = threadIdx_x % 32 A = T.alloc_local((2,), "float16", layout=None) - B = T.decl_buffer((64,), "float16", data=A.data, scope="local", layout=None) - A_local = T.decl_buffer((2,), "float16", data=A.data, scope="local", layout=None) + B = T.decl_tensor((64,), "float16", data=A.data, scope="local", layout=None) + A_local = T.decl_tensor((2,), "float16", data=A.data, scope="local", layout=None) for i in T.vectorized(2): A_local[i] = T.Cast("float16", in_buf_1[threadIdx_x * 2 + i]) - B_1 = T.decl_buffer((64,), "float16", data=A.data, scope="local", layout=None) - A_local_1 = T.decl_buffer((2,), "float16", data=A.data, scope="local", layout=None) + B_1 = T.decl_tensor((64,), "float16", data=A.data, scope="local", layout=None) + A_local_1 = T.decl_tensor((2,), "float16", data=A.data, scope="local", layout=None) for i in T.vectorized(2): out_1[threadIdx_x * 2 + i] = T.Cast("float32", A_local_1[i]) compare(before1, after1, LowerTIRx) @T.prim_func(private=True) - def before2(in_buf: T.Buffer((16, 16), "float32"), out: T.Buffer((16, 16), "float32")) -> None: + def before2(in_buf: T.Tensor((16, 16), "float32"), out: T.Tensor((16, 16), "float32")) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) T.warp_id([1]) @@ -153,7 +153,7 @@ def before2(in_buf: T.Buffer((16, 16), "float32"), out: T.Buffer((16, 16), "floa atom = T.TileLayout(T.S[(1, 2) : (2, 1)]) tile = T.TileLayout(T.S[(2, 2) : (2, 1)]) warp_atom = atom.tile(L_LANE, (8, 4), (1, 2)) - A = T.alloc_buffer( + A = T.alloc_tensor( [4, 2], dtype="float32", scope="local", layout=atom.tile(tile, (2, 2), (1, 2)) ) B_layout = warp_atom.tile(tile, (2, 2), (8, 8)) @@ -170,9 +170,9 @@ def before2(in_buf: T.Buffer((16, 16), "float32"), out: T.Buffer((16, 16), "floa out[lane_id // 4 * 8 + i // 2 * 8 + lane_id % 4, lane_id % 4 * 2 + i % 2] = A_local_1[i] @T.prim_func(private=True) - def after2(in_buf: T.Buffer((16, 16), layout=None), out: T.Buffer((16, 16), layout=None)): - out_1 = T.decl_buffer((256,), data=out.data, layout=None) - in_buf_1 = T.decl_buffer((256,), data=in_buf.data, layout=None) + def after2(in_buf: T.Tensor((16, 16), layout=None), out: T.Tensor((16, 16), layout=None)): + out_1 = T.decl_tensor((256,), data=out.data, layout=None) + in_buf_1 = T.decl_tensor((256,), data=in_buf.data, layout=None) blockIdx_x = T.launch_thread("blockIdx.x", 1) threadIdx_x = T.launch_thread("threadIdx.x", 32) blockIdx_y = T.launch_thread("blockIdx.y", 1) @@ -186,15 +186,15 @@ def after2(in_buf: T.Buffer((16, 16), layout=None), out: T.Buffer((16, 16), layo v: T.let[T.int32] = warp_id_in_cta lane_id: T.let[T.int32] = threadIdx_x % 32 A = T.alloc_local((8,), layout=None) - B = T.decl_buffer((256,), data=A.data, scope="local", layout=None) - A_local = T.decl_buffer((8,), data=A.data, scope="local", layout=None) + B = T.decl_tensor((256,), data=A.data, scope="local", layout=None) + A_local = T.decl_tensor((8,), data=A.data, scope="local", layout=None) for i in T.unroll(4): for j in T.vectorized(2): A_local[i * 2 + j] = in_buf_1[ i // 2 * 128 + threadIdx_x // 4 * 16 + i % 2 * 8 + j + threadIdx_x % 4 ] - B_1 = T.decl_buffer((256,), data=A.data, scope="local", layout=None) - A_local_1 = T.decl_buffer((8,), data=A.data, scope="local", layout=None) + B_1 = T.decl_tensor((256,), data=A.data, scope="local", layout=None) + A_local_1 = T.decl_tensor((8,), data=A.data, scope="local", layout=None) for i in T.vectorized(2): out_1[threadIdx_x // 4 * 128 + threadIdx_x % 4 * 18 + i] = A_local_1[i] @@ -202,7 +202,7 @@ def after2(in_buf: T.Buffer((16, 16), layout=None), out: T.Buffer((16, 16), layo @T.prim_func(private=True) def before3_wgmma_layout( - in_buf: T.Buffer((128, 128), "float32"), out: T.Buffer((128, 128), "float32") + in_buf: T.Tensor((128, 128), "float32"), out: T.Tensor((128, 128), "float32") ) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) @@ -215,7 +215,7 @@ def before3_wgmma_layout( warp_layout = warp_atom.tile(tile, (2, 128 // 8), (8, 8)) L_warp = T.TileLayout(T.S[8 : 1 @ warpid]) layout = warp_layout.tile(L_warp, (8, 1), (16, 128)) - acc = T.alloc_buffer( + acc = T.alloc_tensor( [64], dtype="float32", scope="local", @@ -242,10 +242,10 @@ def before3_wgmma_layout( @T.prim_func(private=True) def after3_wgmma_layout( - in_buf: T.Buffer((128, 128), layout=None), out: T.Buffer((128, 128), layout=None) + in_buf: T.Tensor((128, 128), layout=None), out: T.Tensor((128, 128), layout=None) ): - out_1 = T.decl_buffer((16384,), data=out.data, layout=None) - in_buf_1 = T.decl_buffer((16384,), data=in_buf.data, layout=None) + out_1 = T.decl_tensor((16384,), data=out.data, layout=None) + in_buf_1 = T.decl_tensor((16384,), data=in_buf.data, layout=None) blockIdx_x = T.launch_thread("blockIdx.x", 1) threadIdx_x = T.launch_thread("threadIdx.x", 256) blockIdx_y = T.launch_thread("blockIdx.y", 1) @@ -260,8 +260,8 @@ def after3_wgmma_layout( warp_id_in_wg: T.let[T.int32] = warp_id_in_cta % 4 lane_id: T.let[T.int32] = threadIdx_x % 32 acc = T.alloc_local((64,), layout=None) - B = T.decl_buffer((16384,), data=acc.data, scope="local", layout=None) - acc_local = T.decl_buffer((64,), data=acc.data, scope="local", layout=None) + B = T.decl_tensor((16384,), data=acc.data, scope="local", layout=None) + acc_local = T.decl_tensor((64,), data=acc.data, scope="local", layout=None) for i in range(16): for j in T.unroll(2): for vec in T.vectorized(2): @@ -273,8 +273,8 @@ def after3_wgmma_layout( + threadIdx_x % 4 * 2 + vec ] - B_1 = T.decl_buffer((16384,), data=acc.data, scope="local", layout=None) - acc_local_1 = T.decl_buffer((64,), data=acc.data, scope="local", layout=None) + B_1 = T.decl_tensor((16384,), data=acc.data, scope="local", layout=None) + acc_local_1 = T.decl_tensor((64,), data=acc.data, scope="local", layout=None) for i in range(16): for j in T.unroll(2): for vec in T.vectorized(2): @@ -291,13 +291,13 @@ def after3_wgmma_layout( @T.prim_func(private=True) def before4_multi_view_get( - in_buf: T.Buffer(64, "float32"), out: T.Buffer(64, "float32") + in_buf: T.Tensor(64, "float32"), out: T.Tensor(64, "float32") ) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) T.warp_id([1]) lane_id = T.lane_id([32]) - A = T.alloc_buffer([2], dtype="float16", scope="local", layout=T.TileLayout(T.S[2:1])) + A = T.alloc_tensor([2], dtype="float16", scope="local", layout=T.TileLayout(T.S[2:1])) B_layout = A.layout.tile(L_LANE, (32,), (2,)) B = A.view(64, layout=B_layout) B_1 = A.view(64, layout=B_layout) @@ -315,10 +315,10 @@ def before4_multi_view_get( @T.prim_func(private=True) def after4_multi_view_get( - in_buf: T.Buffer((64,), layout=None), out: T.Buffer((64,), layout=None) + in_buf: T.Tensor((64,), layout=None), out: T.Tensor((64,), layout=None) ): - out_1 = T.decl_buffer((64,), data=out.data, layout=None) - in_buf_1 = T.decl_buffer((64,), data=in_buf.data, layout=None) + out_1 = T.decl_tensor((64,), data=out.data, layout=None) + in_buf_1 = T.decl_tensor((64,), data=in_buf.data, layout=None) blockIdx_x = T.launch_thread("blockIdx.x", 1) threadIdx_x = T.launch_thread("threadIdx.x", 32) blockIdx_y = T.launch_thread("blockIdx.y", 1) @@ -332,17 +332,17 @@ def after4_multi_view_get( v: T.let[T.int32] = warp_id_in_cta lane_id: T.let[T.int32] = threadIdx_x % 32 A = T.alloc_local((2,), "float16", layout=None) - B = T.decl_buffer((64,), "float16", data=A.data, scope="local", layout=None) - B_1 = T.decl_buffer((64,), "float16", data=A.data, scope="local", layout=None) - A_local = T.decl_buffer((2,), "float16", data=A.data, scope="local", layout=None) + B = T.decl_tensor((64,), "float16", data=A.data, scope="local", layout=None) + B_1 = T.decl_tensor((64,), "float16", data=A.data, scope="local", layout=None) + A_local = T.decl_tensor((2,), "float16", data=A.data, scope="local", layout=None) A_local[0] = T.Cast("float16", in_buf_1[threadIdx_x * 2]) - A_local_1 = T.decl_buffer((2,), "float16", data=A.data, scope="local", layout=None) + A_local_1 = T.decl_tensor((2,), "float16", data=A.data, scope="local", layout=None) A_local_1[1] = T.Cast("float16", in_buf_1[threadIdx_x * 2 + 1]) - B_2 = T.decl_buffer((64,), "float16", data=A.data, scope="local", layout=None) - B_3 = T.decl_buffer((64,), "float16", data=A.data, scope="local", layout=None) - A_local_2 = T.decl_buffer((2,), "float16", data=A.data, scope="local", layout=None) + B_2 = T.decl_tensor((64,), "float16", data=A.data, scope="local", layout=None) + B_3 = T.decl_tensor((64,), "float16", data=A.data, scope="local", layout=None) + A_local_2 = T.decl_tensor((2,), "float16", data=A.data, scope="local", layout=None) out_1[threadIdx_x * 2] = T.Cast("float32", A_local_2[0]) - A_local_3 = T.decl_buffer((2,), "float16", data=A.data, scope="local", layout=None) + A_local_3 = T.decl_tensor((2,), "float16", data=A.data, scope="local", layout=None) out_1[threadIdx_x * 2 + 1] = T.Cast("float32", A_local_3[1]) compare(before4_multi_view_get, after4_multi_view_get, LowerTIRx) @@ -635,13 +635,13 @@ def after(): def test_lower_layout(): @T.prim_func(private=True) - def before(A: T.Buffer((128, 32), "float16")) -> None: + def before(A: T.Tensor((128, 32), "float16")) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) T.warp_id([4]) T.lane_id([32]) tid = T.thread_id([128]) - A_smem = T.alloc_buffer( + A_smem = T.alloc_tensor( [128, 32], dtype="float16", scope="shared", @@ -659,8 +659,8 @@ def before(A: T.Buffer((128, 32), "float16")) -> None: compose_q = T.dynamic("compose_q", "int32") @T.prim_func(private=True) - def after(A: T.Buffer((128, 32), "float16", layout=None)) -> None: - A_1 = T.decl_buffer((4096,), "float16", data=A.data, layout=None) + def after(A: T.Tensor((128, 32), "float16", layout=None)) -> None: + A_1 = T.decl_tensor((4096,), "float16", data=A.data, layout=None) blockIdx_x = T.launch_thread("blockIdx.x", 1) threadIdx_x = T.launch_thread("threadIdx.x", 128) blockIdx_y = T.launch_thread("blockIdx.y", 1) @@ -706,12 +706,12 @@ def after(A: T.Buffer((128, 32), "float16", layout=None)) -> None: def test_lower_opcall_fail(): @T.prim_func - def test(A: T.Buffer((64,), "float32", scope="global")) -> None: + def test(A: T.Tensor((64,), "float32", scope="global")) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) T.warp_id([1]) T.lane_id([32]) - A_smem = T.alloc_buffer([64], dtype="float32", scope="shared") + A_smem = T.alloc_tensor([64], dtype="float32", scope="shared") Tx.cta.copy(A[0:64], A_smem[0:64]) for i in range(10): Tx.cta.fill(A_smem[0:64], T.float32(0)) @@ -728,8 +728,8 @@ def before(): T.device_entry() T.cta_id([1]) T.thread_id([128]) - buf = T.alloc_buffer([1024], "uint8", scope="shared.dyn") - A = T.decl_buffer([128], "float16", buf.data, elem_offset=32) + buf = T.alloc_tensor([1024], "uint8", scope="shared.dyn") + A = T.decl_tensor([128], "float16", buf.data, elem_offset=32) T.evaluate(A.access_ptr("rw", ptr_type="float16", offset=A.elem_offset_of([64]))) @T.prim_func(private=True) @@ -741,8 +741,8 @@ def after(): ) v: T.let[T.int32] = blockIdx_x v_1: T.let[T.int32] = threadIdx_x - buf = T.alloc_buffer((1024,), "uint8", scope="shared.dyn", layout=None) - A = T.decl_buffer( + buf = T.alloc_tensor((1024,), "uint8", scope="shared.dyn", layout=None) + A = T.decl_tensor( (128,), "float16", data=buf.data, elem_offset=32, scope="shared.dyn", layout=None ) T.tvm_access_ptr(T.type_annotation("float16"), buf.data, T.Add(32, 64), T.Sub(128, 64), 3) @@ -819,7 +819,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() T.cta_id([1]) @@ -858,7 +858,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() T.cta_id([1]) @@ -903,7 +903,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() T.cta_id([1]) @@ -941,7 +941,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() T.cta_id([1]) @@ -979,7 +979,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() T.cta_id([1]) @@ -1019,7 +1019,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() T.cta_id([1]) @@ -1041,7 +1041,7 @@ def before( def test_lower_exec_context_keeps_plain_predicate_condition(): @T.prim_func(private=True) - def before(A: T.Buffer((1,), "float32", scope="global")): + def before(A: T.Tensor((1,), "float32", scope="global")): T.device_entry() T.cta_id([1]) wg_id = T.warpgroup_id([2]) @@ -1061,7 +1061,7 @@ def before(A: T.Buffer((1,), "float32", scope="global")): def test_lower_exec_context_keeps_plain_scope_predicate_condition(): @T.prim_func(private=True) - def before(A: T.Buffer((1,), "float32", scope="global")): + def before(A: T.Tensor((1,), "float32", scope="global")): T.device_entry() T.cta_id([1]) wg_id = T.warpgroup_id([2]) @@ -1081,7 +1081,7 @@ def before(A: T.Buffer((1,), "float32", scope="global")): def test_simplify_uses_floor_div_scope_predicate_as_context_fact(): @T.prim_func(private=True) - def before(A: T.Buffer((16,), "float32", scope="global")): + def before(A: T.Tensor((16,), "float32", scope="global")): T.device_entry() T.cta_id([1]) wg_id = T.warpgroup_id([2]) @@ -1120,7 +1120,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() T.cta_id([1]) @@ -1143,7 +1143,7 @@ def before( def test_lower_cleanup_accepts_bool_elect_sync_else_path(): @T.prim_func(private=True) - def before(A: T.Buffer((32,), "int32", scope="global")): + def before(A: T.Tensor((32,), "int32", scope="global")): T.device_entry() T.cta_id([1]) T.warp_id([1]) @@ -1180,7 +1180,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() T.cta_id([1]) @@ -1221,7 +1221,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() cbx, cby = T.cta_id_in_cluster([2, 3]) @@ -1267,7 +1267,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() bx = T.cta_id([8]) @@ -1305,7 +1305,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() cbx, cby = T.cta_id_in_cluster([4, 2]) @@ -1340,7 +1340,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() cbx, cby = T.cta_id_in_cluster([4, 2]) @@ -1387,7 +1387,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() T.cta_id_in_cluster([2]) @@ -1425,7 +1425,7 @@ def impl(): @T.prim_func(private=True) def before( - A: T.Buffer((1,), "float32", scope="global"), B: T.Buffer((1,), "float32", scope="global") + A: T.Tensor((1,), "float32", scope="global"), B: T.Tensor((1,), "float32", scope="global") ): T.device_entry() cbx, cby = T.cta_id_in_cluster([3, 2]) @@ -1449,8 +1449,8 @@ def before(): T.device_entry() T.cta_id([1]) T.thread_id([128]) - A = T.alloc_buffer([64, 64], "float16", scope="local") - A0 = T.decl_buffer([64], "float16", A.data, elem_offset=A.elem_offset_of([32, 32])) + A = T.alloc_tensor([64, 64], "float16", scope="local") + A0 = T.decl_tensor([64], "float16", A.data, elem_offset=A.elem_offset_of([32, 32])) T.evaluate(T.address_of(A0[32])) @T.prim_func(private=True) @@ -1463,7 +1463,7 @@ def after(): v: T.let[T.int32] = blockIdx_x v_1: T.let[T.int32] = threadIdx_x A = T.alloc_local((4096,), "float16", layout=None) - A0 = T.decl_buffer( + A0 = T.decl_tensor( (64,), "float16", data=A.data, elem_offset=2080, scope="local", layout=None ) T.address_of(A0[32]) @@ -1477,7 +1477,7 @@ class State: def __init__(self, smem): self.A = T.alloc_local([1], "float16") self.B = T.alloc_local([1], "float16") - self.C = T.decl_buffer([1], "float16", smem, elem_offset=0, scope="shared.dyn") + self.C = T.decl_tensor([1], "float16", smem, elem_offset=0, scope="shared.dyn") @register_mutable_decl("TestMutableCells.int_var1") def int_var1(val): @@ -1498,7 +1498,7 @@ def int_var2(val): @T.prim_func(private=True) def before(): T.device_entry() - smem = T.alloc_buffer([100], "uint8", scope="shared.dyn") + smem = T.alloc_tensor([100], "uint8", scope="shared.dyn") state = State(smem.data) state.A[0] = T.float16(1) state.B[0] = T.float16(2) @@ -1514,10 +1514,10 @@ def before(): @T.prim_func(private=True) def after(): - smem = T.alloc_buffer([100], "uint8", scope="shared.dyn", layout=None) + smem = T.alloc_tensor([100], "uint8", scope="shared.dyn", layout=None) A = T.alloc_local((1,), "float16", layout=None) B = T.alloc_local((1,), "float16", layout=None) - C = T.decl_buffer( + C = T.decl_tensor( (1,), "float16", data=smem.data, elem_offset=0, scope="shared.dyn", layout=None ) A[0] = T.float16(1) @@ -1540,23 +1540,23 @@ def after(): def test_alloc_buffer_with_thread_axis_layout(): - """alloc_buffer with thread-axis layout should lower to 1D physical buffer with memory-axis span.""" # noqa: E501 + """alloc_tensor with thread-axis layout should lower to 1D physical buffer with memory-axis span.""" # noqa: E501 @T.prim_func(private=True) - def before(out: T.Buffer((128, 4), "float32")) -> None: + def before(out: T.Tensor((128, 4), "float32")) -> None: T.device_entry() bx, by, bz = T.cta_id([1, 1, 1]) T.warpgroup_id([1]) warp_id = T.warp_id_in_wg([4]) lane_id = T.lane_id([32]) - reg_wg = T.alloc_buffer((128, 4), "float32", scope="local", layout=wg_local_layout(4)) + reg_wg = T.alloc_tensor((128, 4), "float32", scope="local", layout=wg_local_layout(4)) reg = reg_wg.local(4) for i in T.serial(4): reg[i] = out[lane_id + warp_id * 32, i] @T.prim_func(private=True) - def after(out: T.Buffer((128, 4), layout=None)): - out_1 = T.decl_buffer((512,), data=out.data, layout=None) + def after(out: T.Tensor((128, 4), layout=None)): + out_1 = T.decl_tensor((512,), data=out.data, layout=None) blockIdx_x = T.launch_thread("blockIdx.x", 1) threadIdx_x = T.launch_thread("threadIdx.x", 128) blockIdx_y = T.launch_thread("blockIdx.y", 1) @@ -1571,7 +1571,7 @@ def after(out: T.Buffer((128, 4), layout=None)): warp_id: T.let[T.int32] = warp_id_in_cta % 4 lane_id: T.let[T.int32] = threadIdx_x % 32 reg_wg = T.alloc_local((4,), layout=None) - reg = T.decl_buffer((4,), data=reg_wg.data, scope="local", layout=None) + reg = T.decl_tensor((4,), data=reg_wg.data, scope="local", layout=None) for i in range(4): reg[i] = out_1[warp_id_in_cta % 4 * 128 + threadIdx_x % 32 * 4 + i] @@ -1588,7 +1588,7 @@ def test_scope_id_compliment_no_div_by_zero(): with pytest.raises(Exception): @T.prim_func - def func(A: T.Buffer((1,))): + def func(A: T.Tensor((1,))): T.device_entry() cb_m, cb_n = T.cta_id_in_cluster([2, 2]) bx = T.cta_id([1]) diff --git a/tests/python/tirx/transform/test_transform_naive_allocator.py b/tests/python/tirx/transform/test_transform_naive_allocator.py index d18b83681895..9df664e44a54 100644 --- a/tests/python/tirx/transform/test_transform_naive_allocator.py +++ b/tests/python/tirx/transform/test_transform_naive_allocator.py @@ -32,18 +32,18 @@ def test_one_alloc(): # fmt: off @T.prim_func - def copy(A: T.Buffer(src_shape, 'float32', layout=src_layout)) -> None: + def copy(A: T.Tensor(src_shape, 'float32', layout=src_layout)) -> None: T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) Tx.copy(A_sbuf, A) @T.prim_func - def expected(A: T.Buffer(src_shape, 'float32', layout=src_layout)) -> None: + def expected(A: T.Tensor(src_shape, 'float32', layout=src_layout)) -> None: T.func_attr({"global_symbol": "copy"}) T.device_entry() - A_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout, allocated_addr=[0]) # noqa: E501 + A_sbuf = T.alloc_tensor(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout, allocated_addr=[0]) # noqa: E501 Tx.copy(A_sbuf, A) # fmt: on @@ -57,16 +57,16 @@ def test_two_alloc(): @T.prim_func def copy(A_ptr: T.handle) -> None: T.device_entry() - A_sbuf = T.alloc_buffer([256, 512], "float32", scope="trn.sbuf", layout="PF") - B_sbuf = T.alloc_buffer([512, 512], "float32", scope="trn.sbuf", layout="PF") + A_sbuf = T.alloc_tensor([256, 512], "float32", scope="trn.sbuf", layout="PF") + B_sbuf = T.alloc_tensor([512, 512], "float32", scope="trn.sbuf", layout="PF") Tx.copy(B_sbuf[0:256, :], A_sbuf) @T.prim_func def expected(A_ptr: T.handle) -> None: T.func_attr({"global_symbol": "copy"}) T.device_entry() - A_sbuf = T.alloc_buffer([256, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[0]) # noqa: E501 - B_sbuf = T.alloc_buffer([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[2*512*4]) # noqa: E501 + A_sbuf = T.alloc_tensor([256, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[0]) # noqa: E501 + B_sbuf = T.alloc_tensor([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[2*512*4]) # noqa: E501 Tx.copy(B_sbuf[0:256, :], A_sbuf) # fmt: on @@ -80,16 +80,16 @@ def test_existing_alloc(): @T.prim_func def copy(A_ptr: T.handle) -> None: T.device_entry() - A_sbuf = T.alloc_buffer([256, 512], "float32", scope="trn.sbuf", layout="PF") - B_sbuf = T.alloc_buffer([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[1]) # noqa: E501 + A_sbuf = T.alloc_tensor([256, 512], "float32", scope="trn.sbuf", layout="PF") + B_sbuf = T.alloc_tensor([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[1]) # noqa: E501 Tx.copy(B_sbuf[0:256, :], A_sbuf) @T.prim_func def expected(A_ptr: T.handle) -> None: T.func_attr({"global_symbol": "copy"}) T.device_entry() - A_sbuf = T.alloc_buffer([256, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[4*512*4+1]) # noqa: E501 - B_sbuf = T.alloc_buffer([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[1]) # noqa: E501 + A_sbuf = T.alloc_tensor([256, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[4*512*4+1]) # noqa: E501 + B_sbuf = T.alloc_tensor([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[1]) # noqa: E501 Tx.copy(B_sbuf[0:256, :], A_sbuf) # fmt: on @@ -103,18 +103,18 @@ def test_workspace(): @T.prim_func def copy(A_ptr: T.handle) -> None: T.device_entry() - A_sbuf = T.alloc_buffer([256, 512], "float32", scope="trn.sbuf", layout="PF") - B_sbuf = T.alloc_buffer([512, 512], "float32", scope="trn.sbuf", layout="PF") - C_sbuf = T.alloc_buffer([128, 1024], "float32", scope="trn.sbuf") + A_sbuf = T.alloc_tensor([256, 512], "float32", scope="trn.sbuf", layout="PF") + B_sbuf = T.alloc_tensor([512, 512], "float32", scope="trn.sbuf", layout="PF") + C_sbuf = T.alloc_tensor([128, 1024], "float32", scope="trn.sbuf") Tx.copy(B_sbuf[0:256, :], A_sbuf, workspace={"C": C_sbuf}) @T.prim_func def expected(A_ptr: T.handle) -> None: T.func_attr({"global_symbol": "copy"}) T.device_entry() - A_sbuf = T.alloc_buffer([256, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[0]) # noqa: E501 - B_sbuf = T.alloc_buffer([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[2*512*4]) # noqa: E501 - C_sbuf = T.alloc_buffer([128, 1024], "float32", scope="trn.sbuf", allocated_addr=[2*512*4+4*512*4]) # noqa: E501 + A_sbuf = T.alloc_tensor([256, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[0]) # noqa: E501 + B_sbuf = T.alloc_tensor([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[2*512*4]) # noqa: E501 + C_sbuf = T.alloc_tensor([128, 1024], "float32", scope="trn.sbuf", allocated_addr=[2*512*4+4*512*4]) # noqa: E501 Tx.copy(B_sbuf[0:256, :], A_sbuf, workspace={"C": C_sbuf}) # fmt: on @@ -128,18 +128,18 @@ def test_other_scope_alloc(): @T.prim_func def copy(A_ptr: T.handle) -> None: T.device_entry() - A_sbuf = T.alloc_buffer([256, 512], "float32", scope="trn.sbuf", layout="PF") - B_sbuf = T.alloc_buffer([512, 512], "float32", scope="trn.sbuf", layout="PF") - C_sbuf = T.alloc_buffer([8, 128, 512], "float32", scope="global") + A_sbuf = T.alloc_tensor([256, 512], "float32", scope="trn.sbuf", layout="PF") + B_sbuf = T.alloc_tensor([512, 512], "float32", scope="trn.sbuf", layout="PF") + C_sbuf = T.alloc_tensor([8, 128, 512], "float32", scope="global") Tx.copy(B_sbuf[0:256, :], A_sbuf, workspace={"C": C_sbuf}) @T.prim_func def expected(A_ptr: T.handle) -> None: T.func_attr({"global_symbol": "copy"}) T.device_entry() - A_sbuf = T.alloc_buffer([256, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[0]) # noqa: E501 - B_sbuf = T.alloc_buffer([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[2*512*4]) # noqa: E501 - C_sbuf = T.alloc_buffer([8, 128, 512], "float32", scope="global") + A_sbuf = T.alloc_tensor([256, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[0]) # noqa: E501 + B_sbuf = T.alloc_tensor([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[2*512*4]) # noqa: E501 + C_sbuf = T.alloc_tensor([8, 128, 512], "float32", scope="global") Tx.copy(B_sbuf[0:256, :], A_sbuf, workspace={"C": C_sbuf}) # fmt: on @@ -153,8 +153,8 @@ def test_buffer_views(): @T.prim_func def copy(A_ptr: T.handle) -> None: T.device_entry() - A_sbuf = T.alloc_buffer([256, 512], "float32", scope="trn.sbuf", layout="PF") - B_sbuf = T.alloc_buffer([512, 512], "float32", scope="trn.sbuf", layout="PF") + A_sbuf = T.alloc_tensor([256, 512], "float32", scope="trn.sbuf", layout="PF") + B_sbuf = T.alloc_tensor([512, 512], "float32", scope="trn.sbuf", layout="PF") B_view = B_sbuf.view(2, 256, 512) Tx.copy(B_view[0], A_sbuf) @@ -162,8 +162,8 @@ def copy(A_ptr: T.handle) -> None: def expected(A_ptr: T.handle) -> None: T.func_attr({"global_symbol": "copy"}) T.device_entry() - A_sbuf = T.alloc_buffer([256, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[0]) # noqa: E501 - B_sbuf = T.alloc_buffer([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[2*512*4]) # noqa: E501 + A_sbuf = T.alloc_tensor([256, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[0]) # noqa: E501 + B_sbuf = T.alloc_tensor([512, 512], "float32", scope="trn.sbuf", layout="PF", allocated_addr=[2*512*4]) # noqa: E501 B_view = B_sbuf.view(2, 256, 512) Tx.copy(B_view[0], A_sbuf) # fmt: on