[Metal] M5 Cooperative Tensor T.gemm (#2252)

* [Metal] Add cooperative tensor intrinsics

Expose TileLang-owned cooperative tensor builtins so Metal MPP lowering does not depend on extra TVM fork APIs.

* [Metal] Select cooperative tensor GEMM lowering

Add a shape-aware MPP instruction choice for shared-output Metal GEMM while preserving simdgroup fallback for fragments and unsupported tiles.

* [Metal] Emit MPP cooperative tensor shaders

Generate Metal 4 MPP matmul2d code for cooperative tensor intrinsics and keep source-only codegen separate from runtime compilation.

* [Metal] Lower GEMM through MPP tensor ops

Split Metal GEMM lowering into simdgroup and cooperative tensor emitters so M5 tiles use MPP while fragment accumulators keep the existing path.

* [Metal] Guard generic passes for cooperative tensors

Keep generic allocation and storage rewrites away from opaque Metal cooperative tensor scopes to avoid invalid scope analysis.

* [Metal] Test cooperative tensor GEMM coverage

Add runtime and source-only coverage for non-square MPP GEMM so the new cooperative tensor path is reproducible in CI and on M5.

* [Metal] Update TVM Metal 4 runtime guard

Point the submodule at the macOS SDK guarded Metal 4 runtime update used by cooperative tensor shaders.

* [Metal] Document cooperative tensor GEMM

Add a reference page covering the two Metal GEMM paths, selection rules, current limitations, and planned follow-up work.

* lint

* Improve Metal cooperative tensor GEMM

* Optimize Metal cooperative tensor GEMM

* Update TVM Metal 4 support

* Add Metal backend internals documentation

* Harden Metal cooperative tensor lowering

* test older impl with new framework

* Clean up Metal transform exports

* Use macos-latest for Metal CI

* resolve comments and lint

---------

Co-authored-by: SiriusNEO <chaofan@deepseek.com>
This commit is contained in:
Yichen Yan
2026-07-28 16:17:15 +08:00
committed by GitHub
co-authored by SiriusNEO
parent aaf68d2e0b
commit 1545f0065d
38 changed files with 3591 additions and 265 deletions
+153 -18
View File
@@ -10,10 +10,18 @@ import tilelang.language as T
logging.getLogger("tilelang").setLevel(logging.WARNING)
BLOCK_CONFIGS = [
(16, 16, 16),
(32, 32, 16),
(32, 32, 32),
(64, 64, 32),
("simdgroup", 16, 16, 16, 128, 0, "row"),
("simdgroup", 32, 32, 16, 128, 0, "row"),
("simdgroup", 32, 32, 32, 128, 0, "row"),
("simdgroup", 64, 64, 32, 128, 0, "row"),
("ct_shared", 32, 64, 32, 128, 0, "row"),
("ct_shared", 64, 64, 32, 128, 0, "row"),
# Direct global cooperative tensor path. The K tile is the full problem K,
# so C is accumulated in cooperative-tensor registers and written once.
("ct_global", 64, 64, 0, 64, 0, "row"),
("ct_global", 64, 128, 0, 128, 0, "row"),
("ct_global", 64, 128, 0, 128, 4, "mlx"),
("ct_global", 64, 128, 0, 256, 4, "mlx"),
]
@@ -40,6 +48,95 @@ def matmul_simdgroup(M, N, K, block_M=64, block_N=64, block_K=32, dtype=T.float1
return gemm_kernel
@tilelang.jit
def matmul_cooperative_tensor_shared_c(
M,
N,
K,
block_M=64,
block_N=64,
block_K=32,
dtype=T.float16,
accum_dtype=T.float32,
):
@T.prim_func
def gemm_kernel(
A: T.Tensor((M, K), dtype),
B: T.Tensor((K, N), dtype),
C: T.Tensor((M, N), accum_dtype),
):
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
A_shared = T.alloc_shared((block_M, block_K), dtype, scope="shared")
B_shared = T.alloc_shared((block_K, block_N), dtype, scope="shared")
C_shared = T.alloc_shared((block_M, block_N), accum_dtype, scope="shared")
T.clear(C_shared)
for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=0):
T.copy(A[by * block_M, ko * block_K], A_shared)
T.copy(B[ko * block_K, bx * block_N], B_shared)
T.gemm(A_shared, B_shared, C_shared)
T.copy(C_shared, C[by * block_M, bx * block_N])
return gemm_kernel
@tilelang.jit
def matmul_cooperative_tensor_global(
M,
N,
K,
block_M=64,
block_N=64,
threads=128,
swizzle_panel=0,
swizzle_order="row",
dtype=T.float16,
accum_dtype=T.float32,
):
@T.prim_func
def gemm_kernel(
A: T.Tensor((M, K), dtype),
B: T.Tensor((K, N), dtype),
C: T.Tensor((M, N), accum_dtype),
):
tiles_n = T.ceildiv(N, block_N)
tiles_m = T.ceildiv(M, block_M)
use_mlx_swizzle = swizzle_panel and swizzle_order == "mlx"
grid_n = tiles_n * swizzle_panel if use_mlx_swizzle else tiles_n
grid_m = T.ceildiv(tiles_m, swizzle_panel) if use_mlx_swizzle else tiles_m
with T.Kernel(grid_n, grid_m, threads=threads) as (bx, by):
logical_bx = bx // swizzle_panel if use_mlx_swizzle else bx
logical_by = by * swizzle_panel + bx % swizzle_panel if use_mlx_swizzle else by
if swizzle_panel:
T.use_swizzle(panel_size=swizzle_panel, order=swizzle_order)
if use_mlx_swizzle:
if logical_by < tiles_m:
T.gemm(
A[logical_by * block_M : (logical_by + 1) * block_M, 0:K],
B[0:K, logical_bx * block_N : (logical_bx + 1) * block_N],
C[
logical_by * block_M : (logical_by + 1) * block_M,
logical_bx * block_N : (logical_bx + 1) * block_N,
],
clear_accum=True,
)
else:
T.gemm(
A[logical_by * block_M : (logical_by + 1) * block_M, 0:K],
B[0:K, logical_bx * block_N : (logical_bx + 1) * block_N],
C[
logical_by * block_M : (logical_by + 1) * block_M,
logical_bx * block_N : (logical_bx + 1) * block_N,
],
clear_accum=True,
)
return gemm_kernel
def _tflops(M, N, K, seconds):
return 2.0 * M * N * K / seconds / 1e12
@@ -62,11 +159,40 @@ def bench_torch_mps(M, N, K, warmup, repeats):
return _tflops(M, N, K, avg_s)
def bench_tilelang(M, N, K, block_M, block_N, block_K, warmup, repeats):
kernel = matmul_simdgroup(M, N, K, block_M, block_N, block_K)
def bench_mlx(M, N, K, warmup, repeats):
try:
import mlx.core as mx
except ImportError:
return None
a = mx.random.normal((M, K)).astype(mx.float16)
b = mx.random.normal((K, N)).astype(mx.float16)
mx.eval(a, b)
for _ in range(warmup):
c = a @ b
mx.eval(c)
t0 = time.perf_counter()
for _ in range(repeats):
c = a @ b
mx.eval(c)
return _tflops(M, N, K, (time.perf_counter() - t0) / repeats)
def bench_tilelang(mode, M, N, K, block_M, block_N, block_K, threads, swizzle_panel, swizzle_order, warmup, repeats):
output_dtype = T.float16
if mode == "ct_shared":
kernel = matmul_cooperative_tensor_shared_c(M, N, K, block_M, block_N, block_K, accum_dtype=output_dtype)
elif mode == "ct_global":
kernel = matmul_cooperative_tensor_global(
M, N, K, block_M, block_N, threads, swizzle_panel, swizzle_order, accum_dtype=output_dtype
)
else:
kernel = matmul_simdgroup(M, N, K, block_M, block_N, block_K, accum_dtype=output_dtype)
a = torch.randn(M, K, dtype=torch.float16, device="mps")
b = torch.randn(K, N, dtype=torch.float16, device="mps")
c = torch.zeros(M, N, dtype=torch.float32, device="mps")
c = torch.zeros(M, N, dtype=output_dtype.as_torch(), device="mps")
avg_s = _bench(lambda: kernel(a, b, c), warmup, repeats)
return _tflops(M, N, K, avg_s)
@@ -78,7 +204,7 @@ if __name__ == "__main__":
parser.add_argument("--k", type=int, default=4096)
parser.add_argument("--warmup", type=int, default=10)
parser.add_argument("--repeats", type=int, default=100)
parser.add_argument("--sweep", action="store_true", help="Sweep all block configs instead of using default (64,64,32)")
parser.add_argument("--sweep", action="store_true", help="Sweep all block configs instead of using default CT config")
args = parser.parse_args()
M, N, K = args.m, args.n, args.k
@@ -91,26 +217,35 @@ if __name__ == "__main__":
ref_tflops = bench_torch_mps(M, N, K, args.warmup, args.repeats)
print(f"PyTorch MPS (torch.mm fp16): {ref_tflops:.1f} TFLOPS")
mlx_tflops = bench_mlx(M, N, K, args.warmup, args.repeats)
if mlx_tflops is not None:
print(f"MLX matmul fp16: {mlx_tflops:.1f} TFLOPS")
print()
configs = BLOCK_CONFIGS if args.sweep else [(64, 64, 32)]
configs = BLOCK_CONFIGS if args.sweep else [("ct_global", 64, 128, 0, 128, 0, "row")]
print(f"{'block (M,N,K)':>16s} | {'TileLang':>14s} | {'Ratio':>6s}")
print("-" * 44)
print(f"{'path':>10s} | {'block (M,N,K)':>16s} | {'thr':>4s} | {'swizzle':>8s} | {'TileLang':>14s} | {'vs Torch':>8s} | {'vs MLX':>8s}")
print("-" * 88)
best_tflops = 0.0
best_config = configs[0]
for bM, bN, bK in configs:
for mode, bM, bN, bK, threads, swizzle_panel, swizzle_order in configs:
block_text = f"({bM},{bN},{bK if bK else 'all'})"
swizzle_text = f"{swizzle_panel}:{swizzle_order}" if swizzle_panel else "-"
try:
tl = bench_tilelang(M, N, K, bM, bN, bK, args.warmup, args.repeats)
ratio = tl / ref_tflops * 100
tag = ""
tl = bench_tilelang(mode, M, N, K, bM, bN, bK, threads, swizzle_panel, swizzle_order, args.warmup, args.repeats)
torch_ratio = tl / ref_tflops * 100
mlx_ratio = tl / mlx_tflops * 100 if mlx_tflops else None
if tl > best_tflops:
best_tflops = tl
best_config = (bM, bN, bK)
print(f"{f'({bM},{bN},{bK})':>16s} | {tl:>10.1f} TFLOPS | {ratio:>5.0f}%")
best_config = (mode, bM, bN, bK, threads, swizzle_panel, swizzle_order)
mlx_text = f"{mlx_ratio:>7.0f}%" if mlx_ratio is not None else " N/A"
print(
f"{mode:>10s} | {block_text:>16s} | {threads:>4d} | {swizzle_text:>8s} | "
f"{tl:>10.1f} TFLOPS | {torch_ratio:>7.0f}% | {mlx_text}"
)
except Exception as e:
print(f"{f'({bM},{bN},{bK})':>16s} | {'FAILED':>14s} | {e}")
print(f"{mode:>10s} | {block_text:>16s} | {threads:>4d} | {swizzle_text:>8s} | {'FAILED':>14s} | {e}")
if args.sweep:
print()
@@ -0,0 +1,642 @@
# TileLang Metal Backend Internals
This document explains how the TileLang Metal backend lowers a TileLang kernel
to Metal shader source, how the current GEMM fast path is implemented, which
backend work does not map one-to-one to CUDA, and where the next optimization
work should land.
This is not an Apple GPU architecture overview or a TileLang frontend tutorial.
Metal and Apple GPU background is introduced only where it affects lowering,
codegen, runtime behavior, or performance policy. Reference material is grouped
in Section 5.
How to read this document:
- To understand the current fast path, read Section 0 and Section 1.
- To work on the Metal backend or codegen, read Sections 1, 2, and 3, then use
Section 5 as reference.
- To write TileLang kernels for Metal, read Section 0 and Section 4.
- To find performance data, implementation files, or test commands, go to
Section 5.
## 0. TL;DR
- The current default fast path is **direct-global cooperative tensor**: A/B are
loaded from `device` memory without `threadgroup` staging, C accumulates in a
cooperative tensor destination, and the result is stored back to global C.
- Default benchmark configuration: `ct_global`, `block_M=64`, `block_N=128`,
`threads=128` (4 simdgroups), fp16 input, fp16 output, fp32 accumulation.
- On the current full-tile fp16 GEMM snapshot, TileLang is at or slightly above
MLX and about 5-6% behind PyTorch MPS. See Section 5.6.
- The most important mental model: **Metal is not CUDA**. On Apple unified
memory, `gmem -> smem` is not a performance hint. Shared staging should be
preserved when it carries real semantics such as layout, padding, or reuse.
- The other paths, shared cooperative tensor and simdgroup, still exist for
compatibility and semantic coverage. They are not the current performance
mainline. See Section 3.
## 1. The Fast Path, End to End
This section traces one concrete direct-global cooperative tensor kernel from
TileLang source to generated MSL and then to the performance snapshot. Later
sections explain the mechanisms behind this path.
### 1.1 The Kernel
```python
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
T.gemm(
A[by * block_M:(by + 1) * block_M, 0:K],
B[0:K, bx * block_N:(bx + 1) * block_N],
C[by * block_M:(by + 1) * block_M, bx * block_N:(bx + 1) * block_N],
clear_accum=True,
)
```
A, B, and C are all tile views over `global` tensors. There is no
`alloc_shared`, no `T.copy` staging, and no handwritten K-loop. For full-tile
GEMM, the backend expands the K reduction during lowering. This is the intended
frontend shape: express operand regions and dtypes, and let the backend own lane
mapping, operand construction, and accumulation.
### 1.2 Path Selection
When the backend sees this `T.gemm`, the instruction selection and Metal GEMM
lowering logic pick the direct-global cooperative tensor path. The detailed
selection rules are in Section 2.1. For this kernel, A/B are `global`, C is not
shared, and the tile shape and thread count allow cooperative tensor lowering.
For `M=128, N=256, K=128, block_M=64, block_N=128, threads=128`:
- `threads=128` means 4 simdgroups.
- The 4 simdgroups cover a `64 x 128` C tile as a `2 x 2` partition.
- Each simdgroup owns a `32 x 64` C sub-tile.
- Each `32 x 64` sub-tile is covered by `(32 / 16) x (64 / 32) = 4` op-level
fragments, corresponding to 4 destination cooperative tensors.
### 1.3 The MPP op: `matmul2d(M, N, K)`
The core operation is an MPP tensor op. The generated MSL contains:
```cpp
constexpr auto __pct_desc = mpp::tensor_ops::matmul2d_descriptor(
16, 32, 16, /*trans_a=*/false, /*trans_b=*/false, /*accumulate=*/true,
mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate);
mpp::tensor_ops::matmul2d<__pct_desc, metal::execution_simdgroup> __pct_op;
```
In this document, `matmul2d(M, N, K)` refers to the first three shape parameters
of the descriptor, in M / N / K order. A single MPP op produces an `M x N`
destination fragment and reduces along K. These parameters are not TileLang
`block_M` / `block_N`, and they are not the whole threadgroup tile shape. The
current fast path specializes to `matmul2d(16, 32, 16)`. Larger simdgroup and
threadgroup tiles are covered by multiple op-level fragments, such as the 4
destination cooperative tensors above.
### 1.4 The Generated MSL
The following is a representative snippet for the shape above. Repeated
fragments are omitted for readability. It is generated with
`accum_dtype=T.float16`, so C is `device half*` and the final store contains a
`half4(...)` cast. Existing codegen tests use `accum_dtype=T.float32` by default,
which generates `device float* C` and no fp32-to-fp16 store cast. To reproduce
this fp16-output source exactly, pass `accum_dtype=T.float16`; see Section 5.8.
Kernel signature and execution attributes:
```cpp
[[kernel, max_total_threads_per_threadgroup(128)]] void main_kernel(
const device half* __restrict A [[ buffer(0) ]],
const device half* __restrict B [[ buffer(1) ]],
device half* __restrict C [[ buffer(2) ]],
uint3 blockIdx [[threadgroup_position_in_grid]],
uint3 __gridDim [[threadgroups_per_grid]],
uint3 threadIdx [[thread_position_in_threadgroup]],
uint __simd_group_id [[simdgroup_index_in_threadgroup]]) {
const ushort __lane = __metal_get_thread_index_in_simdgroup(ushort());
...
```
Destination cooperative tensors are created and zeroed before the K-loop. C
storage is elided: C does not round-trip through a `thread` array and instead
persists in cooperative tensor destination objects.
```cpp
auto __pct_c0 = __pct_op.get_destination_cooperative_tensor<...>();
// __pct_c1 / __pct_c2 / __pct_c3 are created the same way.
TILELANG_PRAGMA_UNROLL for (ushort __i = 0; __i < 16; __i++) __pct_c0[__i] = 0.0f;
```
Inside the K-loop, A/B are loaded from `device` memory into thread-private
temporary storage and then moved into MPP cooperative tensor operands. There is
no `threadgroup` allocation and no barrier.
```cpp
for (int k_outer = 0; k_outer < 8; ++k_outer) {
const device half* __src =
(const device half*)(&(A[((blockIdx.y * 8192) +
((__simd_group_id & 1) * 4096)) +
(k_outer * 16)]));
// half4 vector loads into A_local ...
// B loads follow the same pattern ...
auto __ct_a = __pct_op.get_left_input_cooperative_tensor<half, half, float>();
auto __ct_b = __pct_op.get_right_input_cooperative_tensor<half, half, float>();
// A_local / B_local -> __ct_a / __ct_b
__pct_op.run(__ct_a, __ct_b, __pct_c0); // fp32 accumulation
}
```
After the K-loop, C is stored. For fp16 output, the fp32 destination values are
cast on store:
```cpp
device half* __dst = (device half*)(&(C[...]));
*(device half4*)(&__dst[__r0 * 256 + __c0]) =
half4(*(thread float4*)(&__pct_c0[0]));
```
Important source-shape properties:
- A/B are `const device half* __restrict`; C is `device half* __restrict`. This
preserves read-only and noalias information for the Metal compiler.
- simdgroup id and lane id come from Metal execution attributes
(`__simd_group_id`, `__metal_get_thread_index_in_simdgroup`) instead of being
repeatedly derived from `threadIdx`.
- Addresses have the shape `base + tile offset + simdgroup offset + K-loop
induction`.
- `mpp::...matmul2d` destinations accumulate in fp32; fp16 output is handled at
store.
- There is no `threadgroup half` declaration and no barrier.
### 1.5 Why This Is Not a CUDA Translation
This path is not a step-by-step translation of the common CUDA
`global -> shared -> ldmatrix -> mma` pipeline:
```text
CUDA: device A/B -> shared -> ldmatrix -> mma -> shared/reg -> device C
Metal: device A/B ------------> MPP operand -> matmul2d -> destination CT -> device C
```
The Metal backend chooses a source shape that matches MPP directly. On Apple
unified memory, moving A/B through `threadgroup` before feeding the tensor op is
usually an extra copy and barrier when it carries no semantic value. This is why
the shared path in Section 3 is correct but slower for the current full-tile
GEMM snapshot.
## 2. How the Backend Gets There
This section explains how the backend gets from `T.gemm` to the MSL shown in
Section 1. It only covers mechanisms needed for that mainline path.
### 2.1 Path Selection
Path selection has two stages.
**Stage 1: instruction kind (`SelectInst`, `src/metal/op/gemm.cc`).** This stage
only chooses between `metal.cooperative_tensor` and `metal.simdgroup`. It looks
at the lowering target capability, C scope, tile shape, warp policy, and warp
count. It does not look at whether A/B are global or shared.
| Condition | Instruction kind |
| --- | --- |
| C is in `local.fragment` or `metal.simdgroup` | `metal.simdgroup` |
| Target has the `metal4` key and `CanUseCooperativeTensor(policy, M, N, K, warps)` is true | `metal.cooperative_tensor` |
| Otherwise | `metal.simdgroup` as a safe lowering fallback |
`CanUseCooperativeTensor` is a pure shape and policy check: `M % 16 == 0`,
`N % 32 == 0`, `K % 16 == 0`, and the warp policy must be able to split
`num_warps` into a legal `m_warp x n_warp` partition. `threads` must be a
multiple of 32 because each simdgroup has 32 lanes.
`target="metal"` is normalized to include the `metal4` key only when the local
macOS version, SDK, and Apple GPU generation are known to support Metal 4.
Source-only tests or cross-builds can pass an explicit target with
`keys=["metal", "gpu", "metal4"]` when they intentionally want MPP source.
**Stage 2: cooperative tensor dataflow (`GemmMetal.lower`,
`tilelang/metal/op/gemm/gemm_metal.py`).** After Stage 1 selects
`metal.cooperative_tensor`, the Python lowering chooses the concrete dataflow
from A/B scope:
| A/B scope | Dataflow | Notes |
| --- | --- | --- |
| A/B are `global` (`is_gemm_gg`) | direct-global cooperative tensor | Current fast path, Section 1 |
| A/B are `shared` (`is_gemm_ss`) | shared cooperative tensor | Preserves shared semantics, Section 3.1 |
C in cooperative tensor or fragment form uses destination storage elision. C in
shared scope uses writeback. Warp partitioning prefers a balanced grid of
16x32 op-level fragments per simdgroup.
One current capability boundary is important: all of the above choices happen
at **lowering time** based on shape, scope, and policy. They are not runtime or
hardware capability checks. The runtime guard in Section 5.1 is separate: it
only chooses which MSL language version to request when compiling the generated
source. It cannot rewrite already generated MPP source back to simdgroup. If
`SelectInst` chose cooperative tensor, the generated source is MPP source.
Instruction-level runtime capability fallback still needs a separate design.
### 2.2 Surrounding Lowering
The launch, allocation, loop, copy, and swizzle lowering around GEMM determine
what source shape the intrinsic sequence sits in. For the direct-global path:
- **Launch**: `T.Kernel(..., threads=N)` lowers to a Metal threadgroup launch,
emitting `max_total_threads_per_threadgroup(N)`, `blockIdx` / `threadIdx`
style indices, simdgroup id, and lane id. Matrix work is still partitioned by
32-lane simdgroups; there is no CUDA WGMMA-style 128-thread warpgroup.
- **Allocation**: the direct-global path does not allocate `threadgroup`
buffers. A/B use `thread` temporary arrays. C persists in
`metal.cooperative_tensor` destination objects when storage elision applies.
The full scope-to-object mapping is in Section 5.2.
- **Loop**: the K-loop generates affine addresses of the form
`base + tile offset + simdgroup offset + K-loop induction`. It does not rely
on CUDA-style software pipelining or make shared staging the inner-loop
policy. Pointer induction and address hoisting remain optimization targets
(Section 5.5).
- **Copy / sync**: the direct-global path is device load/store plus MPP operand
construction. With no shared dataflow, it avoids threadgroup copy and barrier
operations.
- **Swizzle**: block-level swizzle is supported. `T.use_swizzle(panel_size,
order="mlx")` maps to `rasterization2DMLX`, canonicalizing physical block
index to logical tile index and avoiding repeated physical-index arithmetic in
generated MSL. This is block/grid-level swizzle, not CUDA-style thread-level
or per-thread swizzle.
### 2.3 Metal-Specific Work (No CUDA Counterpart)
The direct-global path includes backend work that does not have a direct CUDA
counterpart:
- **Direct-global cooperative tensor lowering**: A/B become MPP operands from
`device` memory without explicit shared staging.
- **Cooperative tensor destination management**: C accumulates in MPP
destination objects; small C cooperative tensor storage can elide a
thread-local array (Section 1.4).
- **fp16 output / fp32 accumulation split**: MPP destinations use fp32, and the
store to fp16 C performs the cast.
- **MLX-style block rasterization**: physical block indices are canonicalized to
logical tile indices.
- **simdgroup id / lane id lowering**: Metal execution attributes are used
instead of re-deriving them from thread id.
- **Metal 4 target capability gate**: cooperative tensor lowering requires the
`metal4` target key.
- **Metal 4 runtime language-version guard**: MSL 4.0 is requested only when the
SDK, runtime OS, and device family support it.
This work is closer to Metal backend canonicalization and source-shape
construction than to mechanical CUDA op lowering.
## 3. The Other Paths
Besides the fast path, the backend keeps two additional paths. They are correct
and tested, but they are not the current performance mainline.
### 3.1 Shared Cooperative Tensor
If the frontend IR already contains shared staging, the backend lowers it as a
shared cooperative tensor path:
```text
global A/B -> threadgroup A/B -> cooperative tensor operand -> MPP -> C
```
This preserves the shape of CUDA-style kernels and is appropriate when shared
memory carries real semantics: layout transforms, padding or edge zero-fill,
multi-consumer reuse, threadgroup communication, or address patterns that the
direct-global path cannot currently express.
If shared memory is only staging, it is an extra copy plus barrier on Metal. In
the current measurements, it is slower than the direct-global path.
### 3.2 Simdgroup Fallback
The simdgroup path lowers GEMM to simdgroup matrix intrinsics:
```text
threadgroup A/B -> simdgroup_load -> simdgroup_multiply_accumulate -> simdgroup_store
```
It is a compatibility path used when C is in `local.fragment` or
`metal.simdgroup`, or when lowering-time cooperative tensor conditions are not
met. Compared with cooperative tensor, it behaves more like an explicit
fragment backend: the backend manages fragment layout, load/store, and
accumulation/store policy. It is therefore not the M5 / Metal 4 optimization
mainline.
### 3.3 Why These Paths Are Kept
- Existing TileLang kernels may already use CUDA-style staging.
- Shared staging may carry layout transform, padding, edge zero-fill, or
multi-consumer reuse semantics.
- Some non-contiguous or edge-heavy workloads cannot yet be expressed by the
direct-global path.
- Older Apple GPUs and non-Metal-4 runtimes need a complete simdgroup fallback
story, although instruction-level runtime fallback is not implemented yet.
The key requirement is conservatism: a shared buffer can be treated as pure
staging, or bypassed by a future pass, only when the compiler can prove it has
no extra semantics. If that proof fails, the shared path must be preserved.
## 4. Contracts and Guidance
### 4.1 Frontend Contract
For the Metal backend, the frontend should express stable program semantics:
- memory scope;
- GEMM operand region;
- block-level swizzle;
- dtype and output dtype;
- required boundary conditions.
These fields should describe the program, not prescribe the final MSL source
shape. Lane mapping, operand construction, register tile shape, and instruction
selection belong to the backend and codegen. The current Metal control surface
is summarized in the concept map in Section 5.1.
### 4.2 Bad CUDA Assumptions
Do not carry CUDA GEMM assumptions directly into Metal:
- Do not treat `gmem -> smem` as a performance hint. Its semantics and bypass
conditions are described in Section 3.
- Do not add shared staging just to mimic a CUDA tensor-core pipeline.
- Do not treat `local` / rmem staging as free. It can compete with threadgroup
memory for the same class of on-chip storage budget and can increase register
pressure or spilling (Section 5.2).
- Do not expose CUDA warpgroup, thread-level layout, or per-thread swizzle as
Metal frontend semantics.
### 4.3 Lowering and Codegen Rules
Metal lowering should keep path boundaries clear:
- maintain the direct-global cooperative tensor fast path;
- preserve shared path semantics;
- do not mix simdgroup and cooperative tensor policy into one path;
- report unsupported shapes clearly or use a safe fallback.
Metal codegen should focus on MSL source shape:
- emit simple affine addresses;
- use explicit simdgroup id and lane id;
- preserve `const` for read-only buffers and `__restrict` for noalias buffers;
- avoid repeated physical swizzle arithmetic;
- preserve the fp32 accumulation to fp16 output store-cast path.
## 5. Reference
### 5.1 CUDA / Metal Concept Map and Capability
| CUDA / NVIDIA concept | Metal concept | TileLang Metal mapping | Notes |
| --- | --- | --- | --- |
| Grid | Grid of threadgroups | `T.Kernel` grid dims | Mostly direct mapping |
| CTA / thread block | Threadgroup | `T.Kernel(..., threads=N)` | Matrix execution is still split by simdgroup |
| Thread | Thread | thread index lowering | Mostly direct mapping |
| Warp | Simdgroup | simdgroup id / lane id | Current backend assumes 32 lanes |
| Warpgroup (128-thread WGMMA) | No direct exposed equivalent | Not exposed | Cooperative tensor is organized around simdgroup execution context |
| Shared memory | Threadgroup memory | `shared` scope | Semantic scope; performance policy is discussed in Section 3 |
| Register / fragment | thread storage / simdgroup matrix / cooperative tensor destination | `local` / `metal.simdgroup` / `metal.cooperative_tensor` | Do not reuse CUDA fragment assumptions directly |
| Tensor Core MMA | simdgroup matrix or MPP tensor op | `metal.simdgroup` / `metal.cooperative_tensor` | M5+ prefers cooperative tensor |
| `ldmatrix` / shared-to-MMA | cooperative tensor load / MPP operand | codegen intrinsic emission | Direct-global can bypass explicit shared staging |
| Thread-level layout / per-thread swizzle | No exposed equivalent | Not user-controlled | No CUDA-style per-thread data layout or register swizzle control surface |
| Block swizzle | Rasterization | `T.use_swizzle(..., order="mlx")` | Block-level swizzle |
The key difference is that the Metal backend currently does not expose
CUDA-style thread-level layout or per-thread swizzle. It exposes block/grid-level
rasterization swizzle; lane mapping is decided by the backend and the Metal
compiler, not by stable frontend semantics.
The runtime guard in `metal_module.mm` has one job: avoid requesting MSL 4.0
unconditionally. It selects `MTLLanguageVersion4_0` only when the SDK, runtime
OS, and device family support Metal 4; otherwise it compiles with
`MTLLanguageVersion2_3`. This is not an instruction-level fallback. It cannot
rewrite generated MPP source into simdgroup source. Cooperative tensor vs.
simdgroup is a lowering-time target/shape/scope/policy decision (Section 2.1),
separate from runtime language-version selection. A complete runtime fallback
still needs a separate multi-version or retry design.
### 5.2 Memory Scopes and Unified Memory
| TileLang scope | Metal address space / object | Purpose |
| --- | --- | --- |
| `global` | `device` | Input and output tensors |
| `shared` | `threadgroup` | Threadgroup scratchpad / staging; barrier required for cross-thread visibility |
| `local` | `thread` private array | Per-thread temporary array; may occupy registers or spill |
| `local.var` | `thread` scalar | Scalar index or loop state |
| `metal.simdgroup` | `simdgroup_*8x8` object | simdgroup matrix fragment |
| `metal.cooperative_tensor` | MPP operand / destination object | Cooperative tensor path; shape comes from the descriptor and is not fixed 8x8 |
Apple platforms use unified memory, but that does not erase MSL address spaces
or make CUDA memory policy carry over unchanged. One migration assumption to
avoid is physical rmem / smem separation. `thread` storage and `threadgroup`
memory are distinct logical address spaces, but physical placement and spilling
are managed by the Metal compiler and Apple GPU runtime. For TileLang Metal
lowering, treat `local` private arrays and `shared` threadgroup staging as
competing for the same class of on-chip storage budget, rather than assuming a
CUDA-like model where the register file is free and shared memory is budgeted
separately.
### 5.3 Execution Model and Simdgroup Builtins
Execution hierarchy: grid (threadgroup grid), threadgroup (closest to a CUDA
CTA), thread, and simdgroup (32-lane SIMD execution unit). Threadgroup size and
matrix execution group are distinct: TileLang Metal can launch 128 or 256
threads, while cooperative tensor and simdgroup matrix execution are organized
around 32-lane simdgroups. There is no exposed 128-thread warpgroup equivalent.
Supported simdgroup matrix builtins include `metal.simdgroup` scope,
`make_filled_simdgroup_matrix`, `simdgroup_load`,
`simdgroup_multiply_accumulate`, and `simdgroup_store`. `simdgroup_*8x8` is the
simdgroup matrix object shape; it is distinct from the cooperative tensor MPP
descriptor shape described in Section 1.3.
### 5.4 Implemented Features
| Feature | Status | Notes |
| --- | --- | --- |
| Metal 4 target capability gate | Implemented | Cooperative tensor lowering requires the `metal4` target key |
| Metal 4 runtime guard | Implemented | Requests MSL 4.0 only when SDK, OS, and device family support it |
| Cooperative tensor GEMM lowering | Implemented | M5+ mainline |
| Direct-global A/B cooperative tensor | Implemented / default | Current recommended performance path, Section 1 |
| Shared A/B cooperative tensor | Implemented | Compatible with CUDA-style staging, Section 3.1 |
| fp16 output handling | Implemented | fp32 accumulation plus store cast, Section 1.4 |
| C cooperative tensor storage elision | Implemented | Reduces thread-local C round-trip |
| MLX-style block swizzle | Implemented | Block-level rasterization |
| simdgroup GEMM | Implemented | Compatibility path, Section 3.2 |
| simdgroup id / lane id lowering | Implemented | Also used by cooperative tensor kernels |
| `const` / `__restrict` parameter emission | Implemented | Improves MSL alias information |
### 5.5 Known Limitations and Roadmap
Current limitations:
- The direct-global path mainly covers full-tile GEMM.
- Edge masking and partial tile support are not complete.
- The `gmem -> smem` bypass pass is not implemented.
- MLX-style 8-simdgroup lowering is not implemented.
- Pointer induction and address hoisting still have room for improvement.
- CUDA-style thread-level layout and per-thread swizzle control are not exposed
(Section 5.1).
- The simdgroup path is not yet a complete old-device strategy.
- The primary validated path is fp16 input, fp16 output, fp32 accumulation.
- There is no multi-version runtime fallback. Path selection is a lowering-time
target/shape/scope/policy decision, and the runtime guard only selects the MSL
language version (Sections 2.1 and 5.1).
Planned work:
- **Shared-staging bypass pass**: detect pure staging of the form
`global A/B -> shared A/B -> T.gemm`, and rewrite it to the direct-global path
when semantics are provably unchanged. Mathematically, this is conditional
substitution:
```text
S[s] = G[phi(s)]
T.gemm(... S[psi(t)] ...) => T.gemm(... G[phi(psi(t))] ...)
```
The proof must rule out layout transforms, padding, edge fill,
multi-consumer reuse, and threadgroup communication semantics. If the proof
fails, the shared path must be preserved.
- **MLX-style 8-simdgroup cooperative tensor lowering**: `BM=64`,
`BN=128`, 8 simdgroups / 256 threads, each simdgroup owning a `32 x 32` C
tile covered by multiple `matmul2d(16, 32, 16)` op-level fragments.
- **Pointer induction and address hoisting**: hoist A/B/C tile base pointers,
simdgroup offsets, and K-loop pointer increments to reduce repeated affine
expressions in the hot path.
- **Edge masking and partial tile support**: predicated A/B load or zero-fill,
predicated C store, out-of-bounds return for padded physical grids after
swizzle, and branch-free fast paths for aligned cases.
- **Metal-specific policy and tuner**: distinguish at least four policy classes:
4-simdgroup conservative direct-global, 8-simdgroup MLX-style,
shared path for layout / fusion / edge-heavy workloads, and simdgroup
fallback.
### 5.6 Current Performance Snapshot
These numbers are a snapshot of the current development environment and current
commit. They should not be generalized to all Apple GPUs, OS versions, SDKs, or
library versions. Reproduction commands are in Section 5.8.
| Item | Value |
| --- | --- |
| Date | 2026-06-20 |
| Machine | MacBook Air, Mac17,3 |
| Chip | Apple M5 |
| GPU | Apple M5, 10-core GPU |
| Metal support | Metal 4 |
| macOS | 26.5.1 |
| Xcode | 26.4.1 |
| Python | 3.12.13 |
| PyTorch | 2.12.1, MPS enabled |
| MLX | 0.31.2 |
| TileLang branch | `metal-gemm-perf` |
| TileLang measurement commit | `6f952ed9` |
| TVM measurement submodule | `11c1968acf` |
Benchmark conditions: fp16 input, fp16 output, internal fp32 accumulation,
full-tile shapes, default path `ct_global`, `block_M=64`, `block_N=128`,
`threads=128`. Aggregation method: after warmup, timed runs are executed
consecutively, and TFLOPS is computed from the average latency over
`repeats=100`. This is not median or best-of-N.
| Shape | PyTorch MPS | MLX | TileLang | TileLang / MLX |
| --- | ---: | ---: | ---: | ---: |
| 4096 x 4096 x 4096 | 13.6 TFLOPS | 12.7 TFLOPS | 12.9 TFLOPS | 101% |
| 2048 x 2048 x 2048 | 13.6 TFLOPS | 11.5 TFLOPS | 12.8 TFLOPS | 112% |
| 4096 x 2048 x 4096 | 13.6 TFLOPS | 12.5 TFLOPS | 12.8 TFLOPS | 102% |
| 2048 x 4096 x 4096 | 13.4 TFLOPS | 12.4 TFLOPS | 12.6 TFLOPS | 102% |
Conclusion: on these full-tile shapes, TileLang is at or slightly above MLX and
about 5-6% behind PyTorch MPS. The shared staging path is correct but slower.
MLX-style swizzle is correct but is not currently the fastest default strategy.
### 5.7 Implementation Map
| Area | File | Responsibility |
| --- | --- | --- |
| Language API | `tilelang/language/gemm_op.py` | Metal GEMM operand and offset handling |
| Annotation | `tilelang/language/annotations.py` | `T.use_swizzle(..., order="mlx")` mapping |
| Builtins | `tilelang/language/builtin.py` | cooperative tensor and simdgroup builtin wrappers |
| GEMM lowering | `tilelang/metal/op/gemm/gemm_metal.py` | direct-global/shared dataflow split, warp partition, dtype policy |
| Intrinsic helper | `tilelang/metal/intrinsics/metal_macro_generator.py` | simdgroup and cooperative tensor intrinsic emission |
| Metal pipeline | `tilelang/metal/pipeline.py` | Metal-specific pipeline transform placement |
| Fragment rewrite | `tilelang/metal/transform/metal_fragment_to_simdgroup.py` | legacy fragment accumulator to `metal.simdgroup` |
| Core GEMM op | `src/op/gemm.h`, `src/op/gemm.cc` | GEMM op metadata |
| Metal op lowering | `src/metal/op/gemm.cc`, `src/metal/op/utils.h` | `SelectInst`, validation, scope utilities |
| Metal codegen | `src/metal/codegen/codegen_metal.cc`, `.h` | MSL emission |
| TVM runtime | `3rdparty/tvm/src/runtime/metal/metal_module.mm` | guarded `MTLLanguageVersion4_0` selection |
| Runtime tests | `testing/python/metal/test_metal_gemm_v2.py` | Metal correctness |
| Codegen tests | `testing/python/metal/test_metal_gemm_v2_linux.py` | source-level Metal codegen |
| Simdgroup tests | `testing/python/metal/test_metal_simdgroup_store.py` | simdgroup direct store |
| Benchmark | `benchmark/matmul_metal/benchmark_matmul_metal.py` | PyTorch / MLX / TileLang comparison |
### 5.8 Developer Checklist
When changing lowering, check:
- whether `global/shared/local/metal.simdgroup/metal.cooperative_tensor` scope
semantics changed;
- whether CUDA shared staging assumptions leaked into the Metal default path;
- whether the direct-global cooperative tensor fast path is preserved;
- whether unsupported shapes produce a clear error or a safe fallback.
When changing codegen, check:
- whether read-only buffers keep `const` and noalias buffers keep `__restrict`;
- whether simdgroup id and lane id are still explicit;
- whether generated MSL avoids repeated physical swizzle arithmetic;
- whether the fp32 accumulation to fp16 output store-cast path still works.
When changing TVM / Metal runtime code, check:
- whether MSL 4.0 is requested only when SDK, runtime OS, and hardware support
Metal 4;
- whether the MSL 2.3 fallback is preserved and non-Metal-4 devices are not
regressed.
- whether `target="metal"` adds the `metal4` key only on supported local
hardware, and source-only tests use an explicit `metal4` target when needed.
To reproduce the fp16-output MSL snippet in Section 1.4 from the repository root
(codegen-only, no Metal runtime required):
```bash
TILELANG_DISABLE_CACHE=1 python - <<'PY'
import tilelang
import tilelang.language as T
from tilelang import tvm
from testing.python.metal.test_metal_gemm_v2_linux import matmul_gemm_v2_global_c
func = matmul_gemm_v2_global_c(
128, 256, 128, 64, 128,
dtype=T.float16,
accum_dtype=T.float16,
threads=128,
)
metal4_target = tvm.target.Target({"kind": "metal", "keys": ["metal", "gpu", "metal4"]})
with tvm.transform.PassContext(), metal4_target:
print(tilelang.lower(func, target=metal4_target).kernel_source)
PY
```
The important detail is the output buffer dtype: `C` is `T.float16`, so the
generated source uses `device half* C` and emits the fp32-to-fp16 store cast
shown in Section 1.4.
Recommended tests:
```bash
TILELANG_DISABLE_CACHE=1 python -m pytest testing/python/metal/test_metal_gemm_v2.py -q -x
TILELANG_DISABLE_CACHE=1 python -m pytest testing/python/metal/test_metal_gemm_v2_linux.py -q -x
TILELANG_DISABLE_CACHE=1 python -m pytest testing/python/metal/test_metal_simdgroup_store.py -q -x
```
Default benchmark:
```bash
TILELANG_DISABLE_CACHE=1 python benchmark/matmul_metal/benchmark_matmul_metal.py \
--m 4096 --n 4096 --k 4096 --warmup 10 --repeats 100
```
+1
View File
@@ -71,6 +71,7 @@ deeplearning_operators/deepseek_mla
compiler_internals/letstmt_inline
compiler_internals/inject_fence_proxy
compiler_internals/tensor_checks
compiler_internals/metal_tilelang_development
:::
:::{toctree}
File diff suppressed because it is too large Load Diff
+44 -1
View File
@@ -28,6 +28,7 @@
#include <string>
#include <unordered_map>
#include <unordered_set>
#include "target/source/codegen_c.h"
@@ -37,6 +38,7 @@ namespace codegen {
class CodeGenTileLangMetal final : public CodeGenC {
public:
explicit CodeGenTileLangMetal(Target target);
std::string Finish() final;
// override print thread tag.
void PrintArgUnionDecl();
void AddFunction(const GlobalVar &gvar, const PrimFunc &func) final;
@@ -46,6 +48,9 @@ public:
void PrintStorageSync(const CallNode *op) final; // NOLINT(*)
void PrintType(DataType t, std::ostream &os) final; // NOLINT(*)
void BindThreadIndex(const IterVar &iv) final; // NOLINT(*)
void VisitExpr_(const BufferLoadNode *op,
std::ostream &os) final; // NOLINT(*)
void VisitStmt_(const BufferStoreNode *op) final; // NOLINT(*)
// print load of single element
void PrintVecElemLoad(const std::string &vec, DataType t, int i,
std::ostream &os) final; // NOLINT(*)
@@ -54,8 +59,12 @@ public:
const std::string &value) final;
// overload visitor
void VisitStmt_(const AllocBufferNode *op) final; // NOLINT(*)
void VisitStmt_(const AttrStmtNode *op) final; // NOLINT(*)
void VisitStmt_(const ForNode *op) final; // NOLINT(*)
void VisitExpr_(const SelectNode *op, std::ostream &os) final; // NOLINT(*)
void VisitExpr_(const BroadcastNode *op, std::ostream &os) final; // NOLINT(*)
void VisitExpr_(const AddNode *op, std::ostream &os) final; // NOLINT(*)
void VisitExpr_(const CastNode *op, std::ostream &os) final; // NOLINT(*)
void VisitExpr_(const CallNode *op, std::ostream &os) final; // NOLINT(*)
void VisitExpr_(const FloatImmNode *op, std::ostream &os) final; // NOLINT(*)
@@ -63,7 +72,41 @@ public:
using CodeGenC::PrintType;
private:
std::unordered_map<const VarNode *, std::string> simdgroup_dtype_;
std::string GetAddrSpaceOf(const PrimExpr &ptr_expr) const;
std::string GetPointeeTypeOf(const PrimExpr &ptr_expr,
const std::string &fallback);
bool IsThreadIdxXExpr(const PrimExpr &expr) const;
bool IsBlockIdxExpr(const PrimExpr &expr, int dim) const;
bool IsConstIntExpr(const PrimExpr &expr, int64_t value) const;
bool IsMlxPanelRemainderExpr(const PrimExpr &expr) const;
bool IsMlxPanelRowExpr(const PrimExpr &expr) const;
bool IsMlxLogicalBlockXExpr(const PrimExpr &expr) const;
bool IsMlxLogicalBlockYExpr(const PrimExpr &expr) const;
bool TryPrintMlxLogicalYAffineExpr(const PrimExpr &expr, std::ostream &os);
bool TryPrintMlxSwizzleExpr(const PrimExpr &expr, std::ostream &os);
bool TryPrintSimdgroupIndexExpr(const CallNode *op, std::ostream &os);
void PrintSimdgroupIndexExpr(int64_t group_mask, int64_t group_shift,
std::ostream &os) const;
void EnsureFragmentLaneVars();
void EnsureCooperativeTensorBuffer(const Var &var);
std::unordered_map<Var, std::string, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
simdgroup_dtype_;
std::unordered_map<Var, std::string, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
cooperative_tensor_dtype_;
std::unordered_set<Var, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
ct_c_inlined_;
std::unordered_set<Var, ffi::ObjectPtrHash, ffi::ObjectPtrEqual>
ct_c_storage_elided_;
Var thread_idx_x_var_;
Var block_idx_x_var_;
Var block_idx_y_var_;
int active_mlx_swizzle_panel_{0};
int active_mlx_swizzle_log_{0};
bool emitted_metal_simdgroup_id_{false};
bool emitted_frag_lane_vars_{false};
bool uses_cooperative_tensor_{false};
bool needs_fragment_lane_vars_{false};
int thread_index_bits_{32};
int thread_work_dim_{0};
Target target_;
+131
View File
@@ -7,6 +7,7 @@
#include "metal/op/utils.h"
#include "metal/target_utils.h"
#include "op/builtin.h"
#include "op/utils.h"
#include <tvm/tirx/builtin.h>
@@ -25,6 +26,11 @@ bool CheckSIMDGroupCopy(const CopyNode &op) {
(IsSharedBuffer(op.dst) || IsGlobalBuffer(op.dst));
}
bool CheckCooperativeTensorCopy(const CopyNode &op) {
return IsCooperativeTensorBuffer(op.src) &&
(IsSharedBuffer(op.dst) || IsGlobalBuffer(op.dst));
}
Stmt LowerSIMDGroupCopy(const CopyNode &op, const LowerArgs &lower_args,
arith::Analyzer *analyzer) {
(void)analyzer;
@@ -104,6 +110,128 @@ Stmt LowerSIMDGroupCopy(const CopyNode &op, const LowerArgs &lower_args,
return SeqStmt(stmts);
}
Stmt LowerCooperativeTensorCopy(const CopyNode &op, const LowerArgs &lower_args,
arith::Analyzer *analyzer) {
(void)analyzer;
TVM_FFI_ICHECK(IsCooperativeTensorBuffer(op.src));
int total_elements = 1;
for (auto s : op.src->shape) {
auto imm = s.as<IntImmNode>();
TVM_FFI_ICHECK(imm) << "cooperative_tensor buffer must have constant shape";
total_elements *= imm->value;
}
constexpr int kTileSize = 16;
constexpr int kTileElems = kTileSize * kTileSize;
TVM_FFI_ICHECK(total_elements % kTileElems == 0)
<< "cooperative_tensor buffer size must be multiple of " << kTileElems
<< ", got " << total_elements;
TVM_FFI_ICHECK(op.dst_range.size() == 2)
<< "Expected 2D destination for cooperative_tensor store";
PrimExpr dst_row_base = op.dst_range[0]->min;
PrimExpr dst_col_base = op.dst_range[1]->min;
PrimExpr dst_stride = op.dst->shape[op.dst->shape.size() - 1];
int warp_size = TargetMetalGetWarpSize(lower_args.target);
const auto *block_size_imm =
lower_args.thread_bounds->extent.as<IntImmNode>();
TVM_FFI_ICHECK(block_size_imm)
<< "cooperative_tensor copy requires constant thread bounds";
int block_size = block_size_imm->value;
int num_warps = block_size / warp_size;
PrimExpr warp_id = FloorDiv(lower_args.thread_index, warp_size);
const auto *m_imm = op.src_range[0]->extent.as<IntImmNode>();
const auto *n_imm = op.src_range[1]->extent.as<IntImmNode>();
TVM_FFI_ICHECK(m_imm && n_imm)
<< "cooperative_tensor copy requires constant extents";
int M = m_imm->value;
int N = n_imm->value;
int kMPerWarp = kTileSize;
int kNPerWarp = kTileSize * 2;
int m_warp = 1, n_warp = num_warps;
int max_m = M / kMPerWarp;
int max_n = N / kNPerWarp;
float ideal = N > 0 ? static_cast<float>(M) / N : 1.f;
float best_score = std::numeric_limits<float>::max();
for (int m = 1; m <= std::min(num_warps, max_m); ++m) {
if (num_warps % m != 0) {
continue;
}
int n = num_warps / m;
if (n > max_n) {
continue;
}
float m_per = static_cast<float>(M) / (m * kMPerWarp);
float n_per = static_cast<float>(N) / (n * kNPerWarp);
float score = std::abs(m_per / n_per - ideal);
if (score < best_score) {
best_score = score;
m_warp = m;
n_warp = n;
}
}
int elems_per_thread = total_elements / (num_warps * warp_size);
int warp_M = M / m_warp;
int warp_N = N / n_warp;
int warp_tiles = elems_per_thread / (kTileSize * kTileSize / warp_size);
int kTileN = warp_N;
int kTileM = kTileSize;
if (warp_tiles > 0 && warp_M > kTileSize) {
kTileN = warp_N;
kTileM = kTileSize;
}
if (kTileN > warp_N) {
kTileN = warp_N;
}
int warp_row_tiles = warp_M / kTileM;
int warp_col_tiles = warp_N / kTileN;
TVM_FFI_ICHECK(warp_row_tiles > 0 && warp_col_tiles > 0)
<< "Cannot partition " << M << "x" << N << " matrix across " << m_warp
<< "x" << n_warp << " warps";
int tile_elems_per_thread = kTileM * kTileN / warp_size;
TVM_FFI_ICHECK(warp_row_tiles * warp_col_tiles * tile_elems_per_thread ==
elems_per_thread)
<< "Tile partition inconsistent with buffer size: " << warp_row_tiles
<< "x" << warp_col_tiles << " tiles of " << kTileM << "x" << kTileN
<< " = " << warp_row_tiles * warp_col_tiles * tile_elems_per_thread
<< " elems/thread, expected " << elems_per_thread;
PrimExpr warp_m = FloorMod(warp_id, m_warp);
PrimExpr warp_n = FloorDiv(warp_id, m_warp);
Array<Stmt> stmts;
for (int i = 0; i < warp_row_tiles; i++) {
for (int j = 0; j < warp_col_tiles; j++) {
int tile_idx = i * warp_col_tiles + j;
PrimExpr row = dst_row_base + warp_m * warp_M + i * kTileM;
PrimExpr col = dst_col_base + warp_n * warp_N + j * kTileN;
PrimExpr ptr = Call(DataType::Handle(), builtin::address_of(),
{BufferLoad(op.dst, {row, col})});
int kMMAK = kTileSize;
stmts.push_back(Evaluate(Call(
DataType::Handle(), cooperative_tensor_store(),
{op.src->data, IntImm(DataType::Int(32), tile_idx), ptr, dst_stride,
IntImm(DataType::Int(32), kTileM), IntImm(DataType::Int(32), kTileN),
Cast(DataType::Bool(), IntImm(DataType::Int(32), 0)),
IntImm(DataType::Int(32), kTileM), IntImm(DataType::Int(32), kTileN),
IntImm(DataType::Int(32), kMMAK), IntImm(DataType::Int(32), 2)})));
}
}
if (stmts.size() == 1) {
return stmts[0];
}
return SeqStmt(stmts);
}
} // namespace
struct Copy {
@@ -118,6 +246,9 @@ struct Copy {
if (CheckSIMDGroupCopy(op)) {
return LowerSIMDGroupCopy(op, lower_args, analyzer);
}
if (CheckCooperativeTensorCopy(op)) {
return LowerCooperativeTensorCopy(op, lower_args, analyzer);
}
return LowerNormalCopy(op, lower_args, analyzer);
}
};
+31
View File
@@ -8,6 +8,7 @@
#include "backend/common/target_utils.h"
#include "metal/op/utils.h"
#include "op/builtin.h"
#include "op/utils.h"
#include "transform/loop_partition.h"
#include "transform/loop_vectorize.h"
@@ -24,6 +25,36 @@ namespace metal {
struct Fill {
static Stmt Lower(const FillNode &op, const LowerArgs &lower_args,
arith::Analyzer *analyzer) {
if (IsCooperativeTensorBuffer(op.dst)) {
int region_elements = 1;
for (auto r : op.region) {
auto imm = r->extent.as<IntImmNode>();
TVM_FFI_ICHECK(imm)
<< "cooperative_tensor fill region must have constant extents";
region_elements *= imm->value;
}
constexpr int kTileM = 16;
constexpr int kTileN = 32;
constexpr int kTileElems = kTileM * kTileN;
TVM_FFI_ICHECK(region_elements % kTileElems == 0)
<< "cooperative_tensor buffer size must be multiple of " << kTileElems
<< ", got " << region_elements;
int num_tiles = region_elements / kTileElems;
PrimExpr fill_value = Cast(op.dst->dtype, op.value);
Array<Stmt> stmts;
for (int i = 0; i < num_tiles; i++) {
stmts.push_back(
Evaluate(Call(DataType::Handle(), cooperative_tensor_fill(),
{op.dst->data, IntImm(DataType::Int(32), i),
fill_value, IntImm(DataType::Int(32), kTileM),
IntImm(DataType::Int(32), kTileN)})));
}
if (stmts.size() == 1) {
return stmts[0];
}
return SeqStmt(stmts);
}
if (IsSIMDGroupBuffer(op.dst)) {
int region_elements = 1;
for (auto r : op.region) {
+72 -5
View File
@@ -22,12 +22,21 @@ namespace metal {
namespace {
constexpr const char *kMetalSIMDGroup = "metal.simdgroup";
constexpr const char *kMetalCooperativeTensor = "metal.cooperative_tensor";
std::pair<int, int> ComputeMetalWarpPartition(const GemmWarpPolicyNode &policy,
int M, int N, int num_warps) {
int M, int N, int num_warps,
String gemm_inst) {
int m_warp = 1, n_warp = 1;
constexpr int kMPerWarp = 8;
constexpr int kNPerWarp = 8;
int kMPerWarp, kNPerWarp;
if (gemm_inst == kMetalCooperativeTensor) {
kMPerWarp = 16;
kNPerWarp = 32;
} else {
// kMetalSIMDGroup: keep existing 8x8 micro tile
kMPerWarp = 8;
kNPerWarp = 8;
}
TVM_FFI_ICHECK(M % kMPerWarp == 0)
<< "M must be divisible by " << kMPerWarp << ", but got " << M;
@@ -71,6 +80,43 @@ std::pair<int, int> ComputeMetalWarpPartition(const GemmWarpPolicyNode &policy,
return {m_warp, n_warp};
}
bool CanUseCooperativeTensor(const GemmWarpPolicyNode &policy, int M, int N,
int K, int num_warps) {
constexpr int kMPerWarp = 16;
constexpr int kNPerWarp = 32;
if (M % kMPerWarp != 0 || N % kNPerWarp != 0 || K % 16 != 0) {
return false;
}
int max_m = M / kMPerWarp;
int max_n = N / kNPerWarp;
if (policy.IsFullRow()) {
int m_warp = num_warps;
if (M % (m_warp * kMPerWarp) != 0) {
m_warp = max_m;
}
return m_warp > 0 && num_warps % m_warp == 0 && num_warps / m_warp <= max_n;
}
if (policy.IsFullCol()) {
int n_warp = num_warps;
if (N % (n_warp * kNPerWarp) != 0) {
n_warp = max_n;
}
return n_warp > 0 && num_warps % n_warp == 0 && num_warps / n_warp <= max_m;
}
if (policy.IsSquare()) {
for (int m = 1; m <= std::min(num_warps, max_m); ++m) {
if (num_warps % m != 0) {
continue;
}
int n = num_warps / m;
if (n <= max_n && M % (m * kMPerWarp) == 0 && N % (n * kNPerWarp) == 0) {
return true;
}
}
}
return false;
}
} // namespace
struct Gemm {
@@ -81,22 +127,43 @@ struct Gemm {
"Metal target "
<< target->str();
}
if (op.c_.scope() == "local.fragment" ||
op.c_.scope() == "metal.simdgroup") {
return kMetalSIMDGroup;
}
int num_warps = block_size / TargetMetalGetWarpSize(target);
if (TargetMetalSupportsMetal4(target) &&
CanUseCooperativeTensor(*op.policy_.operator->(), op.m_, op.n_, op.k_,
num_warps)) {
return kMetalCooperativeTensor;
}
return kMetalSIMDGroup;
}
static std::pair<int, int>
ComputeWarpPartition(const GemmWarpPolicyNode &policy, int M, int N,
int block_size, Target target, String gemm_inst) {
TVM_FFI_ICHECK(gemm_inst == kMetalSIMDGroup)
TVM_FFI_ICHECK(gemm_inst == kMetalSIMDGroup ||
gemm_inst == kMetalCooperativeTensor)
<< "Unsupported Metal GEMM instruction: " << gemm_inst;
int num_warps = block_size / TargetMetalGetWarpSize(target);
return ComputeMetalWarpPartition(policy, M, N, num_warps);
return ComputeMetalWarpPartition(policy, M, N, num_warps, gemm_inst);
}
static bool ReuseExistingSharedLayout(String gemm_inst) {
(void)gemm_inst;
return false;
}
static String InstructionKind(String gemm_inst) {
if (gemm_inst == kMetalSIMDGroup) {
return "metal_simdgroup";
}
if (gemm_inst == kMetalCooperativeTensor) {
return "metal_cooperative_tensor";
}
return "unknown";
}
};
} // namespace metal
+4 -2
View File
@@ -20,8 +20,8 @@ inline bool IsSIMDGroupBuffer(const Buffer &buffer) {
return buffer.defined() && buffer.scope() == "metal.simdgroup";
}
inline bool IsRegisterBuffer(const Buffer &buffer) {
return IsFragmentBuffer(buffer) || IsSIMDGroupBuffer(buffer);
inline bool IsCooperativeTensorBuffer(const Buffer &buffer) {
return buffer.defined() && buffer.scope() == "metal.cooperative_tensor";
}
inline std::pair<int, int> ComputeSquareWarpPartition(int num_warps, int M,
@@ -39,6 +39,8 @@ inline std::pair<int, int> ComputeSquareWarpPartition(int num_warps, int M,
int n = num_warps / m;
if (n > max_n)
continue;
if (M % (m * kMPerWarp) != 0 || N % (n * kNPerWarp) != 0)
continue;
float m_per = static_cast<float>(M) / (m * kMPerWarp);
float n_per = static_cast<float>(N) / (n * kNPerWarp);
+12 -1
View File
@@ -21,13 +21,24 @@ int TargetMetalGetWarpSize(Target target) {
return 32;
}
bool TargetMetalSupportsMetal4(Target target) {
for (const auto &key : target->keys) {
if (key == "metal4") {
return true;
}
}
return false;
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::GlobalDef()
.def("tl.TargetIsMetal",
[](Target target) { return TargetIsMetal(target); })
.def("tl.TargetMetalGetWarpSize",
[](Target target) { return TargetMetalGetWarpSize(target); });
[](Target target) { return TargetMetalGetWarpSize(target); })
.def("tl.TargetMetalSupportsMetal4",
[](Target target) { return TargetMetalSupportsMetal4(target); });
}
} // namespace tl
+1
View File
@@ -13,6 +13,7 @@ namespace tl {
bool TargetIsMetal(Target target);
int TargetMetalGetWarpSize(Target target);
bool TargetMetalSupportsMetal4(Target target);
} // namespace tl
} // namespace tvm
+20
View File
@@ -212,6 +212,26 @@ TIR_DEFINE_TL_BUILTIN(tma_store_scatter4)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
TIR_DEFINE_TL_BUILTIN(cooperative_tensor_fill)
.set_num_inputs(5)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
TIR_DEFINE_TL_BUILTIN(cooperative_tensor_load)
.set_num_inputs(11)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
TIR_DEFINE_TL_BUILTIN(cooperative_tensor_store)
.set_num_inputs(11)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
TIR_DEFINE_TL_BUILTIN(cooperative_tensor_multiply_accumulate)
.set_num_inputs(13)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
TIR_DEFINE_TL_BUILTIN(ptx_fence_barrier_init)
.set_num_inputs(-1)
.set_attr<TCallEffectKind>("TCallEffectKind",
+5
View File
@@ -377,6 +377,11 @@ TVM_DLL const Op &tma_load_gather4();
*/
TVM_DLL const Op &tma_store_scatter4();
TVM_DLL const Op &cooperative_tensor_fill();
TVM_DLL const Op &cooperative_tensor_load();
TVM_DLL const Op &cooperative_tensor_store();
TVM_DLL const Op &cooperative_tensor_multiply_accumulate();
/*!
* \brief tvm intrinsics for barrier initialization fence
*
+2 -1
View File
@@ -74,7 +74,8 @@ void RegisterGemmImpl(GemmImpl impl) {
* expected layout is:
* [Aptr, Bptr, Cptr, trans_A (Bool), trans_B (Bool),
* M (Int), N (Int), K (Int), policy (Int), clear_accum (Bool),
* stride_A (Int), stride_B (Int), offset_A (Int), offset_B (Int),
* stride_A (Int), stride_B (Int), offset_A (PrimExpr),
* offset_B (PrimExpr),
* (optional) kPack (Int), (optional) internal wg_wait (Int),
* (optional) mbar (BufferLoad), cCoord_y (PrimExpr), cCoord_x (PrimExpr)]
*/
+25 -9
View File
@@ -17,6 +17,7 @@
#include <algorithm>
#include <deque>
#include <memory>
#include <optional>
#include <queue>
#include "../layout/layout.h"
@@ -128,14 +129,23 @@ public:
"required for layout inference.";
// Run InferLayout
auto updates = next->InferLayout(LayoutInferArgs{target_,
thread_bounds,
layout_map,
cur_analyzer,
{},
bind_var_to_expr_,
false},
level);
LayoutMap updates;
try {
updates = next->InferLayout(LayoutInferArgs{target_,
thread_bounds,
layout_map,
cur_analyzer,
{},
bind_var_to_expr_,
false},
level);
} catch (const std::bad_optional_access &e) {
LOG(FATAL) << "bad_optional_access while inferring layout for op "
<< cur_infer_id << " (" << next->GetTypeKey() << ") at level "
<< InferLevelToString(level)
<< "\nthread_bounds=" << thread_bounds
<< "\nstmt=" << infer_list_stmt_[cur_infer_id];
}
// Process the returned updates
for (const auto &[buffer, layout] : updates) {
@@ -523,7 +533,13 @@ private:
if (op->op.as<GlobalVarNode>())
return;
auto p = ParseOperator(GetRef<Call>(op));
TileOperator p;
try {
p = ParseOperator(GetRef<Call>(op));
} catch (const std::bad_optional_access &e) {
LOG(FATAL) << "bad_optional_access while parsing tile op call: "
<< GetRef<Call>(op);
}
if (p.defined()) {
for (const auto &arg : op->args) {
if (auto buffer = getBufferFromAccessPtr(arg)) {
+7 -2
View File
@@ -164,8 +164,13 @@ private:
}
void VisitStmt_(const AllocBufferNode *op) final {
auto storage_scope =
runtime::StorageScope::Create(GetPtrStorageScope(op->buffer->data));
auto scope = GetPtrStorageScope(op->buffer->data);
if (scope == "metal.cooperative_tensor") {
StmtVisitor::VisitStmt_(op);
return;
}
auto storage_scope = runtime::StorageScope::Create(scope);
if (storage_scope.rank == runtime::StorageRank::kShared &&
storage_scope.tag == ".dyn") {
ICHECK(!dyn_shmem_size.defined())
+4
View File
@@ -69,6 +69,10 @@ private:
public:
void VisitStmt_(const AllocBufferNode *op) final {
if (op->buffer.scope() == "metal.cooperative_tensor") {
StmtExprVisitor::VisitStmt_(op);
return;
}
if (IsDynamicSharedMemory(op->buffer->data)) {
dyn_shmem_allocs_[op->buffer->data.get()] = op;
} else if (IsStaticSharedMemory(op->buffer->data)) {
@@ -24,6 +24,7 @@
#include "support/check.h"
#include <tvm/ir/cast.h>
#include <tvm/ir/type.h>
#include <tvm/s_tir/analysis.h>
#include <tvm/tirx/analysis.h>
#include <tvm/tirx/stmt.h>
@@ -31,6 +32,8 @@
#include <tvm/tirx/transform.h>
#include <tvm/tirx/var.h>
#include <unordered_map>
#include "../op/utils.h"
#include "tir/transforms/ir_utils.h"
@@ -52,7 +55,74 @@ using namespace tirx::transform;
using tirx::SBlockNode;
using tirx::SBlockRealizeNode;
// Use TVM's tir analysis API for LCA detection.
// TVM's LCA detector parses StorageScope for every touched buffer, while this
// TVM submodule does not know TileLang's Metal cooperative tensor scope yet.
// Cooperative tensor allocations are preserved in their original blocks, so the
// LCA query can use an analysis-only copy where those buffers look local.
class MetalCooperativeTensorLCASanitizer : public StmtExprMutator {
public:
struct Result {
PrimFunc func;
std::unordered_map<const StmtNode *, const StmtNode *> stmt_remap;
};
static Result Sanitize(PrimFunc func) {
MetalCooperativeTensorLCASanitizer sanitizer;
auto fptr = func.CopyOnWrite();
fptr->body = sanitizer.VisitStmt(func->body);
return {func, std::move(sanitizer.stmt_remap_)};
}
private:
Stmt VisitStmt(const Stmt &stmt) final {
if (!stmt.defined()) {
return stmt;
}
const StmtNode *original = stmt.get();
Stmt result = StmtExprMutator::VisitStmt(stmt);
stmt_remap_.emplace(original, original);
stmt_remap_.emplace(result.get(), original);
return result;
}
Buffer VisitBufferDef(const Buffer &buffer, bool alloc_data) final {
Buffer new_buffer = StmtExprMutator::VisitBufferDef(buffer, alloc_data);
return SanitizeBuffer(new_buffer);
}
Buffer VisitBufferUse(const Buffer &buffer) final {
Buffer new_buffer = StmtExprMutator::VisitBufferUse(buffer);
return SanitizeBuffer(new_buffer);
}
Buffer SanitizeBuffer(const Buffer &buffer) {
if (!buffer.defined() || buffer.scope() != "metal.cooperative_tensor") {
return buffer;
}
auto it = sanitized_buffers_.find(buffer.get());
if (it != sanitized_buffers_.end()) {
return it->second;
}
const auto *ptr_type = buffer->data->type_annotation.as<PointerTypeNode>();
ICHECK(ptr_type != nullptr)
<< "Expected cooperative tensor buffer data to have pointer type: "
<< buffer;
Var local_data(buffer->data->name_hint,
PointerType(ptr_type->element_type, "local"),
buffer->data->span);
Buffer sanitized(local_data, buffer->dtype, buffer->shape, buffer->strides,
buffer->elem_offset, buffer->name, buffer->data_alignment,
buffer->offset_factor, buffer->buffer_type,
buffer->axis_separators, buffer->span);
sanitized_buffers_.emplace(buffer.get(), sanitized);
buffer_remap_.Set(buffer, sanitized);
return sanitized;
}
std::unordered_map<const BufferNode *, Buffer> sanitized_buffers_;
std::unordered_map<const StmtNode *, const StmtNode *> stmt_remap_;
};
class CollectManagedAllocations : public StmtExprVisitor {
public:
@@ -194,8 +264,12 @@ class BufferAllocationLocator : public StmtExprMutator {
public:
explicit BufferAllocationLocator(const PrimFunc &func) {
// Use TVM's tir LCA detection implementation
Map<Buffer, Optional<Stmt>> buffer_lca = tirx::DetectBufferAccessLCA(func);
Map<Var, Optional<Stmt>> var_lca = tirx::DetectBufferVarAccessLCA(func);
MetalCooperativeTensorLCASanitizer::Result lca_input =
MetalCooperativeTensorLCASanitizer::Sanitize(func);
Map<Buffer, Optional<Stmt>> buffer_lca =
tirx::DetectBufferAccessLCA(lca_input.func);
Map<Var, Optional<Stmt>> var_lca =
tirx::DetectBufferVarAccessLCA(lca_input.func);
// The buffer_alloc_recorder Array is used to keep the buffer allocation
// order since the buffer_lca Map is unordered.
@@ -221,7 +295,7 @@ public:
// barrier_init annotation remains attached to the block that owns the
// initialization. Moving them into injected opaque blocks causes
// LowerSharedBarrier to see barrier buffers without local annotations.
if (IsBarrierBuffer(buffer)) {
if (ShouldPreserveOriginalBlock(buffer)) {
continue;
}
// Prefer the LCA derived from the underlying data var. If missing, fall
@@ -229,11 +303,11 @@ public:
const StmtNode *stmt = nullptr;
auto vit = var_lca.find(buffer->data);
if (vit != var_lca.end()) {
stmt = (*vit).second.get();
stmt = RemapAnalysisStmt((*vit).second.get(), lca_input.stmt_remap);
} else {
auto bit = buffer_lca.find(buffer);
if (bit != buffer_lca.end()) {
stmt = (*bit).second.get();
stmt = RemapAnalysisStmt((*bit).second.get(), lca_input.stmt_remap);
}
}
stmt = ResolveAllocationSite(buffer->data.get(), stmt);
@@ -254,6 +328,16 @@ public:
}
private:
static const StmtNode *RemapAnalysisStmt(
const StmtNode *stmt,
const std::unordered_map<const StmtNode *, const StmtNode *> &remap) {
if (stmt == nullptr) {
return nullptr;
}
auto it = remap.find(stmt);
return it == remap.end() ? stmt : it->second;
}
// Maintain a stack of Buffers per data var to correctly handle cases
// where multiple Buffer objects share the same underlying data Var.
void PushBinding(const Var &v, const Buffer &buf) {
@@ -351,11 +435,11 @@ private:
Stmt VisitStmt_(const SBlockNode *op) final {
ICHECK(!op->init.defined());
Array<Buffer> alloc_buffers;
Array<Buffer> preserved_barrier_buffers;
Array<Buffer> preserved_original_block_buffers;
for (const Buffer &buf : op->alloc_buffers) {
if (IsBarrierBuffer(buf)) {
if (ShouldPreserveOriginalBlock(buf)) {
alloc_buffers.push_back(buf);
preserved_barrier_buffers.push_back(buf);
preserved_original_block_buffers.push_back(buf);
PushBinding(buf->data, buf);
}
}
@@ -389,7 +473,7 @@ private:
PopBinding(buf->data);
}
}
for (const Buffer &buf : preserved_barrier_buffers) {
for (const Buffer &buf : preserved_original_block_buffers) {
PopBinding(buf->data);
}
@@ -445,9 +529,10 @@ private:
*/
std::unordered_set<const VarNode *> managed_allocations_;
static bool IsBarrierBuffer(const Buffer &buffer) {
static bool ShouldPreserveOriginalBlock(const Buffer &buffer) {
String scope = buffer.scope();
return scope == "shared.barrier" || scope == "shared.cluster_barrier";
return scope == "shared.barrier" || scope == "shared.cluster_barrier" ||
scope == "metal.cooperative_tensor";
}
};
+23
View File
@@ -23,6 +23,7 @@
* Re-write data access to enable memory sharing when possible.
*/
#include "common/attr.h"
#include "metal/op/utils.h"
#include "support/check.h"
#include <tvm/arith/analyzer.h>
#include <tvm/ir/attrs.h>
@@ -56,6 +57,17 @@ using runtime::StorageScope;
using namespace tirx;
using namespace ffi;
namespace {
// Backend-managed allocation scopes whose AllocBuffer nodes must remain in
// place. StorageRewrite should still visit their bodies, but must not plan,
// hoist, merge, or remap the allocation itself.
bool IsStorageRewriteOpaqueAllocBuffer(const Buffer &buffer) {
return metal::IsCooperativeTensorBuffer(buffer);
}
} // namespace
/*!
* \brief Perform data type legalization on the given BufferLoadNode pointer.
* Equal to BufferLoadNode::LegalizeDType, but operates on a pointer.
@@ -130,6 +142,10 @@ public:
};
void VisitStmt_(const AllocBufferNode *op) final {
if (IsStorageRewriteOpaqueAllocBuffer(op->buffer)) {
StmtExprVisitor::VisitStmt_(op);
return;
}
size_t level = scope_.size();
const VarNode *buf = op->buffer->data.get();
@@ -583,6 +599,9 @@ public:
}
Stmt VisitStmt_(const AllocBufferNode *op) final {
if (IsStorageRewriteOpaqueAllocBuffer(op->buffer)) {
return StmtExprMutator::VisitStmt_(op);
}
// AllocBuffer combines allocation and buffer declaration.
// Storage rewrite may merge this allocation with others.
if (auto it = alloc_map_.find(op->buffer->data.get());
@@ -1364,6 +1383,10 @@ public:
}
void VisitStmt_(const AllocBufferNode *op) final {
if (IsStorageRewriteOpaqueAllocBuffer(op->buffer)) {
StmtExprVisitor::VisitStmt_(op);
return;
}
const Array<PrimExpr> &shape = op->buffer->shape;
PrimExpr extent = !shape.empty() ? shape[shape.size() - 1] : PrimExpr(0);
OnArrayDeclaration(op->buffer->data, op->buffer->dtype, extent,
+155
View File
@@ -9,6 +9,15 @@ from tilelang import tvm as tvm
import tilelang.testing
import tilelang.language as T
import torch
import pytest
from tilelang.metal.target import check_metal4_availability
requires_metal4 = pytest.mark.skipif(
not check_metal4_availability(),
reason="direct global-C Metal GEMM requires Metal 4 cooperative tensor support",
)
@tilelang.jit
@@ -66,6 +75,107 @@ def assert_gemm_v2(
)
@tilelang.jit
def matmul_gemm_v2_global_c(
M,
N,
K,
block_M,
block_N,
dtype=T.float16,
accum_dtype=T.float32,
threads=128,
swizzle_panel=0,
swizzle_order="row",
):
@T.prim_func
def gemm_kernel(
A: T.Tensor((M, K), dtype),
B: T.Tensor((K, N), dtype),
C: T.Tensor((M, N), accum_dtype),
):
tiles_n = T.ceildiv(N, block_N)
tiles_m = T.ceildiv(M, block_M)
use_mlx_swizzle = swizzle_panel and swizzle_order == "mlx"
grid_n = tiles_n * swizzle_panel if use_mlx_swizzle else tiles_n
grid_m = T.ceildiv(tiles_m, swizzle_panel) if use_mlx_swizzle else tiles_m
with T.Kernel(grid_n, grid_m, threads=threads) as (bx, by):
logical_bx = bx // swizzle_panel if use_mlx_swizzle else bx
logical_by = by * swizzle_panel + bx % swizzle_panel if use_mlx_swizzle else by
if swizzle_panel:
T.use_swizzle(panel_size=swizzle_panel, order=swizzle_order)
if use_mlx_swizzle:
if logical_by < tiles_m:
T.gemm(
A[logical_by * block_M : (logical_by + 1) * block_M, 0:K],
B[0:K, logical_bx * block_N : (logical_bx + 1) * block_N],
C[
logical_by * block_M : (logical_by + 1) * block_M,
logical_bx * block_N : (logical_bx + 1) * block_N,
],
clear_accum=True,
)
else:
T.gemm(
A[logical_by * block_M : (logical_by + 1) * block_M, 0:K],
B[0:K, logical_bx * block_N : (logical_bx + 1) * block_N],
C[
logical_by * block_M : (logical_by + 1) * block_M,
logical_bx * block_N : (logical_bx + 1) * block_N,
],
clear_accum=True,
)
return gemm_kernel
def assert_gemm_v2_global_c(
M,
N,
K,
block_M,
block_N,
dtype=T.float16,
accum_dtype=T.float32,
threads=128,
swizzle_panel=0,
swizzle_order="row",
atol=1e-2,
rtol=1e-2,
):
jit_kernel = matmul_gemm_v2_global_c(
M,
N,
K,
block_M,
block_N,
dtype=dtype,
accum_dtype=accum_dtype,
threads=threads,
swizzle_panel=swizzle_panel,
swizzle_order=swizzle_order,
)
torch_dtype = dtype.as_torch()
torch_accum_dtype = accum_dtype.as_torch()
a = torch.randn(M, K, dtype=torch_dtype, device="mps")
b = torch.randn(K, N, dtype=torch_dtype, device="mps")
c = torch.zeros(M, N, dtype=torch_accum_dtype, device="mps")
jit_kernel(a, b, c)
if accum_dtype == T.float16:
ref = (a.to(torch.float32) @ b.to(torch.float32)).to(torch_accum_dtype)
else:
ref = a.to(torch_accum_dtype) @ b.to(torch_accum_dtype)
assert torch.allclose(ref, c, atol=atol, rtol=rtol), (
f"Result mismatch for direct global C, M={M}, N={N}, K={K}, "
f"block=({block_M},{block_N}), dtype={dtype}\n"
f"max diff: {(ref - c).abs().max().item()}"
)
@tilelang.testing.requires_metal
def test_gemm_v2_16x16x16():
assert_gemm_v2(128, 128, 128, 16, 16, 16)
@@ -81,6 +191,51 @@ def test_gemm_v2_large():
assert_gemm_v2(128, 128, 128, 32, 32, 32)
@tilelang.testing.requires_metal
def test_gemm_v2_cooperative_tensor_non_square():
assert_gemm_v2(128, 128, 128, 32, 64, 32)
@tilelang.testing.requires_metal
@requires_metal4
def test_gemm_v2_cooperative_tensor_global_c():
assert_gemm_v2_global_c(256, 256, 256, 64, 128)
@tilelang.testing.requires_metal
@requires_metal4
def test_gemm_v2_cooperative_tensor_global_c_fp16_mlx_swizzle():
assert_gemm_v2_global_c(
256,
256,
256,
64,
128,
accum_dtype=T.float16,
swizzle_panel=4,
swizzle_order="mlx",
atol=1e-1,
rtol=1e-2,
)
@tilelang.testing.requires_metal
@requires_metal4
def test_gemm_v2_cooperative_tensor_global_c_fp16_mlx_swizzle_multirow():
assert_gemm_v2_global_c(
512,
256,
256,
64,
128,
accum_dtype=T.float16,
swizzle_panel=4,
swizzle_order="mlx",
atol=1e-1,
rtol=1e-2,
)
@tilelang.testing.requires_metal
def test_gemm_v2_1024():
assert_gemm_v2(1024, 1024, 1024, 16, 16, 16, atol=1.0)
@@ -13,6 +13,9 @@ import tilelang.testing
import tilelang.language as T
from tilelang.metal.intrinsics.metal_macro_generator import MPSIntrinEmitter
METAL_TARGET = tvm.target.Target({"kind": "metal", "keys": ["metal", "gpu"]})
METAL4_TARGET = tvm.target.Target({"kind": "metal", "keys": ["metal", "gpu", "metal4"]})
def matmul_gemm_v2(M, N, K, block_M, block_N, block_K, dtype=T.float16, accum_dtype=T.float32):
@T.prim_func
@@ -39,6 +42,85 @@ def matmul_gemm_v2(M, N, K, block_M, block_N, block_K, dtype=T.float16, accum_dt
return main
def matmul_gemm_v2_shared_c(M, N, K, block_M, block_N, block_K, dtype=T.float16, accum_dtype=T.float32):
@T.prim_func
def main(
A: T.Tensor((M, K), dtype),
B: T.Tensor((K, N), dtype),
C: T.Tensor((M, N), accum_dtype),
):
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
A_shared = T.alloc_shared((block_M, block_K), dtype, scope="shared")
B_shared = T.alloc_shared((block_K, block_N), dtype, scope="shared")
C_shared = T.alloc_shared((block_M, block_N), accum_dtype, scope="shared")
T.clear(C_shared)
for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=0):
T.copy(A[by * block_M, ko * block_K], A_shared, coalesced_width=2)
T.copy(B[ko * block_K, bx * block_N], B_shared, coalesced_width=2)
T.gemm(A_shared, B_shared, C_shared)
T.copy(C_shared, C[by * block_M, bx * block_N], coalesced_width=2)
return main
def matmul_gemm_v2_global_c(
M,
N,
K,
block_M,
block_N,
dtype=T.float16,
accum_dtype=T.float32,
threads=128,
swizzle_panel=0,
swizzle_order="row",
):
@T.prim_func
def main(
A: T.Tensor((M, K), dtype),
B: T.Tensor((K, N), dtype),
C: T.Tensor((M, N), accum_dtype),
):
tiles_n = T.ceildiv(N, block_N)
tiles_m = T.ceildiv(M, block_M)
use_mlx_swizzle = swizzle_panel and swizzle_order == "mlx"
grid_n = tiles_n * swizzle_panel if use_mlx_swizzle else tiles_n
grid_m = T.ceildiv(tiles_m, swizzle_panel) if use_mlx_swizzle else tiles_m
with T.Kernel(grid_n, grid_m, threads=threads) as (bx, by):
logical_bx = bx // swizzle_panel if use_mlx_swizzle else bx
logical_by = by * swizzle_panel + bx % swizzle_panel if use_mlx_swizzle else by
if swizzle_panel:
T.use_swizzle(panel_size=swizzle_panel, order=swizzle_order)
if use_mlx_swizzle:
if logical_by < tiles_m:
T.gemm(
A[logical_by * block_M : (logical_by + 1) * block_M, 0:K],
B[0:K, logical_bx * block_N : (logical_bx + 1) * block_N],
C[
logical_by * block_M : (logical_by + 1) * block_M,
logical_bx * block_N : (logical_bx + 1) * block_N,
],
clear_accum=True,
)
else:
T.gemm(
A[logical_by * block_M : (logical_by + 1) * block_M, 0:K],
B[0:K, logical_bx * block_N : (logical_bx + 1) * block_N],
C[
logical_by * block_M : (logical_by + 1) * block_M,
logical_bx * block_N : (logical_bx + 1) * block_N,
],
clear_accum=True,
)
return main
def assert_metal_gemm_v2_codegen(
M,
N,
@@ -55,13 +137,89 @@ def assert_metal_gemm_v2_codegen(
src_code = artifact.kernel_source
assert src_code is not None
assert "kernel void" in src_code
assert "main_kernel" in src_code
# Verify simdgroup matrix operations are present
assert "simdgroup_multiply_accumulate" in src_code
assert "simdgroup_load" in src_code
assert "simdgroup_store" in src_code
def assert_metal_gemm_v2_cooperative_tensor_codegen(
M,
N,
K,
block_M,
block_N,
block_K,
dtype=T.float16,
accum_dtype=T.float32,
):
func = matmul_gemm_v2_shared_c(M, N, K, block_M, block_N, block_K, dtype=dtype, accum_dtype=accum_dtype)
with tvm.transform.PassContext(), METAL4_TARGET:
artifact = tilelang.lower(func, target=METAL4_TARGET)
src_code = artifact.kernel_source
assert src_code is not None
assert "main_kernel" in src_code
assert "mpp::tensor_ops::matmul2d" in src_code
assert "cooperative_tensor" in src_code
def assert_metal_gemm_v2_global_cooperative_tensor_codegen(
M,
N,
K,
block_M,
block_N,
dtype=T.float16,
accum_dtype=T.float32,
threads=128,
swizzle_panel=0,
swizzle_order="row",
):
func = matmul_gemm_v2_global_c(
M,
N,
K,
block_M,
block_N,
dtype=dtype,
accum_dtype=accum_dtype,
threads=threads,
swizzle_panel=swizzle_panel,
swizzle_order=swizzle_order,
)
with tvm.transform.PassContext(), METAL4_TARGET:
artifact = tilelang.lower(func, target=METAL4_TARGET)
src_code = artifact.kernel_source
assert src_code is not None
assert "main_kernel" in src_code
assert "mpp::tensor_ops::matmul2d" in src_code
assert "const device half* __restrict A" in src_code
assert "const device half* __restrict B" in src_code
assert "const device half* __src" in src_code
assert "[[simdgroup_index_in_threadgroup]]" in src_code
assert "__metal_get_thread_index_in_simdgroup" in src_code
assert "max_total_threads_per_threadgroup(128)" in src_code
assert "threadgroup half" not in src_code
assert "thread float C_ct" not in src_code
assert "blockIdx.x" in src_code
if swizzle_order != "mlx":
assert "blockIdx.y" in src_code
if accum_dtype == T.float16:
assert "device half* __restrict C" in src_code
assert "half4(*(thread float4*)" in src_code
if swizzle_order == "mlx":
assert "__physical_blockIdx" in src_code
assert "__physical_blockIdx.x >> 2" in src_code
assert "__physical_blockIdx.x & 3u" in src_code
assert "blockIdx.x) >> 2" not in src_code
assert "blockIdx.x) & 3" not in src_code
if (M + block_M - 1) // block_M > swizzle_panel:
assert f"* {block_M * K * swizzle_panel}" not in src_code
def test_metal_gemm_v2_float16():
assert_metal_gemm_v2_codegen(64, 64, 64, 16, 16, 16, dtype=T.float16)
@@ -74,6 +232,40 @@ def test_metal_gemm_v2_larger():
assert_metal_gemm_v2_codegen(128, 128, 128, 32, 32, 32, dtype=T.float16)
def test_metal_gemm_v2_cooperative_tensor_codegen():
assert_metal_gemm_v2_cooperative_tensor_codegen(128, 128, 128, 32, 64, 32, dtype=T.float16)
def test_metal_gemm_v2_without_metal4_uses_simdgroup():
func = matmul_gemm_v2_shared_c(128, 128, 128, 32, 64, 32, dtype=T.float16)
with tvm.transform.PassContext(), METAL_TARGET:
artifact = tilelang.lower(func, target=METAL_TARGET)
src_code = artifact.kernel_source
assert src_code is not None
assert "mpp::tensor_ops::matmul2d" not in src_code
assert "MetalPerformancePrimitives" not in src_code
assert "simdgroup_multiply_accumulate" in src_code
def test_metal_gemm_v2_global_cooperative_tensor_codegen():
assert_metal_gemm_v2_global_cooperative_tensor_codegen(128, 256, 128, 64, 128, dtype=T.float16)
def test_metal_gemm_v2_global_cooperative_tensor_mlx_swizzle_codegen():
assert_metal_gemm_v2_global_cooperative_tensor_codegen(
512,
256,
128,
64,
128,
dtype=T.float16,
accum_dtype=T.float16,
swizzle_panel=4,
swizzle_order="mlx",
)
def test_metal_gemm_v2_small_blocks():
"""Test with blocks where warp_rows > 1 and warp_cols > 1, which previously
produced incorrect results due to swizzle padding changing the stride.
@@ -63,9 +63,10 @@ def assert_simdgroup_store_codegen(M, N, K, block_M, block_N, block_K, dtype=T.f
src = artifact.kernel_source
assert src is not None
assert "kernel void" in src
assert "[[kernel" in src or "kernel void" in src
assert "simdgroup_multiply_accumulate" in src
assert "make_filled_simdgroup_matrix" in src
assert "MetalPerformancePrimitives" not in src
assert "simdgroup_float8x8" in src or "simdgroup_half8x8" in src, "Expected simdgroup_float8x8 or simdgroup_half8x8 for C accumulator"
@@ -1,5 +1,6 @@
from __future__ import annotations
from tvm import DataType
from tvm.tirx import IndexMap
import tilelang.language as T
@@ -39,6 +40,22 @@ def ldmatrix_32x16_to_shared_16x32_layout_b(thread_id, local_id):
return row, col
def metal_ct_store_32x16_to_16x32_layout(thread_id, local_id):
lane = thread_id % 32
qid = lane >> 2
base_row = (qid & 4) | ((lane >> 1) & 3)
base_col = ((qid & 2) | (lane & 1)) * 4
frag = local_id // 8
frag_local = local_id % 8
row = base_row + (frag_local // 4) * 8
col = base_col + frag * 16 + frag_local % 4
return row, col
def metal_ct_store_index_map():
return IndexMap.from_func(metal_ct_store_32x16_to_16x32_layout, index_dtype=T.int32)
def mma_store_32x8_to_shared_16x16_layout(thread_id, local_id):
row = 8 * (local_id % 4 // 2) + (thread_id // 4)
col = 8 * (local_id // 4) + (thread_id % 4) * 2 + (local_id % 2)
+1 -1
View File
@@ -49,7 +49,7 @@ def register_metal_postproc(func: Callable[[str, Target], str], override: bool =
and returns the processed code (str).
override: Whether to override existing registered function. Defaults to True.
"""
tvm_ffi.register_global_func("tvm_callback_metal_compile", f=func, override=override)
tvm_ffi.register_global_func("tilelang_callback_metal_postproc", f=func, override=override)
def register_cuda_postproc_callback(func: Callable | bool = None, override: bool = True):
+8 -1
View File
@@ -20,7 +20,14 @@ __all__ = [
def use_swizzle(panel_size: int, order: str = "row", enable: bool = True):
"""Annotate a kernel to use a specific threadblock swizzle pattern."""
device_func = "rasterization2DRow" if order == "row" else "rasterization2DColumn"
if order == "row":
device_func = "rasterization2DRow"
elif order == "column":
device_func = "rasterization2DColumn"
elif order == "mlx":
device_func = "rasterization2DMLX"
else:
raise ValueError(f"Unsupported swizzle order: {order}")
if not enable:
return None
return attr(None, "threadblock_swizzle_pattern", tvm_tuple(device_func, panel_size))
+114
View File
@@ -1260,6 +1260,120 @@ def increase_descriptor_offset(descriptor: PrimExpr, offset: PrimExpr) -> PrimEx
return evaluate(tirx.call_intrin("handle", tirx.op.Op.get("tl.increase_descriptor_offset"), descriptor, offset))
def cooperative_tensor_fill(data, idx, value, rows: int, cols: int):
return evaluate(
tirx.call_intrin(
"handle",
tirx.op.Op.get("tl.cooperative_tensor_fill"),
data,
idx,
value,
rows,
cols,
)
)
def cooperative_tensor_load(
data,
idx,
ptr,
stride,
rows: int,
cols: int,
transposed,
tile_m: int,
tile_n: int,
tile_k: int,
operand_role: int,
):
return evaluate(
tirx.call_intrin(
"handle",
tirx.op.Op.get("tl.cooperative_tensor_load"),
data,
idx,
ptr,
stride,
rows,
cols,
transposed,
tile_m,
tile_n,
tile_k,
operand_role,
)
)
def cooperative_tensor_store(
data,
idx,
ptr,
stride,
rows: int,
cols: int,
transposed,
tile_m: int,
tile_n: int,
tile_k: int,
operand_role: int,
):
return evaluate(
tirx.call_intrin(
"handle",
tirx.op.Op.get("tl.cooperative_tensor_store"),
data,
idx,
ptr,
stride,
rows,
cols,
transposed,
tile_m,
tile_n,
tile_k,
operand_role,
)
)
def cooperative_tensor_multiply_accumulate(
c_data,
c_idx,
a_data,
a_idx,
b_data,
b_idx,
d_data,
d_idx,
m: int,
n: int,
k: int,
trans_a,
trans_b,
):
return evaluate(
tirx.call_intrin(
"handle",
tirx.op.Op.get("tl.cooperative_tensor_multiply_accumulate"),
c_data,
c_idx,
a_data,
a_idx,
b_data,
b_idx,
d_data,
d_idx,
m,
n,
k,
trans_a,
trans_b,
)
)
def loop_break():
"""Break out of the innermost loop."""
return tirx.call_intrin("handle", tirx.op.Op.get("tl.loop_break"))
-2
View File
@@ -101,8 +101,6 @@ def _gemm_impl(
A_offset = retrieve_offset(A_region)
B_offset = retrieve_offset(B_region)
assert A_offset[-2] == 0, "The offset of the first dimension of A must be 0"
assert B_offset[-2] == 0, "The offset of the first dimension of B must be 0"
offset_a = A_offset[-1]
offset_b = B_offset[-1]
+1
View File
@@ -2,6 +2,7 @@ from . import codegen # noqa: F401
from . import intrinsics # noqa: F401
from . import op # noqa: F401
from . import pipeline # noqa: F401
from .utils import is_metal_cooperative_tensor, is_metal_simdgroup # noqa: F401
from . import target # noqa: F401
from . import execution_backend # noqa: F401
from . import transform # noqa: F401
+2 -1
View File
@@ -8,6 +8,7 @@ from tilelang.backend.host_codegen import HostCodegenHook, register_host_codegen
_build_metal = global_func_device_codegen("target.build.tilelang_metal")
_build_metal_without_compile = global_func_device_codegen("target.build.tilelang_metal_without_compile")
def _mark_host_metal_context(mod: IRModule, target_host: Target, target: Target) -> IRModule:
@@ -21,7 +22,7 @@ register_device_codegen(
DeviceCodegen(
"metal",
build=_build_metal,
build_without_compile=_build_metal,
build_without_compile=_build_metal_without_compile,
),
override=True,
)
@@ -5,9 +5,13 @@ from tvm import arith
from tvm import tirx as tir
from tvm.tirx import Buffer, BufferRegion
OPERAND_LEFT = 0
OPERAND_RIGHT = 1
OPERAND_DEST = 2
class MPSIntrinEmitter:
"""Metal simdgroup MMA intrinsic emitter for GEMM operations."""
"""Metal simdgroup/cooperative tensor intrinsic emitter for GEMM operations."""
WARP_SIZE = 32
@@ -23,9 +27,13 @@ class MPSIntrinEmitter:
warp_row_tiles: int = 8,
warp_col_tiles: int = 8,
chunk: int = 32,
thread_var: tir.Var | None = None,
thread_var: tir.PrimExpr | None = None,
a_stride_override: int | None = None,
b_stride_override: int | None = None,
inner_k_steps: int = 1,
use_cooperative_tensor: bool = True,
):
"""Initialize the Metal simdgroup MMA emitter."""
"""Initialize the Metal GEMM intrinsic emitter."""
self.a_dtype = a_dtype
self.b_dtype = b_dtype
self.accum_dtype = accum_dtype
@@ -37,13 +45,20 @@ class MPSIntrinEmitter:
self.warp_col_tiles = warp_col_tiles
self.chunk = chunk
self.thread_var = thread_var
self.a_stride_override = a_stride_override
self.b_stride_override = b_stride_override
self.inner_k_steps = inner_k_steps
self.use_cooperative_tensor = use_cooperative_tensor
# Metal simdgroup matrix size (always 8x8)
self.micro_size_x = 8
self.micro_size_y = 8
self.micro_size_k = 8
if use_cooperative_tensor:
self.micro_size_x = 16
self.micro_size_y = 32
self.micro_size_k = 16
else:
self.micro_size_x = 8
self.micro_size_y = 8
self.micro_size_k = 8
# Number of 8x8 tiles per warp
self.warp_rows = warp_row_tiles // self.micro_size_x
self.warp_cols = warp_col_tiles // self.micro_size_y
@@ -53,18 +68,13 @@ class MPSIntrinEmitter:
current_frame = T.KernelLaunchFrame.Current()
assert current_frame is not None, "Must be called in a T.Kernel Frame"
return current_frame.get_thread_binding()
else:
return self.thread_var
return self.thread_var
def _get_warp_indices(self):
"""Compute (warp_m, warp_n) from the current thread binding."""
thread_binding = self.get_thread_binding()
WARP_SIZE = self.WARP_SIZE
block_row_warps = self.block_row_warps
block_col_warps = self.block_col_warps
warp_m = (thread_binding // WARP_SIZE) % block_row_warps
warp_n = (thread_binding // (WARP_SIZE * block_row_warps)) % block_col_warps
warp_m = (thread_binding // self.WARP_SIZE) % self.block_row_warps
warp_n = (thread_binding // (self.WARP_SIZE * self.block_row_warps)) % self.block_col_warps
return warp_m, warp_n
@staticmethod
@@ -74,166 +84,268 @@ class MPSIntrinEmitter:
For 2D buffers, extra_indices is an empty tuple.
For 3D+ buffers (e.g. after pipeline expansion), extra_indices contains
the leading dimension offsets so callers can index correctly.
Metal simdgroup matrix operations accept a row stride but require the
Metal matrix operations accept a row stride but require the
innermost dimension to be contiguous.
"""
if isinstance(buf, BufferRegion):
buffer = buf.buffer
extra = tuple(r.min for r in buf.region[:-2])
off_row = buf.region[-2].min
off_col = buf.region[-1].min
extra = tuple(r.min for r in buf.region[:-2])
else:
buffer = buf
extra = ()
off_row = 0
off_col = 0
extra = ()
if buffer.strides:
inner_stride = buffer.strides[-1]
if not arith.Analyzer().can_prove_equal(inner_stride, 1):
raise ValueError(
f"Metal simdgroup matrix operations require a contiguous innermost dimension (stride 1), but got stride {inner_stride}"
f"Metal matrix operations require a contiguous innermost dimension (stride 1), but got stride {inner_stride}"
)
stride = buffer.strides[-2]
else:
stride = buffer.shape[-1]
return buffer, extra, off_row, off_col, stride
def ldmatrix_a(self, A_local_buf, A_shared_buf: Buffer | BufferRegion, ki):
"""Load matrix A tiles from shared memory into simdgroup local buffers."""
def ldmatrix_a(self, A_local_buf, A_shared_buf: Buffer | BufferRegion, ki, k_inner: int = 0):
"""Load matrix A tiles from memory into simdgroup/cooperative tensor buffers."""
warp_rows = self.warp_rows
warp_row_tiles = self.warp_row_tiles
micro_size_x = self.micro_size_x
micro_size_y = self.micro_size_y
micro_size_k = self.micro_size_k
a_transposed = self.a_transposed
use_cooperative_tensor = self.use_cooperative_tensor
warp_m, _ = self._get_warp_indices()
buffer, extra, offset_m, offset_k, stride = self._parse_buffer_nd(A_shared_buf)
if self.a_stride_override is not None:
stride = self.a_stride_override
@T.macro
def _warp_ldmatrix_a(A_local_buf, buffer, offset_m, offset_k, stride, warp_m, ki):
"""Load A matrix tiles via simdgroup_load per warp row."""
for i in T.serial(warp_rows):
if a_transposed:
row_idx = offset_k + ki * micro_size_k
col_idx = offset_m + warp_m * (self.warp_row_tiles) + i * micro_size_x
col_idx = offset_m + warp_m * warp_row_tiles + i * micro_size_x
else:
row_idx = offset_m + warp_m * (self.warp_row_tiles) + i * micro_size_x
row_idx = offset_m + warp_m * warp_row_tiles + i * micro_size_x
col_idx = offset_k + ki * micro_size_k
indices = extra + (row_idx, col_idx)
ptr = T.access_ptr(buffer[indices], "r")
T.simdgroup_load(
A_local_buf.data,
i,
ptr,
stride,
micro_size_x,
micro_size_k,
T.bool(a_transposed),
)
ptr = T.access_ptr(buffer[extra + (row_idx, col_idx)], "r")
if use_cooperative_tensor:
T.cooperative_tensor_load(
A_local_buf.data,
k_inner * warp_rows + i,
ptr,
stride,
micro_size_x,
micro_size_k,
T.bool(a_transposed),
micro_size_x,
micro_size_y,
micro_size_k,
OPERAND_LEFT,
)
else:
T.simdgroup_load(
A_local_buf.data,
i,
ptr,
stride,
micro_size_x,
micro_size_k,
T.bool(a_transposed),
)
return _warp_ldmatrix_a(A_local_buf, buffer, offset_m, offset_k, stride, warp_m, ki)
def ldmatrix_b(self, B_local_buf, B_shared_buf: Buffer | BufferRegion, ki):
"""Load matrix B tiles from shared memory into simdgroup local buffers."""
def ldmatrix_b(self, B_local_buf, B_shared_buf: Buffer | BufferRegion, ki, k_inner: int = 0):
"""Load matrix B tiles from memory into simdgroup/cooperative tensor buffers."""
warp_cols = self.warp_cols
warp_col_tiles = self.warp_col_tiles
micro_size_x = self.micro_size_x
micro_size_y = self.micro_size_y
micro_size_k = self.micro_size_k
b_transposed = self.b_transposed
use_cooperative_tensor = self.use_cooperative_tensor
_, warp_n = self._get_warp_indices()
buffer, extra, offset_k, offset_n, stride = self._parse_buffer_nd(B_shared_buf)
if self.b_stride_override is not None:
stride = self.b_stride_override
@T.macro
def _warp_ldmatrix_b(B_local_buf, buffer, offset_k, offset_n, stride, warp_n, ki):
"""Load B matrix tiles via simdgroup_load per warp column."""
for j in T.serial(warp_cols):
if b_transposed:
row_idx = offset_n + warp_n * (self.warp_col_tiles) + j * micro_size_y
row_idx = offset_n + warp_n * warp_col_tiles + j * micro_size_y
col_idx = offset_k + ki * micro_size_k
else:
row_idx = offset_k + ki * micro_size_k
col_idx = offset_n + warp_n * (self.warp_col_tiles) + j * micro_size_y
indices = extra + (row_idx, col_idx)
ptr = T.access_ptr(buffer[indices], "r")
T.simdgroup_load(
B_local_buf.data,
j,
ptr,
stride,
micro_size_k,
micro_size_y,
T.bool(b_transposed),
)
col_idx = offset_n + warp_n * warp_col_tiles + j * micro_size_y
ptr = T.access_ptr(buffer[extra + (row_idx, col_idx)], "r")
if use_cooperative_tensor:
T.cooperative_tensor_load(
B_local_buf.data,
k_inner * warp_cols + j,
ptr,
stride,
micro_size_k,
micro_size_y,
T.bool(b_transposed),
micro_size_x,
micro_size_y,
micro_size_k,
OPERAND_RIGHT,
)
else:
T.simdgroup_load(
B_local_buf.data,
j,
ptr,
stride,
micro_size_k,
micro_size_y,
T.bool(b_transposed),
)
return _warp_ldmatrix_b(B_local_buf, buffer, offset_k, offset_n, stride, warp_n, ki)
def mma(self, A_local_buf, B_local_buf, C_local_buf):
"""Perform simdgroup matrix multiply-accumulate: C += A * B."""
def mma(self, A_local_buf, B_local_buf, C_local_buf, k_inner: int = 0):
"""Perform matrix multiply-accumulate: C += A * B."""
warp_rows = self.warp_rows
warp_cols = self.warp_cols
micro_size_x = self.micro_size_x
micro_size_y = self.micro_size_y
micro_size_k = self.micro_size_k
a_transposed = self.a_transposed
b_transposed = self.b_transposed
use_cooperative_tensor = self.use_cooperative_tensor
@T.macro
def _warp_mma(A_local_buf, B_local_buf, C_local_buf):
"""Execute simdgroup matrix multiply-accumulate across all warp tiles."""
for i, j in T.grid(warp_rows, warp_cols):
T.simdgroup_multiply_accumulate(
C_local_buf.data,
i * warp_cols + j,
A_local_buf.data,
i,
B_local_buf.data,
j,
C_local_buf.data,
i * warp_cols + j,
)
index_c = i * warp_cols + j
if use_cooperative_tensor:
T.cooperative_tensor_multiply_accumulate(
C_local_buf.data,
index_c,
A_local_buf.data,
k_inner * warp_rows + i,
B_local_buf.data,
k_inner * warp_cols + j,
C_local_buf.data,
index_c,
micro_size_x,
micro_size_y,
micro_size_k,
T.bool(a_transposed),
T.bool(b_transposed),
)
else:
T.simdgroup_multiply_accumulate(
C_local_buf.data,
index_c,
A_local_buf.data,
i,
B_local_buf.data,
j,
C_local_buf.data,
index_c,
)
return _warp_mma(A_local_buf, B_local_buf, C_local_buf)
def simdgroup_copy(self, C_simd_buf, C_dst, is_store=True):
"""Copy between simdgroup local buffers and shared/global memory."""
"""Copy between register-backed Metal matrix buffers and memory."""
warp_rows = self.warp_rows
warp_cols = self.warp_cols
warp_row_tiles = self.warp_row_tiles
warp_col_tiles = self.warp_col_tiles
micro_size_x = self.micro_size_x
micro_size_y = self.micro_size_y
micro_size_k = self.micro_size_k
use_cooperative_tensor = self.use_cooperative_tensor
warp_m, warp_n = self._get_warp_indices()
buffer, extra, offset_m, offset_n, stride = self._parse_buffer_nd(C_dst)
ct_op = T.cooperative_tensor_store if is_store else T.cooperative_tensor_load
simd_op = T.simdgroup_store if is_store else T.simdgroup_load
access_mode = "w" if is_store else "r"
@T.macro
def _simdgroup_copy(C_simd_buf, buffer, offset_m, offset_n, stride, warp_m, warp_n):
"""Copy tiles between simdgroup local and shared memory."""
for i, j in T.grid(warp_rows, warp_cols):
row = offset_m + warp_m * self.warp_row_tiles + i * micro_size_x
col = offset_n + warp_n * self.warp_col_tiles + j * micro_size_y
row = offset_m + warp_m * warp_row_tiles + i * micro_size_x
col = offset_n + warp_n * warp_col_tiles + j * micro_size_y
index_c = i * warp_cols + j
indices = extra + (row, col)
simd_op(
C_simd_buf.data,
index_c,
T.access_ptr(buffer[indices], access_mode),
stride,
micro_size_x,
micro_size_y,
T.bool(False),
)
ptr = T.access_ptr(buffer[extra + (row, col)], access_mode)
if use_cooperative_tensor:
ct_op(
C_simd_buf.data,
index_c,
ptr,
stride,
micro_size_x,
micro_size_y,
T.bool(False),
micro_size_x,
micro_size_y,
micro_size_k,
OPERAND_DEST,
)
else:
simd_op(
C_simd_buf.data,
index_c,
ptr,
stride,
micro_size_x,
micro_size_y,
T.bool(False),
)
return _simdgroup_copy(C_simd_buf, buffer, offset_m, offset_n, stride, warp_m, warp_n)
def make_cooperative_tensor_store_layout(self, local_buf):
from tilelang.utils.language import is_fragment
assert is_fragment(local_buf), f"{local_buf} must be a fragment"
shape = local_buf.shape
def inverse_metal_ct_index(i, j):
row = i % self.micro_size_x
col = j % self.micro_size_y
row_low = row % 8
col_group = (col % 16) // 4
lane_id = (col_group % 2) + (row_low % 4) * 2 + (col_group // 2) * 8 + (row_low // 4) * 16
local_id = (col // 16) * 8 + (row // 8) * 4 + (col % 4)
return lane_id, local_id
def forward_thread(i: int, j: int) -> int:
warp_m = (i // self.micro_size_x) // self.warp_rows
warp_n = (j // self.micro_size_y) // self.warp_cols
mma_i = i % self.micro_size_x
mma_j = j % self.micro_size_y
lane_id, _ = inverse_metal_ct_index(mma_i, mma_j)
return warp_m * (self.block_col_warps * self.WARP_SIZE) + warp_n * self.WARP_SIZE + lane_id
def forward_index(i: int, j: int) -> int:
warp_i = (i // self.micro_size_x) % self.warp_rows
warp_j = (j // self.micro_size_y) % self.warp_cols
mma_i = i % self.micro_size_x
mma_j = j % self.micro_size_y
_, local_id = inverse_metal_ct_index(mma_i, mma_j)
return warp_i * (self.warp_cols * 16) + warp_j * 16 + local_id
return T.Fragment(shape, forward_thread_fn=forward_thread, forward_index_fn=forward_index)
def simd_store(self, C_simd_buf, C_dst):
"""Store simdgroup local buffer to shared/global memory."""
"""Store simdgroup/cooperative tensor local buffer to memory."""
return self.simdgroup_copy(C_simd_buf, C_dst, is_store=True)
def simd_load(self, C_simd_buf, C_src):
"""Load shared/global memory into simdgroup local buffer."""
"""Load memory into simdgroup/cooperative tensor local buffer."""
return self.simdgroup_copy(C_simd_buf, C_src, is_store=False)
+18 -1
View File
@@ -4,11 +4,28 @@ from __future__ import annotations
from tilelang.language.common import * # noqa: F401,F403
from tilelang.language.common import __all__ as _COMMON_ALL
from tilelang.language.builtin import ( # noqa: F401
cooperative_tensor_fill,
cooperative_tensor_load,
cooperative_tensor_multiply_accumulate,
cooperative_tensor_store,
)
from .tir import * # noqa: F401,F403
from .tir import __all__ as _TIR_ALL
__tilelang_dialect__ = "metal"
__all__ = tuple(dict.fromkeys((*_COMMON_ALL, *_TIR_ALL)))
__all__ = tuple(
dict.fromkeys(
(
*_COMMON_ALL,
*_TIR_ALL,
"cooperative_tensor_fill",
"cooperative_tensor_load",
"cooperative_tensor_multiply_accumulate",
"cooperative_tensor_store",
)
)
)
del _COMMON_ALL, _TIR_ALL
+9 -2
View File
@@ -2,11 +2,18 @@ from __future__ import annotations
from tilelang.tileop.gemm.registry import register_gemm_impl
from tilelang.metal.target import target_is_metal
from .gemm_metal import GEMM_INST_METAL, GemmMetal
from .gemm_metal import (
GEMM_INST_METAL,
GEMM_INST_METAL_COOPERATIVE_TENSOR,
GemmMetal,
GemmMetalSimdGroup,
)
def _match_metal(target) -> bool:
return target_is_metal(target)
register_gemm_impl("metal.simdgroup", GEMM_INST_METAL, _match_metal, GemmMetal)
register_gemm_impl(GEMM_INST_METAL, GEMM_INST_METAL, _match_metal, GemmMetalSimdGroup)
register_gemm_impl(GEMM_INST_METAL_COOPERATIVE_TENSOR, GEMM_INST_METAL_COOPERATIVE_TENSOR, _match_metal, GemmMetal)
+251 -45
View File
@@ -1,19 +1,41 @@
from __future__ import annotations
from tilelang.tileop.gemm.gemm_base import GemmBase
from tilelang.utils.language import is_shared, is_full_region, is_fragment, is_metal_simdgroup
from tilelang import tvm as tvm
from tvm.target import Target
from tvm.ir import Range
from tvm import tirx as tir
from tilelang.layout import Layout
from tilelang.metal import language as T
from tilelang.metal.utils import (
is_metal_cooperative_tensor,
is_metal_simdgroup,
)
from tilelang.tileop.gemm.gemm_base import GemmBase
from tilelang.transform.simplify import _Simplify
from tilelang.utils.language import (
is_fragment,
is_full_region,
is_global,
is_shared,
)
from tvm import tirx as tir
from tvm.ir import Range
from tvm.target import Target
GEMM_INST_METAL = "metal.simdgroup"
GEMM_INST_METAL_COOPERATIVE_TENSOR = "metal.cooperative_tensor"
class GemmMetal(GemmBase):
def _make_padded_layout(buffer):
shape = buffer.shape
stride = int(shape[-2])
continuous = int(shape[-1])
element_bits = int(tvm.DataType(buffer.dtype).bits)
padded = continuous
if (element_bits * continuous) % 256 == 0:
padded += 128 // element_bits
return Layout([stride, continuous], lambda i, j: i * padded + j)
class GemmMetalSimdGroup(GemmBase):
def is_gemm_ss(self) -> bool:
return is_shared(self.A) and is_shared(self.B)
@@ -47,6 +69,7 @@ class GemmMetal(GemmBase):
warp_col_tiles=warp_col_tiles,
chunk=self.chunk,
thread_var=thread_index,
use_cooperative_tensor=False,
)
a_dtype = self.a_dtype
@@ -61,9 +84,7 @@ class GemmMetal(GemmBase):
A_region = self.ARegion
B_region = self.BRegion
C_region = self.CRegion
C_buf = C_region.buffer
clear_accum = self.clear_accum
c_in_register = is_fragment(C_buf) or is_metal_simdgroup(C_buf)
@@ -73,41 +94,226 @@ class GemmMetal(GemmBase):
f"Metal GEMM requires C in local.fragment, metal.simdgroup or shared scope, got {C_buf.scope()}"
)
if self.is_gemm_ss():
if c_in_register:
@T.prim_func
def _gemm_ss_simdgroup() -> None:
A_local = T.alloc_local((warp_rows * 64), a_dtype, scope="metal.simdgroup")
B_local = T.alloc_local((warp_cols * 64), b_dtype, scope="metal.simdgroup")
if clear_accum:
for _i in T.serial(num_simd_c):
T.make_filled_simdgroup_matrix(C_buf.data, _i, T.cast(0, accum_dtype))
for ki in T.serial(0, (block_K // micro_size_k)):
mps_emitter.ldmatrix_a(A_local, A_region, ki)
mps_emitter.ldmatrix_b(B_local, B_region, ki)
mps_emitter.mma(A_local, B_local, C_buf)
return _Simplify(_gemm_ss_simdgroup, inline_let=True)
else:
@T.prim_func
def _gemm_ss_shared() -> None:
A_local = T.alloc_local((warp_rows * 64), a_dtype, scope="metal.simdgroup")
B_local = T.alloc_local((warp_cols * 64), b_dtype, scope="metal.simdgroup")
C_simd = T.alloc_local((num_simd_c * 64), accum_dtype, scope="metal.simdgroup")
if clear_accum:
for _i in T.serial(num_simd_c):
T.make_filled_simdgroup_matrix(C_simd.data, _i, T.cast(0, accum_dtype))
else:
mps_emitter.simd_load(C_simd, C_buf)
for ki in T.serial(0, (block_K // micro_size_k)):
mps_emitter.ldmatrix_a(A_local, A_region, ki)
mps_emitter.ldmatrix_b(B_local, B_region, ki)
mps_emitter.mma(A_local, B_local, C_simd)
mps_emitter.simd_store(C_simd, C_buf)
return _Simplify(_gemm_ss_shared, inline_let=True)
else:
if not self.is_gemm_ss():
raise ValueError(f"Unsupported gemm combination, A: {self.A.scope()}, B: {self.B.scope()}")
if c_in_register:
@T.prim_func
def _gemm_ss_simdgroup() -> None:
A_local = T.alloc_local((warp_rows * 64), a_dtype, scope="metal.simdgroup")
B_local = T.alloc_local((warp_cols * 64), b_dtype, scope="metal.simdgroup")
if clear_accum:
for _i in T.serial(num_simd_c):
T.make_filled_simdgroup_matrix(C_buf.data, _i, T.cast(0, accum_dtype))
for ki in T.serial(0, (block_K // micro_size_k)):
mps_emitter.ldmatrix_a(A_local, A_region, ki)
mps_emitter.ldmatrix_b(B_local, B_region, ki)
mps_emitter.mma(A_local, B_local, C_buf)
return _Simplify(_gemm_ss_simdgroup, inline_let=True)
@T.prim_func
def _gemm_ss_shared() -> None:
A_local = T.alloc_local((warp_rows * 64), a_dtype, scope="metal.simdgroup")
B_local = T.alloc_local((warp_cols * 64), b_dtype, scope="metal.simdgroup")
C_simd = T.alloc_local((num_simd_c * 64), accum_dtype, scope="metal.simdgroup")
if clear_accum:
for _i in T.serial(num_simd_c):
T.make_filled_simdgroup_matrix(C_simd.data, _i, T.cast(0, accum_dtype))
else:
mps_emitter.simd_load(C_simd, C_buf)
for ki in T.serial(0, (block_K // micro_size_k)):
mps_emitter.ldmatrix_a(A_local, A_region, ki)
mps_emitter.ldmatrix_b(B_local, B_region, ki)
mps_emitter.mma(A_local, B_local, C_simd)
mps_emitter.simd_store(C_simd, C_buf)
return _Simplify(_gemm_ss_shared, inline_let=True)
class GemmMetal(GemmBase):
def is_gemm_ss(self) -> bool:
return is_shared(self.A) and is_shared(self.B)
def is_gemm_gg(self) -> bool:
return is_global(self.A) and is_global(self.B)
@staticmethod
def _valid_gg_warp_partitions(M: int, N: int, num_warps: int):
for m_warp in range(1, num_warps + 1):
if num_warps % m_warp != 0:
continue
n_warp = num_warps // m_warp
if M % (m_warp * 16) == 0 and N % (n_warp * 32) == 0:
yield m_warp, n_warp
def _make_mps_emitter(self, target: Target, thread_nums: int):
from tilelang.metal.intrinsics.metal_macro_generator import MPSIntrinEmitter
m_warp, n_warp = self.policy.compute_warp_partition(self.M, self.N, thread_nums, target, GEMM_INST_METAL_COOPERATIVE_TENSOR)
if self.is_gemm_gg():
if int(thread_nums) % 32 != 0:
raise ValueError(f"Metal cooperative tensor GG requires threads to be a multiple of 32, got {thread_nums}")
num_warps = int(thread_nums) // 32
if num_warps <= 0:
raise ValueError(f"Metal cooperative tensor GG requires at least one warp, got {thread_nums} threads")
candidates = list(self._valid_gg_warp_partitions(int(self.M), int(self.N), num_warps))
if not candidates:
raise ValueError(
"Metal cooperative tensor GG requires a warp partition "
f"where M is divisible by m_warp*16 and N by n_warp*32; "
f"got tile ({self.M}, {self.N}) with {num_warps} warps"
)
# Prefer partitions where each simdgroup owns a balanced grid of
# 16x32 cooperative tensor operations. This keeps A/B cooperative
# tensor load counts balanced for direct GG tiles.
m_warp, n_warp = min(
candidates,
key=lambda part: (
abs(int(self.M) // (part[0] * 16) - int(self.N) // (part[1] * 32)),
-part[1],
),
)
warp_row_tiles = int(self.M // m_warp)
warp_col_tiles = int(self.N // n_warp)
return (
MPSIntrinEmitter(
a_dtype=self.a_dtype,
b_dtype=self.b_dtype,
accum_dtype=self.accum_dtype,
a_transposed=self.trans_A,
b_transposed=self.trans_B,
block_row_warps=m_warp,
block_col_warps=n_warp,
warp_row_tiles=warp_row_tiles,
warp_col_tiles=warp_col_tiles,
chunk=self.chunk,
),
m_warp,
n_warp,
)
def infer_layout(self, target: Target, thread_nums: int):
result = {}
if self.is_gemm_ss():
result[self.A] = _make_padded_layout(self.A)
result[self.B] = _make_padded_layout(self.B)
if is_fragment(self.C):
emitter, _, _ = self._make_mps_emitter(target, thread_nums)
result[self.C] = emitter.make_cooperative_tensor_store_layout(self.C)
return result
@staticmethod
def _get_padded_stride(buffer):
continuous = int(buffer.shape[-1])
element_bits = int(tvm.DataType(buffer.dtype).bits)
padded = continuous
if (element_bits * continuous) % 256 == 0:
padded += 128 // element_bits
return padded
def lower(
self,
layout_map: dict,
target: Target,
thread_bounds: Range,
thread_index: tir.PrimExpr,
mbar_phase_expr: tir.PrimExpr | None = None,
):
thread_nums = thread_bounds.extent
_, m_warp, n_warp = self._make_mps_emitter(target, int(thread_nums))
warp_row_tiles = int(self.M // m_warp)
warp_col_tiles = int(self.N // n_warp)
from tilelang.metal.intrinsics.metal_macro_generator import MPSIntrinEmitter
a_stride = self._get_padded_stride(self.A) if self.is_gemm_ss() else None
b_stride = self._get_padded_stride(self.B) if self.is_gemm_ss() else None
c_bytes_per_thread = warp_row_tiles * warp_col_tiles * 64
inner_k_steps = 2 if c_bytes_per_thread <= 128 else 1
output_dtype = self.accum_dtype
accum_dtype = T.float32 if self.is_gemm_gg() and str(output_dtype) in ("float16", "bfloat16") else output_dtype
mps_emitter = MPSIntrinEmitter(
a_dtype=self.a_dtype,
b_dtype=self.b_dtype,
accum_dtype=accum_dtype,
a_transposed=self.trans_A,
b_transposed=self.trans_B,
block_row_warps=m_warp,
block_col_warps=n_warp,
warp_row_tiles=warp_row_tiles,
warp_col_tiles=warp_col_tiles,
chunk=self.chunk,
thread_var=thread_index,
a_stride_override=a_stride,
b_stride_override=b_stride,
inner_k_steps=inner_k_steps,
)
a_dtype = self.a_dtype
b_dtype = self.b_dtype
warp_rows = mps_emitter.warp_rows
warp_cols = mps_emitter.warp_cols
num_simd_c = warp_rows * warp_cols
block_K = mps_emitter.chunk
micro_size_x = mps_emitter.micro_size_x
micro_size_y = mps_emitter.micro_size_y
micro_size_k = mps_emitter.micro_size_k
inner_k_steps = mps_emitter.inner_k_steps
a_tile_elems = micro_size_x * micro_size_k
b_tile_elems = micro_size_k * micro_size_y
c_tile_elems = micro_size_x * micro_size_y
A_region = self.ARegion
B_region = self.BRegion
C_region = self.CRegion
C_buf = C_region.buffer
clear_accum = self.clear_accum
c_in_cooperative_tensor = is_metal_cooperative_tensor(C_buf) or is_fragment(C_buf)
assert block_K >= micro_size_k, f"block_K ({block_K}) must be >= micro_size_k ({micro_size_k})"
if not (self.is_gemm_ss() or self.is_gemm_gg()):
raise ValueError(f"Unsupported gemm combination, A: {self.A.scope()}, B: {self.B.scope()}")
if c_in_cooperative_tensor:
assert is_full_region(C_region), "Fragment output C must be a full region"
@T.prim_func
def _gemm_cooperative_tensor() -> None:
A_local = T.alloc_local((warp_rows * a_tile_elems * inner_k_steps), a_dtype, scope="metal.cooperative_tensor")
B_local = T.alloc_local((warp_cols * b_tile_elems * inner_k_steps), b_dtype, scope="metal.cooperative_tensor")
if clear_accum:
for _i in T.serial(num_simd_c):
T.cooperative_tensor_fill(C_buf.data, _i, T.cast(0, accum_dtype), micro_size_x, micro_size_y)
for k_outer in T.serial(0, (block_K // (micro_size_k * inner_k_steps))):
for k_inner in T.serial(0, inner_k_steps):
ki = k_outer * inner_k_steps + k_inner
mps_emitter.ldmatrix_a(A_local, A_region, ki, k_inner)
mps_emitter.ldmatrix_b(B_local, B_region, ki, k_inner)
for k_inner in T.serial(0, inner_k_steps):
mps_emitter.mma(A_local, B_local, C_buf, k_inner)
return _Simplify(_gemm_cooperative_tensor, inline_let=True)
@T.prim_func
def _gemm_with_c_writeback() -> None:
A_local = T.alloc_local((warp_rows * a_tile_elems * inner_k_steps), a_dtype, scope="metal.cooperative_tensor")
B_local = T.alloc_local((warp_cols * b_tile_elems * inner_k_steps), b_dtype, scope="metal.cooperative_tensor")
C_ct = T.alloc_local((num_simd_c * c_tile_elems), accum_dtype, scope="metal.cooperative_tensor")
if clear_accum:
for _i in T.serial(num_simd_c):
T.cooperative_tensor_fill(C_ct.data, _i, T.cast(0, accum_dtype), micro_size_x, micro_size_y)
else:
mps_emitter.simd_load(C_ct, C_region)
for k_outer in T.serial(0, (block_K // (micro_size_k * inner_k_steps))):
for k_inner in T.serial(0, inner_k_steps):
ki = k_outer * inner_k_steps + k_inner
mps_emitter.ldmatrix_a(A_local, A_region, ki, k_inner)
mps_emitter.ldmatrix_b(B_local, B_region, ki, k_inner)
for k_inner in T.serial(0, inner_k_steps):
mps_emitter.mma(A_local, B_local, C_ct, k_inner)
mps_emitter.simd_store(C_ct, C_region)
return _Simplify(_gemm_with_c_writeback, inline_let=True)
+62 -3
View File
@@ -1,10 +1,12 @@
from __future__ import annotations
from platform import mac_ver
import re
import subprocess
from tvm.target import Target
from tilelang.backend.target import register_target_detector
from tilelang.backend.target import TargetLike, register_target_detector, register_target_normalizer
def _target_ffi_api():
@@ -21,14 +23,71 @@ def check_metal_availability() -> bool:
return arch == "arm64"
def _detect_metal_target() -> str | None:
def _parse_major_version(version: str) -> int:
try:
return int(version.split(".", 1)[0])
except (TypeError, ValueError):
return 0
def _command_stdout(cmd: list[str]) -> str | None:
try:
proc = subprocess.run(cmd, check=False, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, text=True)
except (OSError, subprocess.SubprocessError):
return None
if proc.returncode != 0:
return None
return proc.stdout.strip()
def check_metal4_availability() -> bool:
mac_release, _, arch = mac_ver()
if arch != "arm64" or _parse_major_version(mac_release) < 26:
return False
sdk_version = _command_stdout(["xcrun", "-sdk", "macosx", "--show-sdk-version"])
if sdk_version is None or _parse_major_version(sdk_version) < 26:
return False
cpu_brand = _command_stdout(["sysctl", "-n", "machdep.cpu.brand_string"]) or ""
match = re.search(r"\bApple M(\d+)\b", cpu_brand)
return match is not None and int(match.group(1)) >= 5
def _metal_target_config(enable_metal4: bool) -> dict[str, object]:
keys = ["metal", "gpu"]
if enable_metal4:
keys.append("metal4")
return {"kind": "metal", "keys": keys}
def _detect_metal_target() -> Target | str | None:
if check_metal_availability():
return "metal"
return Target(_metal_target_config(check_metal4_availability()))
return None
def normalize_metal_target(target: TargetLike) -> Target | None:
if isinstance(target, Target):
if target.kind.name == "metal":
return target
return None
if isinstance(target, dict):
if target.get("kind") != "metal":
return None
return Target(target)
if target.strip() != "metal":
return None
return Target(_metal_target_config(check_metal4_availability()))
def target_is_metal(target: Target) -> bool:
return _target_ffi_api().TargetIsMetal(target)
def target_metal_supports_metal4(target: Target) -> bool:
return _target_ffi_api().TargetMetalSupportsMetal4(target)
register_target_detector("metal", _detect_metal_target, override=True)
register_target_normalizer("metal", normalize_metal_target, override=True)
@@ -1,8 +1,8 @@
"""Rewrite local.fragment to metal.simdgroup for GEMM accumulator buffers on Metal.
"""Rewrite local.fragment to metal.simdgroup for legacy Metal GEMM accumulators.
This pass runs after pipelining and before LayoutInference, so that
simdgroup matrices (which are hardware-opaque and have no explicit
thread-level layout) are never seen by LayoutInference.
M5 cooperative-tensor GEMM keeps fragment/local buffers in regular scopes until
Metal codegen sees explicit tl.cooperative_tensor_* builtins. Only legacy
simdgroup GEMM requires changing the accumulator scope before layout inference.
"""
from __future__ import annotations
@@ -18,17 +18,11 @@ from tvm.tirx.transform import prim_func_pass
@lru_cache(maxsize=1)
def _get_gemm_ops():
return frozenset(
{
Op.get("tl.tileop.gemm"),
}
)
return frozenset({Op.get("tl.tileop.gemm")})
def _extract_buffer_var_from_region(region_call):
if not isinstance(region_call, tir.Call):
return None
if len(region_call.args) < 1:
if not isinstance(region_call, tir.Call) or len(region_call.args) < 1:
return None
buf_load = region_call.args[0]
if isinstance(buf_load, tir.BufferLoad):
@@ -36,6 +30,26 @@ def _extract_buffer_var_from_region(region_call):
return None
def _get_num_warps_from_body(body: tir.Stmt) -> int:
warp_size = 32
num_threads = None
def _visitor(stmt):
nonlocal num_threads
if (
isinstance(stmt, tir.AttrStmt)
and stmt.attr_key == "thread_extent"
and hasattr(stmt.node, "thread_tag")
and "threadIdx.x" in str(stmt.node.thread_tag)
):
val = stmt.value
if isinstance(val, tir.IntImm):
num_threads = val.value
tir.stmt_functor.post_order_visit(body, _visitor)
return num_threads // warp_size if num_threads is not None else 1
def _collect_fragment_gemm_accum_vars(body: tir.Stmt) -> set:
accum_vars: set = set()
gemm_ops = _get_gemm_ops()
@@ -54,13 +68,17 @@ def _collect_fragment_gemm_accum_vars(body: tir.Stmt) -> set:
return accum_vars
def _remap_buffer(buf, var_map):
def _remap_buffer(buf, var_map, num_warps=1):
old_data = buf.data
new_data = var_map.get(old_data, None)
if new_data is None:
return buf
total = 1
for s in buf.shape:
total *= s.value if isinstance(s, tir.IntImm) else s
new_total = total // num_warps if num_warps > 1 else total
return tir.decl_buffer(
buf.shape,
[tir.IntImm("int32", new_total)],
buf.dtype,
buf.name,
data=new_data,
@@ -71,7 +89,6 @@ def _remap_buffer(buf, var_map):
def _remap_buffer_region(region, buf_map):
"""Remap buffer inside a BufferRegion using buf_map."""
if region is None:
return region
new_buf = buf_map.get(region.buffer, None)
@@ -81,27 +98,17 @@ def _remap_buffer_region(region, buf_map):
def _remap_match_buffer(match, buf_map):
"""Remap buffer inside a MatchBufferRegion using buf_map."""
if match is None:
return match
new_buf = buf_map.get(match.buffer, None)
new_src = _remap_buffer_region(match.source, buf_map)
if new_buf is None and new_src is match.source:
return match
return tir.MatchBufferRegion(
new_buf if new_buf is not None else match.buffer,
new_src,
)
return tir.MatchBufferRegion(new_buf if new_buf is not None else match.buffer, new_src)
def _rewrite_scope(body, var_map):
# Phase 1: Substitute Var references in the entire body, including inside
# BufferRegion objects nested within call_intrin arguments (e.g., gemm op
# calls). The ir_transform approach below only visits SBlock/AllocBuffer
# nodes and the per-SBlock substitute() does not reach buffer vars embedded
# inside opaque Call args processed by subsequent lowering passes.
def _rewrite_scope(body, var_map, num_warps=1):
body = tir.stmt_functor.substitute(body, var_map)
buf_map = {}
def _pre_order(stmt):
@@ -109,7 +116,7 @@ def _rewrite_scope(body, var_map):
new_alloc_bufs = []
changed = False
for buf in stmt.alloc_buffers:
new_buf = _remap_buffer(buf, var_map)
new_buf = _remap_buffer(buf, var_map, num_warps)
new_alloc_bufs.append(new_buf)
if not new_buf.same_as(buf):
buf_map[buf] = new_buf
@@ -130,7 +137,7 @@ def _rewrite_scope(body, var_map):
stmt.annotations,
)
elif isinstance(stmt, tir.AllocBuffer):
new_buf = _remap_buffer(stmt.buffer, var_map)
new_buf = _remap_buffer(stmt.buffer, var_map, num_warps)
if not new_buf.same_as(stmt.buffer):
buf_map[stmt.buffer] = new_buf
return tir.AllocBuffer(new_buf, stmt.annotations, stmt.span)
@@ -148,15 +155,14 @@ def _metal_fragment_to_simdgroup(func: tir.PrimFunc, mod: IRModule, ctx) -> tir.
if not accum_vars:
return func
num_warps = _get_num_warps_from_body(func.body)
var_map: dict = {}
for var in accum_vars:
ptr_type = var.type_annotation
new_ptr = PointerType(ptr_type.element_type, "metal.simdgroup")
new_var = tir.Var(var.name, new_ptr)
var_map[var] = new_var
var_map[var] = tir.Var(var.name, new_ptr)
new_body = _rewrite_scope(func.body, var_map)
return func.with_body(new_body)
return func.with_body(_rewrite_scope(func.body, var_map, num_warps))
MetalFragmentToSimdgroup = prim_func_pass(
+16
View File
@@ -0,0 +1,16 @@
from __future__ import annotations
from tilelang._typing import BufferLikeType
from tilelang.utils.language import _get_buffer
def is_metal_cooperative_tensor(buffer: BufferLikeType) -> bool:
"""Check if the buffer is in the Metal cooperative tensor scope."""
buffer = _get_buffer(buffer)
return buffer.scope() == "metal.cooperative_tensor"
def is_metal_simdgroup(buffer: BufferLikeType) -> bool:
"""Check if the buffer is in the Metal simdgroup scope."""
buffer = _get_buffer(buffer)
return buffer.scope() == "metal.simdgroup"
+9 -3
View File
@@ -75,14 +75,20 @@ from tvm.tirx.transform import prim_func_pass
# Cache the Op for if_then_else to avoid repeated lookups
_IF_THEN_ELSE_OP = Op.get("tirx.if_then_else")
from tilelang.utils.language import is_fragment, is_global, is_local, is_local_var, is_shared, is_metal_simdgroup
from tilelang.utils.language import (
is_fragment,
is_global,
is_local,
is_local_var,
is_shared,
)
def is_local_buffer(buffer: Buffer) -> bool:
"""Check if a buffer is local (register-level), including local.var and metal.simdgroup."""
"""Check if a buffer is local/register-level."""
if buffer is None:
return False
return is_local(buffer) or is_fragment(buffer) or is_local_var(buffer) or is_metal_simdgroup(buffer)
return is_local(buffer) or is_fragment(buffer) or is_local_var(buffer)
def is_global_or_shared_buffer(buffer: Buffer) -> bool:
-14
View File
@@ -119,20 +119,6 @@ def is_fragment(buffer: BufferLikeType) -> bool:
return buffer.scope().startswith("local.fragment")
def is_metal_simdgroup(buffer: BufferLikeType) -> bool:
"""
Check if the buffer is in the Metal simdgroup scope.
Args:
buffer: The TVM buffer, BufferLoad, or BufferRegion to check.
Returns:
bool: True if the buffer is in metal.simdgroup scope, False otherwise.
"""
buffer = _get_buffer(buffer)
return buffer.scope() == "metal.simdgroup"
def is_local_var(buffer: BufferLikeType) -> bool:
"""
Check if the buffer is in the local.var memory scope.