mirror of
https://github.com/tile-ai/tilelang.git
synced 2026-10-02 06:34:36 +08:00
[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:
@@ -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
|
||||
```
|
||||
@@ -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}
|
||||
|
||||
+1218
-16
File diff suppressed because it is too large
Load Diff
@@ -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_;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
};
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -13,6 +13,7 @@ namespace tl {
|
||||
|
||||
bool TargetIsMetal(Target target);
|
||||
int TargetMetalGetWarpSize(Target target);
|
||||
bool TargetMetalSupportsMetal4(Target target);
|
||||
|
||||
} // namespace tl
|
||||
} // namespace tvm
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
@@ -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)]
|
||||
*/
|
||||
|
||||
@@ -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)) {
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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,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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user