Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
Show all changes
39 commits
Select commit Hold shift + click to select a range
9f9677a
support 2-cta alloc, dealloc, umma_arrive
Rachmanino Feb 26, 2026
fbc6bec
draft tma load 2sm support
Rachmanino Feb 26, 2026
4dc2901
draft tcgen5 2sm support
Rachmanino Mar 2, 2026
9b048a4
add draft lower_blackwell_2sm pass
Rachmanino Mar 3, 2026
2fc3064
support threadblock swizzle annotation for cluster launch
Rachmanino Mar 6, 2026
8b27b40
enhance cuda arch restriction for cluster template functions
Rachmanino Mar 6, 2026
ae83639
Change return type in block_rank_in_cluster function from uint32 to i…
Rachmanino Mar 6, 2026
d45865a
fix threadblock swizzle for cluster
Rachmanino Mar 6, 2026
d2d7111
introduce tcgen05.fence::{before, after}_thread_sync
Rachmanino Mar 6, 2026
20f01be
update lower_blackwell_2sm pass
Rachmanino Mar 6, 2026
ebb8b76
update b_continuity calculation in GemmTCGEN5 for compatiblity with 2cta
Rachmanino Mar 6, 2026
150fb9b
use m/n per_cta for offset calculation in 2cta tcgen5 and introduce a…
Rachmanino Mar 6, 2026
6ddea2a
draft refactor of tcgen05mma
Rachmanino Mar 6, 2026
ec9cef9
successfully generated correct 2sm code
Rachmanino Mar 6, 2026
dcc1bae
lint
Rachmanino Mar 6, 2026
d09fc92
fix bug caused by rebase
Rachmanino Mar 9, 2026
a76f221
upd 2sm persistent kernel
Rachmanino Mar 9, 2026
28617da
reorgnize tcgen5mma examples, fix cross-wave bugs, optimize to 1670T
Rachmanino Mar 12, 2026
d229491
lint
Rachmanino Mar 12, 2026
dbc9388
general support for 2cta in GetTCGEN5MMAMeta
Rachmanino Mar 12, 2026
a01e62d
drop the change of make_tcgen05mma_swizzled_layout
Rachmanino Mar 12, 2026
5ddea83
support 2cta for more dtypes and gemm_ts
Rachmanino Mar 12, 2026
d832d14
add check for cluster_dims when lowering 2cta tcgen5
Rachmanino Mar 12, 2026
97beb2c
refactor 2sm conditional lowering
Rachmanino Mar 13, 2026
9572c64
remove unnecessary syncthreads
Rachmanino Mar 13, 2026
15e3048
support Layout B and add maint test
Rachmanino Mar 13, 2026
d901fed
lint
Rachmanino Mar 13, 2026
605d9dc
Restore support for cp.async.mbarrier and cleaning up unused code
Rachmanino Mar 16, 2026
2ec5fe8
Merge branch 'main' into wt/2sm
Rachmanino Mar 20, 2026
4131990
Merge branch 'main' into wt/2sm
Rachmanino Mar 20, 2026
33ac055
fix tma injection pass change
Rachmanino Mar 20, 2026
1caa424
typo
Rachmanino Mar 20, 2026
6c09a53
Merge branch 'main' into wt/2sm
Rachmanino Mar 23, 2026
d174892
support tma_load_2sm lowering after merge main
Rachmanino Mar 23, 2026
6e698b7
refactor to expose use_2cta option in `T.tcgen05_gemm` and remove 2ct…
Rachmanino Mar 23, 2026
3795413
lint
Rachmanino Mar 23, 2026
d42abc8
fix cutedsl codegen for threadblock swizzle
Rachmanino Mar 23, 2026
e7e03b6
lint
Rachmanino Mar 23, 2026
28f4224
upd maint script for tcgen05
Rachmanino Mar 23, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Prev Previous commit
Next Next commit
support Layout B and add maint test
  • Loading branch information
Rachmanino committed Mar 18, 2026
commit 15e3048ea67bcbc9af14855e93be4da3c53cc892
174 changes: 174 additions & 0 deletions maint/gemm_v2/correctness_evaluation_tcgen05_2cta.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,174 @@
# pytest correctness_evaluation_tcgen05_2cta.py -n 32
import pytest
from tilelang import tvm as tvm
import tilelang
import tilelang.testing
import tilelang.language as T
tilelang.disable_cache()

