Skip to content

Commit c309e4e

Browse files
authored
[TIR] Add cooperative_tensor builtins and metal.cooperative_tensor storage scope (#19423)
part of tile-ai/tilelang#1869 ## Summary Add TIR builtins and storage scope for Metal cooperative_tensor operations (MetalPerformancePrimitives / Metal 4). ## Motivation Apple Metal 4 introduces MetalPerformancePrimitives (MPP) with `matmul2d` using `cooperative_tensor` operands. On M5, this routes to NAX tensor cores; on M1-M4, it falls back to simdgroup matrix instructions. These TIR primitives enable backend codegen to emit MPP calls. ## Changes ### New TIR builtins - `cooperative_tensor_fill(d, index, value, rows, cols)` - `cooperative_tensor_load(d, index, ptr, stride, rows, cols, transpose)` - `cooperative_tensor_store(d, index, ptr, stride, rows, cols, transpose)` - `cooperative_tensor_multiply_accumulate(d, di, a, ai, b, bi, c, ci, M, N, K, trans_a, trans_b)` ### New storage scope - `metal.cooperative_tensor` (`StorageRank::kMetalCooperativeTensor`) ### Files changed - `include/tvm/tirx/builtin.h` — Op declarations - `src/tirx/op/builtin.cc` — Op registrations - `python/tvm/tirx/op.py` — Python wrappers - `python/tvm/script/ir_builder/tirx/ir.py` — Script parser exports - `src/runtime/thread_storage_scope.h` — StorageRank enum + scope parsing These builtins mirror the existing `simdgroup_*` builtins for the older Metal simdgroup matrix API, extended with M/N/K dimension parameters for the matmul2d descriptor.
1 parent 2c76c79 commit c309e4e

5 files changed

Lines changed: 176 additions & 0 deletions

File tree

‎include/tvm/tirx/builtin.h‎

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -782,6 +782,51 @@ TVM_DLL const Op& simdgroup_store();
782782
*/
783783
TVM_DLL const Op& simdgroup_multiply_accumulate();
784784

785+
// Metal cooperative_tensor intrinsics (MetalPerformancePrimitives / Metal 4)
786+
787+
/*!
788+
* \brief Fill a cooperative_tensor with a given value.
789+
*
790+
* void cooperative_tensor_fill(Var d, PrimExpr index, PrimExpr value,
791+
* int rows, int cols);
792+
*/
793+
TVM_DLL const Op& cooperative_tensor_fill();
794+
795+
/*!
796+
* \brief Load data from device or threadgroup memory into a cooperative_tensor.
797+
*
798+
* void cooperative_tensor_load(Var d, PrimExpr index, PrimExpr ptr,
799+
* PrimExpr stride, int rows, int cols,
800+
* bool transpose_matrix,
801+
* int mma_M, int mma_N, int mma_K,
802+
* int operand_role);
803+
* operand_role: 0=left(A), 1=right(B), 2=destination(C)
804+
*/
805+
TVM_DLL const Op& cooperative_tensor_load();
806+
807+
/*!
808+
* \brief Store data from a cooperative_tensor to device or threadgroup memory.
809+
*
810+
* void cooperative_tensor_store(Var d, PrimExpr index, PrimExpr ptr,
811+
* PrimExpr stride, int rows, int cols,
812+
* bool transpose_matrix,
813+
* int mma_M, int mma_N, int mma_K,
814+
* int operand_role);
815+
* operand_role: 0=left(A), 1=right(B), 2=destination(C)
816+
*/
817+
TVM_DLL const Op& cooperative_tensor_store();
818+
819+
/*!
820+
* \brief Multiply and accumulate two matrices using cooperative_tensor
821+
* (MetalPerformancePrimitives matmul2d).
822+
*
823+
* void cooperative_tensor_multiply_accumulate(
824+
* Var d, PrimExpr index_d, Var a, PrimExpr index_a,
825+
* Var b, PrimExpr index_b, Var c, PrimExpr index_c,
826+
* int M, int N, int K, bool transpose_a, bool transpose_b);
827+
*/
828+
TVM_DLL const Op& cooperative_tensor_multiply_accumulate();
829+
785830
// TODO(tvm-team) replace the usage of the vector operations by Shuffle.
786831
/*!
787832
* \brief Get the high level half of the vector

‎python/tvm/tirx/op.py‎

Lines changed: 104 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1793,6 +1793,110 @@ def simdgroup_multiply_accumulate(
17931793
)
17941794

17951795

1796+
def cooperative_tensor_fill(
1797+
d: Var,
1798+
index: PrimExpr,
1799+
value: PrimExpr,
1800+
rows: int,
1801+
cols: int,
1802+
):
1803+
return call_intrin("handle", "tirx.cooperative_tensor_fill", d, index, value, rows, cols)
1804+
1805+
1806+
def cooperative_tensor_load(
1807+
d: Var,
1808+
index: PrimExpr,
1809+
ptr: PrimExpr,
1810+
stride: PrimExpr,
1811+
rows: int,
1812+
cols: int,
1813+
transpose_matrix: bool = False,
1814+
mma_M: int = 0,
1815+
mma_N: int = 0,
1816+
mma_K: int = 0,
1817+
operand_role: int = 0,
1818+
):
1819+
return call_intrin(
1820+
"handle",
1821+
"tirx.cooperative_tensor_load",
1822+
d,
1823+
index,
1824+
ptr,
1825+
stride,
1826+
rows,
1827+
cols,
1828+
transpose_matrix,
1829+
mma_M,
1830+
mma_N,
1831+
mma_K,
1832+
operand_role,
1833+
)
1834+
1835+
1836+
def cooperative_tensor_store(
1837+
d: PrimExpr,
1838+
index: PrimExpr,
1839+
ptr: PrimExpr,
1840+
stride: PrimExpr,
1841+
rows: int,
1842+
cols: int,
1843+
transpose_matrix: bool = False,
1844+
mma_M: int = 0,
1845+
mma_N: int = 0,
1846+
mma_K: int = 0,
1847+
operand_role: int = 0,
1848+
):
1849+
return call_intrin(
1850+
"handle",
1851+
"tirx.cooperative_tensor_store",
1852+
d,
1853+
index,
1854+
ptr,
1855+
stride,
1856+
rows,
1857+
cols,
1858+
transpose_matrix,
1859+
mma_M,
1860+
mma_N,
1861+
mma_K,
1862+
operand_role,
1863+
)
1864+
1865+
1866+
def cooperative_tensor_multiply_accumulate(
1867+
d: Var,
1868+
index_d: PrimExpr,
1869+
a: Var,
1870+
index_a: PrimExpr,
1871+
b: Var,
1872+
index_b: PrimExpr,
1873+
c: Var,
1874+
index_c: PrimExpr,
1875+
M: int,
1876+
N: int,
1877+
K: int,
1878+
transpose_a: bool = False,
1879+
transpose_b: bool = False,
1880+
):
1881+
return call_intrin(
1882+
"handle",
1883+
"tirx.cooperative_tensor_multiply_accumulate",
1884+
d,
1885+
index_d,
1886+
a,
1887+
index_a,
1888+
b,
1889+
index_b,
1890+
c,
1891+
index_c,
1892+
M,
1893+
N,
1894+
K,
1895+
transpose_a,
1896+
transpose_b,
1897+
)
1898+
1899+
17961900
def vectorlow(dtype, vec):
17971901
"""Get the low level half of the vector
17981902

‎python/tvm/tirx/script/builder/ir.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1965,6 +1965,10 @@ def wrapped(*args, **kwargs) -> T:
19651965
simdgroup_load = _op_wrapper(_tir_op.simdgroup_load)
19661966
simdgroup_store = _op_wrapper(_tir_op.simdgroup_store)
19671967
simdgroup_multiply_accumulate = _op_wrapper(_tir_op.simdgroup_multiply_accumulate)
1968+
cooperative_tensor_fill = _op_wrapper(_tir_op.cooperative_tensor_fill)
1969+
cooperative_tensor_load = _op_wrapper(_tir_op.cooperative_tensor_load)
1970+
cooperative_tensor_store = _op_wrapper(_tir_op.cooperative_tensor_store)
1971+
cooperative_tensor_multiply_accumulate = _op_wrapper(_tir_op.cooperative_tensor_multiply_accumulate)
19681972
create_barriers = _op_wrapper(_tir_op.create_barriers)
19691973
assume = _op_wrapper(_tir_op.assume)
19701974
undef = _op_wrapper(_tir_op.undef)
@@ -2255,6 +2259,10 @@ def wrapped(*args, **kwargs):
22552259
"simdgroup_load",
22562260
"simdgroup_store",
22572261
"simdgroup_multiply_accumulate",
2262+
"cooperative_tensor_fill",
2263+
"cooperative_tensor_load",
2264+
"cooperative_tensor_store",
2265+
"cooperative_tensor_multiply_accumulate",
22582266
"create_barriers",
22592267
"mma_store",
22602268
"mma_fill",

‎src/runtime/thread_storage_scope.h‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -71,6 +71,8 @@ enum class StorageRank {
7171
kMMAMatrixC = 11,
7272
/*! \brief Metal SIMD group memory */
7373
kMetalSimdGroup = 12,
74+
/*! \brief Metal cooperative_tensor memory (MetalPerformancePrimitives) */
75+
kMetalCooperativeTensor = 13,
7476
};
7577

