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