def matmul_2cta(
M,
N,
K,
block_M,
block_N,
block_K,
in_dtype,
out_dtype,
accum_dtype,
num_stages,
):
@T.prim_func
def main(
A: T.Tensor((M, K), in_dtype),
B: T.Tensor((K, N), in_dtype),
C: T.Tensor((M, N), out_dtype),
):
with T.Kernel(T.ceildiv(M, block_M), T.ceildiv(N, block_N), threads=128, cluster_dims=2) as (bx, by):
A_shared = T.alloc_shared((num_stages, block_M, block_K), in_dtype)
B_shared = T.alloc_shared((num_stages, block_K, block_N // 2), in_dtype) # each CTA holds half of B
C_tmem = T.alloc_tmem([block_M, block_N], accum_dtype)
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
C_shared = T.alloc_shared((block_M, block_N), out_dtype)
loaded = T.alloc_cluster_barrier([32 * 2] * num_stages)
consumed = T.alloc_cluster_barrier([1] * num_stages)
tmem_full = T.alloc_barrier([1])

tx = T.get_thread_binding()
cta_id = T.block_rank_in_cluster()
T.assume(cta_id < 2)

T.use_swizzle(16)

if tx < 32: # warp 0: issue TMA loads
for k in T.serial(T.ceildiv(K, block_K)):
T.mbarrier_wait_parity(consumed[k % num_stages], ((k // num_stages) & 1) ^ 1)
T.copy(
A[bx * block_M:(bx + 1) * block_M, k * block_K:(k + 1) * block_K],
A_shared[k % num_stages, :, :],
)
T.copy(
B[k * block_K:(k + 1) * block_K,
(by * 2 + cta_id) * (block_N // 2):(by * 2 + cta_id + 1) * (block_N // 2)],
B_shared[k % num_stages, :, :],
)
T.mbarrier_arrive(loaded[k % num_stages], 0) # arrive on leader CTA's barrier
elif cta_id == 0 and tx < 64: # warp 1 on leader CTA: issue tcgen5 MMA
for k in T.serial(T.ceildiv(K, block_K)):
T.mbarrier_wait_parity(loaded[k % num_stages], (k // num_stages) & 1)
T.gemm(
A_shared[k % num_stages, :, :],
B_shared[k % num_stages, :, :],
C_tmem,
mbar=consumed[k % num_stages],
wg_wait=-1,
clear_accum=k == 0,
)
T.tcgen05_mma_arrive(tmem_full, arrive_2cta=True)

T.mbarrier_wait_parity(tmem_full, 0)
T.copy(C_tmem, C_local)
T.copy(C_local, C_shared)
T.copy(C_shared, C[bx * block_M, by * block_N])

return main


def _compile_and_check(program, out_dtype):
kernel = tilelang.compile(program, out_idx=[2], execution_backend='cython')

print(kernel.get_kernel_source())

profiler = kernel.get_profiler(tensor_supply_type=tilelang.TensorSupplyType.Normal)

def ref_program(A, B):
import torch

C = torch.matmul(A.to(torch.float), B.to(torch.float))
return C.to(torch.__getattribute__(out_dtype))

profiler.assert_allclose(ref_program, atol=1e-2, rtol=1e-2)
print("assert_allclose passed")


def run_gemm(
M,
N,
K,
in_dtype,
out_dtype,
accum_dtype,
block_M,
block_N,
block_K,
num_stages=4,
):
program = matmul_2cta(
M,
N,
K,
block_M,
block_N,
block_K,
in_dtype,
out_dtype,
accum_dtype,
num_stages,
)
_compile_and_check(program, out_dtype)


M_VALUES = [64, 128, 256]
N_VALUES = [64, 128, 256]

# atom_k=16 for fp16/bf16 (K%16==0), atom_k=32 for fp8/int8 (K%32==0)
K_VALUES_16 = [16, 32, 64, 128]
K_VALUES_32 = [32, 64, 128]

# Dtype cases: (block_K, in_dtype, out_dtype, accum_dtype)
FP16_CASES = [
pytest.param(k, T.float16, T.float32, T.float32, id=f"K{k}-fp16-fp32-fp32")
for k in K_VALUES_16
]

FP8_E5M2_CASES = [
pytest.param(k, T.float8_e5m2, T.float32, T.float32, id=f"K{k}-fp8e5m2-fp32-fp32")
for k in K_VALUES_32
]

INT8_CASES = [
pytest.param(k, T.int8, T.int32, T.int32, id=f"K{k}-int8-int32-int32")
for k in K_VALUES_32
]

ALL_DTYPE_CASES = FP16_CASES + FP8_E5M2_CASES + INT8_CASES


@pytest.mark.parametrize("m", M_VALUES, ids=lambda v: f"M{v}")
@pytest.mark.parametrize("n", N_VALUES, ids=lambda v: f"N{v}")
@pytest.mark.parametrize("block_k,in_dtype,out_dtype,accum_dtype", ALL_DTYPE_CASES)
def test_gemm_2cta(m, n, block_k, in_dtype, out_dtype, accum_dtype):
import torch

for attr in {in_dtype, out_dtype, accum_dtype}:
if not hasattr(torch, attr):
pytest.skip(f"Torch does not expose dtype {attr}")

# M = 2 * block_M so ceildiv(M, block_M) = 2 (cluster needs >= 2 tiles in M dim)
# K = 3 * block_K to exercise multi-iteration pipelining
k = block_k * 3
run_gemm(
m * 2,
n,
k,
in_dtype,
out_dtype,
accum_dtype,
block_M=m,
block_N=n,
block_K=block_k,
)


if __name__ == "__main__":
tilelang.testing.main()
19 changes: 13 additions & 6 deletions tilelang/intrinsics/tcgen05_macro_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -622,12 +622,19 @@ def forward(i: PrimExpr, j: PrimExpr):
aj + atom_idx * atom_n,
]
if atom_m == 128:
# Layout D
print(f"Layout D: ai={ai}, aj={aj}, atom_idx={atom_idx}")
return [
ai,
aj + atom_idx * atom_n,
]
if enable_2cta:
# Layout B
half_atom_n = atom_n // 2
return [
ai + (aj // half_atom_n) * 64,
(aj % half_atom_n) + atom_idx * half_atom_n,
]
else:
# Layout D
return [
ai,
aj + atom_idx * atom_n,
]
if atom_m == 64:
# Layout E (.ws variant)
half_atom_n = atom_n // 2
Expand Down