7678
/*!
@@ -129,6 +131,8 @@ struct StorageScope {
129131
return "m16n8k8.matrixC" + tag;
130132
case StorageRank::kMetalSimdGroup:
131133
return "metal.simdgroup" + tag;
134+
case StorageRank::kMetalCooperativeTensor:
135+
return "metal.cooperative_tensor" + tag;
132136
default:
133137
TVM_FFI_THROW(InternalError) << "unknown storage scope";
134138
return "";
@@ -182,6 +186,9 @@ struct StorageScope {
182186
} else if (s.compare(0, 15, "metal.simdgroup") == 0) {
183187
r.rank = StorageRank::kMetalSimdGroup;
184188
r.tag = s.substr(15, std::string::npos);
189+
} else if (s.compare(0, 24, "metal.cooperative_tensor") == 0) {
190+
r.rank = StorageRank::kMetalCooperativeTensor;
191+
r.tag = s.substr(24, std::string::npos);
185192
} else {
186193
TVM_FFI_THROW(InternalError) << "unknown storage scope " << s;
187194
}

‎src/tirx/op/builtin.cc‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -345,6 +345,18 @@ TIR_DEFINE_BUILTIN_FUNC(simdgroup_store)
345345
TIR_DEFINE_BUILTIN_FUNC(simdgroup_multiply_accumulate)
346346
.set_attr<TCallEffectKind>("TCallEffectKind", Integer(CallEffectKind::kOpaque));
347347

348+
TIR_DEFINE_BUILTIN_FUNC(cooperative_tensor_fill)
349+
.set_attr<TCallEffectKind>("TCallEffectKind", Integer(CallEffectKind::kOpaque));
350+
351+
TIR_DEFINE_BUILTIN_FUNC(cooperative_tensor_load)
352+
.set_attr<TCallEffectKind>("TCallEffectKind", Integer(CallEffectKind::kOpaque));
353+
354+
TIR_DEFINE_BUILTIN_FUNC(cooperative_tensor_store)
355+
.set_attr<TCallEffectKind>("TCallEffectKind", Integer(CallEffectKind::kOpaque));
356+
357+
TIR_DEFINE_BUILTIN_FUNC(cooperative_tensor_multiply_accumulate)
358+
.set_attr<TCallEffectKind>("TCallEffectKind", Integer(CallEffectKind::kOpaque));
359+
348360
TIR_DEFINE_BUILTIN_FUNC(vectorhigh)
349361
.set_attr<TCallEffectKind>("TCallEffectKind", Integer(CallEffectKind::kPure))
350362
.set_attr<TScriptDtypePrintLocation>("TScriptDtypePrintLocation",

0 commit comments

Comments
 (0)