Skip to content

Commit 358addc

Browse files
committed
fix Metal GEMM dispatch and update tests for cooperative_tensor MMA
- gemm.cc: add TargetIsMetal case in getGemmInst (was falling through to fatal error after upstream rebase removed catch-all) - test_metal_gemm_v2_linux: update block sizes to be valid for MMA(16,32,16), check for matmul2d or simdgroup intrinsics - test_metal_simdgroup_store: update block sizes and codegen assertions for cooperative_tensor path, use dynamic thread count
1 parent 0ec9cb4 commit 358addc

3 files changed

Lines changed: 27 additions & 32 deletions

File tree

‎src/op/gemm.cc‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -189,6 +189,8 @@ GemmInst GemmNode::getGemmInst(int block_size, Target target) const {
189189
return GemmInst::kWMMA;
190190
} else if (TargetIsCuda(target)) {
191191
return GemmInst::kMMA;
192+
} else if (TargetIsMetal(target)) {
193+
return GemmInst::kMetalExp;
192194
} else if (TargetIsCPU(target)) {
193195
return GemmInst::kScalar;
194196
} else {

‎testing/python/metal/test_metal_gemm_v2_linux.py‎

Lines changed: 11 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,8 @@ def main(
1818
B: T.Tensor((K, N), dtype),
1919
C: T.Tensor((M, N), accum_dtype),
2020
):
21-
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
21+
num_threads = max(32, (block_M // 16) * (block_N // 32) * 32)
22+
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=num_threads) as (bx, by):
2223
A_shared = T.alloc_shared((block_M, block_K), dtype, scope="shared")
2324
B_shared = T.alloc_shared((block_K, block_N), dtype, scope="shared")
2425
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
@@ -53,29 +54,27 @@ def assert_metal_gemm_v2_codegen(
5354
src_code = artifact.kernel_source
5455
assert src_code is not None
5556
assert "kernel void" in src_code
56-
# Verify simdgroup matrix operations are present
57-
assert "simdgroup_multiply_accumulate" in src_code
58-
assert "simdgroup_load" in src_code
59-
assert "simdgroup_store" in src_code
57+
# Verify matrix operations are present (cooperative_tensor/matmul2d or simdgroup)
58+
has_cooperative = "matmul2d" in src_code or "cooperative_tensor" in src_code
59+
has_simdgroup = "simdgroup_multiply_accumulate" in src_code
60+
assert has_cooperative or has_simdgroup, f"Expected matmul2d or simdgroup_multiply_accumulate in Metal source"
6061

6162

6263
def test_metal_gemm_v2_float16():
63-
assert_metal_gemm_v2_codegen(64, 64, 64, 16, 16, 16, dtype=T.float16)
64+
assert_metal_gemm_v2_codegen(128, 128, 128, 32, 32, 32, dtype=T.float16)
6465

6566

6667
def test_metal_gemm_v2_float32():
67-
assert_metal_gemm_v2_codegen(64, 64, 64, 16, 16, 16, dtype=T.float32, accum_dtype=T.float32)
68+
assert_metal_gemm_v2_codegen(128, 128, 128, 32, 32, 32, dtype=T.float32, accum_dtype=T.float32)
6869

6970

7071
def test_metal_gemm_v2_larger():
71-
assert_metal_gemm_v2_codegen(128, 128, 128, 32, 32, 32, dtype=T.float16)
72+
assert_metal_gemm_v2_codegen(128, 128, 128, 64, 64, 32, dtype=T.float16)
7273

7374

7475
def test_metal_gemm_v2_small_blocks():
75-
"""Test with blocks where warp_rows > 1 and warp_cols > 1, which previously
76-
produced incorrect results due to swizzle padding changing the stride.
77-
"""
78-
assert_metal_gemm_v2_codegen(16, 16, 16, 16, 16, 16, dtype=T.float16)
76+
"""Test with minimum valid block sizes for cooperative_tensor MMA(16,32,16)."""
77+
assert_metal_gemm_v2_codegen(32, 32, 32, 32, 32, 32, dtype=T.float16)
7978

8079

8180
if __name__ == "__main__":

‎testing/python/metal/test_metal_simdgroup_store.py‎

Lines changed: 14 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,8 @@ def gemm_kernel(
2020
B: T.Tensor((K, N), dtype),
2121
C: T.Tensor((M, N), accum_dtype),
2222
):
23-
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
23+
num_threads = max(32, (block_M // 16) * (block_N // 32) * 32)
24+
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=num_threads) as (bx, by):
2425
A_shared = T.alloc_shared((block_M, block_K), dtype, scope="shared")
2526
B_shared = T.alloc_shared((block_K, block_N), dtype, scope="shared")
2627
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
@@ -64,53 +65,46 @@ def assert_simdgroup_store_codegen(M, N, K, block_M, block_N, block_K, dtype=T.f
6465
src = artifact.kernel_source
6566
assert src is not None
6667
assert "kernel void" in src
67-
assert "simdgroup_multiply_accumulate" in src
68-
assert "make_filled_simdgroup_matrix" in src
69-
70-
assert "simdgroup_float8x8" in src or "simdgroup_half8x8" in src, "Expected simdgroup_float8x8 or simdgroup_half8x8 for C accumulator"
71-
72-
store_to_device = src.count("simdgroup_store(C_local")
73-
assert store_to_device > 0, "Expected simdgroup_store of C_local to device memory"
74-
75-
load_c_from_shared = [line for line in src.split("\n") if "simdgroup_load" in line and "C_local" in line]
76-
assert len(load_c_from_shared) == 0, f"C_local should not be loaded from shared memory, but found: {load_c_from_shared}"
68+
has_cooperative = "matmul2d" in src
69+
has_simdgroup = "simdgroup_multiply_accumulate" in src
70+
assert has_cooperative or has_simdgroup, "Expected matmul2d or simdgroup_multiply_accumulate"
7771

7872

7973
# --- Codegen tests (cross-platform) ---
8074

8175

8276
def test_codegen_square_small():
83-
assert_simdgroup_store_codegen(64, 64, 64, 16, 16, 16)
77+
assert_simdgroup_store_codegen(64, 64, 64, 32, 32, 32)
8478

8579

8680
def test_codegen_square_large():
87-
assert_simdgroup_store_codegen(128, 128, 128, 32, 32, 32)
81+
assert_simdgroup_store_codegen(128, 128, 128, 64, 64, 32)
8882

8983

9084
def test_codegen_non_square():
91-
assert_simdgroup_store_codegen(128, 128, 128, 32, 64, 16)
85+
assert_simdgroup_store_codegen(128, 128, 128, 32, 64, 32)
9286

9387

9488
def test_codegen_float32_accum():
95-
assert_simdgroup_store_codegen(64, 64, 64, 16, 16, 16, dtype=T.float32, accum_dtype=T.float32)
89+
assert_simdgroup_store_codegen(64, 64, 64, 32, 32, 32, dtype=T.float32, accum_dtype=T.float32)
9690

9791

9892
# --- Correctness tests (require Metal hardware) ---
9993

10094

10195
@tilelang.testing.requires_metal
102-
def test_correctness_16x16x16():
103-
assert_simdgroup_store_correctness(128, 128, 128, 16, 16, 16)
96+
def test_correctness_32x32x32():
97+
assert_simdgroup_store_correctness(128, 128, 128, 32, 32, 32)
10498

10599

106100
@tilelang.testing.requires_metal
107-
def test_correctness_32x32x32():
108-
assert_simdgroup_store_correctness(128, 128, 128, 32, 32, 32)
101+
def test_correctness_64x64x32():
102+
assert_simdgroup_store_correctness(128, 128, 128, 64, 64, 32)
109103

110104

111105
@tilelang.testing.requires_metal
112106
def test_correctness_non_square_block():
113-
assert_simdgroup_store_correctness(128, 128, 128, 32, 64, 16)
107+
assert_simdgroup_store_correctness(128, 128, 128, 32, 64, 32)
114108

115109

116110
@tilelang.testing.requires_metal

0 commit comments

Comments
 (0)