Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
78 changes: 72 additions & 6 deletions examples/gemm_sm100/gemm_tcgen5mma_ws.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
# Non-persistent, 1-SM GEMM
# Non-persistent

import torch
import tilelang
Expand Down Expand Up @@ -59,8 +59,72 @@ def gemm(A, B, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_

# Wait for all tcgen5 to finish
T.mbarrier_wait_parity(tmem_full, 0)
T.copy(C_tmem, C_local)
if use_tma_store:
T.copy(C_local, C_shared)
T.copy(C_shared, C[by * block_M, bx * block_N])
else:
T.copy(C_local, C_local_cast)
T.copy(C_local_cast, C[by * block_M, bx * block_N]) # STG256
return C


@tilelang.jit
def gemm_2cta(A, B, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages, use_tma_store=True):
M, N, K = T.const("M, N, K")

k_iters = T.ceildiv(K, block_K)

A: T.Tensor[[M, K], in_dtype]
B: T.Tensor[[K, N], in_dtype]
C = T.empty((M, N), out_dtype)

with T.Kernel(T.ceildiv(M, block_M), T.ceildiv(N, block_N), threads=128, cluster_dims=2) as (by, bx):
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 hold 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)
C_local_cast = T.alloc_fragment((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])
Comment thread
coderabbitai[bot] marked this conversation as resolved.

tx = T.get_thread_binding()
cta_id = T.block_rank_in_cluster()
T.assume(cta_id < 2) # todo: automatically assume this

T.sync_threads() # TileLang won't generate this if not annotated
T.use_swizzle(16) # TL will perform auto threadblock swizzle with cluster

if tx < 32: # warp 0: issue tma
for k in T.serial(k_iters):
T.mbarrier_wait_parity(consumed[k % num_stages], ((k // num_stages) & 1) ^ 1)
T.tma_copy(
A[by * block_M : (by + 1) * block_M, k * block_K : (k + 1) * block_K],
A_shared[k % num_stages, :, :],
barrier=loaded[k % num_stages],
)
T.tma_copy(
B[k * block_K : (k + 1) * block_K, (bx * 2 + cta_id) * (block_N // 2) : (bx * 2 + cta_id + 1) * (block_N // 2)],
B_shared[k % num_stages, :, :],
barrier=loaded[k % num_stages],
)
T.mbarrier_arrive(loaded[k % num_stages], 0) # arrive on leader cta's barrier
elif cta_id == 0 and tx < 64: # Only warp 1 on leader cta issues tcgen5
for k in T.serial(k_iters):
T.mbarrier_wait_parity(loaded[k % num_stages], (k // num_stages) & 1)
T.tcgen05_gemm(
A_shared[k % num_stages, :, :],
B_shared[k % num_stages, :, :],
C_tmem,
mbar=consumed[k % num_stages],
clear_accum=k == 0,
use_2cta=True,
)
T.tcgen05_mma_arrive(tmem_full, arrive_2cta=True)

# Wait for all tcgen5 to finish
T.mbarrier_wait_parity(tmem_full, 0)
T.copy(C_tmem, C_local)
if use_tma_store:
T.copy(C_local, C_shared)
Expand All @@ -75,18 +139,20 @@ def main():
M, N, K = 8192, 8192, 8192
block_M, block_N, block_K = 128, 256, 64
in_dtype, out_dtype, accum_dtype = T.bfloat16, T.bfloat16, T.float
num_stages = 4
enable_2cta_tcgen5mma = True
num_stages = 6 if enable_2cta_tcgen5mma else 4 # Each cta only needs to load half of B, enabling larger stages
kernel = gemm_2cta if enable_2cta_tcgen5mma else gemm

a = torch.randn(M, K, device="cuda", dtype=torch.bfloat16)
b = torch.randn(K, N, device="cuda", dtype=torch.bfloat16)
c = gemm(a, b, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages)
print(gemm.get_kernel_source(a, b, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages))
c = kernel(a, b, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages)
print(kernel.get_kernel_source(a, b, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages))

ref_c = (a.to(torch.float) @ b.to(torch.float)).to(torch.bfloat16)
torch.testing.assert_close(c, ref_c, rtol=1e-2, atol=1e-2)
print("All checks passed. ✅")

tl_latency = do_bench(lambda: gemm(a, b, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages), backend="cupti")
tl_latency = do_bench(lambda: kernel(a, b, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages), backend="cupti")
torch_latency = do_bench(lambda: a @ b, backend="cupti")
print(f"Tilelang latency: {tl_latency} ms")
print(f"Flops: {2 * M * N * K / (tl_latency / 1e3) / 1e12} TFLOPS")
Expand Down
171 changes: 156 additions & 15 deletions examples/gemm_sm100/gemm_tcgen5mma_ws_persistent.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
# Persistent, 1-SM, num_epi_stages = 2
# Persistent, num_epi_stages = 2

import torch
import tilelang
Expand All @@ -8,7 +8,7 @@


@tilelang.jit
def gemm(
def gemm_persistent(
A,
B,
block_M,
Expand All @@ -34,7 +34,7 @@ def gemm(
k_blocks = T.ceildiv(K, block_K)
waves = T.ceildiv(m_blocks * n_blocks, sm_num)
group_size = 8
assert n_blocks % group_size == 0
assert n_blocks % (2 * group_size) == 0 # Please adjust group_size if not satisfied

with T.Kernel(sm_num, threads=256) as (block_id):
A_shared = T.alloc_shared((num_stages, block_M, block_K), in_dtype)
Expand All @@ -59,18 +59,19 @@ def gemm(

if bx * block_M < M and by * block_N < N:
for k in T.serial(k_blocks):
T.mbarrier_wait_parity(consumed[k % num_stages], ((k // num_stages) & 1) ^ 1)
phase = w * k_blocks + k
T.mbarrier_wait_parity(consumed[phase % num_stages], ((phase // num_stages) & 1) ^ 1)
T.tma_copy(
A[bx * block_M : (bx + 1) * block_M, k * block_K : (k + 1) * block_K],
A_shared[k % num_stages, :, :],
barrier=loaded[k % num_stages],
A_shared[phase % num_stages, :, :],
barrier=loaded[phase % num_stages],
)
T.tma_copy(
B[k * block_K : (k + 1) * block_K, by * block_N : (by + 1) * block_N],
B_shared[k % num_stages, :, :],
barrier=loaded[k % num_stages],
B_shared[phase % num_stages, :, :],
barrier=loaded[phase % num_stages],
)
T.mbarrier_arrive(loaded[k % num_stages])
T.mbarrier_arrive(loaded[phase % num_stages])

elif tx < 64: # warp 1: issue tcgen5
for w in T.unroll(waves):
Expand All @@ -81,7 +82,8 @@ def gemm(
if bx * block_M < M and by * block_N < N:
T.mbarrier_wait_parity(tmem_empty[w & 1], ((w // 2) & 1) ^ 1)
for k in T.serial(k_blocks):
T.mbarrier_wait_parity(loaded[k % num_stages], (k // num_stages) & 1)
phase = w * k_blocks + k
T.mbarrier_wait_parity(loaded[phase % num_stages], (phase // num_stages) & 1)
if w & 1 == 0:
T.tcgen05_gemm(
A_shared[k % num_stages, :, :],
Expand Down Expand Up @@ -129,24 +131,163 @@ def gemm(
return C


@tilelang.jit
def gemm_persistent_2cta(
A,
B,
block_M,
block_N,
store_block_N, # block_N for C_shared
block_K,
in_dtype,
out_dtype,
accum_dtype,
num_stages,
use_tma_store=True,
):
M, N, K = T.const("M, N, K")

A: T.Tensor[[M, K], in_dtype]
B: T.Tensor[[K, N], in_dtype]
C = T.empty((M, N), out_dtype)

sm_num = driver.get_num_sms()
num_clusters = sm_num // 2
m_blocks = T.ceildiv(M, block_M)
m_clusters = m_blocks // 2
n_blocks = T.ceildiv(N, block_N)
assert K % (2 * block_K) == 0 # for simplicity
k_blocks = T.ceildiv(K, block_K)
waves = T.ceildiv(m_blocks * n_blocks, sm_num)
group_size = 8 # in cluster
assert n_blocks % (2 * group_size) == 0 # Please adjust group_size if not satisfied

with T.Kernel(sm_num, threads=256, cluster_dims=2) as (block_id):
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)
C_tmem_0 = T.alloc_tmem([block_M, block_N], accum_dtype)
C_tmem_1 = T.alloc_tmem([block_M, block_N], accum_dtype)
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
C_local_cast = T.alloc_fragment((block_M, block_N), out_dtype)
C_shared = T.alloc_shared((block_M, store_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_cluster_barrier([1] * 2)
tmem_empty = T.alloc_cluster_barrier([128 * 2] * 2)

tx = T.get_thread_binding()
cta_id = T.block_rank_in_cluster()
T.assume(cta_id < 2) # todo: automatically assume this

if tx < 32: # warp 0: issue tma
for w in T.unroll(waves):
# manual threadblock swizzle
cluster_id = block_id // 2
tile_id = num_clusters * w + cluster_id
bx_cluster = (tile_id // group_size) % m_clusters
bx = bx_cluster * 2 + cta_id
by = (tile_id % group_size) + (tile_id // group_size) // m_clusters * group_size

if bx * block_M < M and by * block_N < N:
for k in T.serial(k_blocks):
phase = w * k_blocks + k
T.mbarrier_wait_parity(consumed[phase % num_stages], ((phase // num_stages) & 1) ^ 1)
T.tma_copy(
A[bx * block_M : (bx + 1) * block_M, k * block_K : (k + 1) * block_K],
A_shared[phase % num_stages, :, :],
barrier=loaded[phase % num_stages],
)

T.tma_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[phase % num_stages, :, :],
barrier=loaded[phase % num_stages],
)
T.mbarrier_arrive(loaded[phase % num_stages], 0)

elif tx < 64 and cta_id == 0: # warp 1: issue tcgen5
for w in T.unroll(waves):
# manual threadblock swizzle
cluster_id = block_id // 2
tile_id = num_clusters * w + cluster_id
bx_cluster = (tile_id // group_size) % m_clusters
bx = bx_cluster * 2 + cta_id
by = (tile_id % group_size) + (tile_id // group_size) // m_clusters * group_size

if bx * block_M < M and by * block_N < N:
T.mbarrier_wait_parity(tmem_empty[w & 1], ((w // 2) & 1) ^ 1)
for k in T.serial(k_blocks):
phase = w * k_blocks + k
T.mbarrier_wait_parity(loaded[phase % num_stages], (phase // num_stages) & 1)
if w & 1 == 0:
T.tcgen05_gemm(
A_shared[phase % num_stages, :, :],
B_shared[phase % num_stages, :, :],
C_tmem_0,
mbar=consumed[phase % num_stages],
clear_accum=k == 0,
use_2cta=True,
)
else:
T.tcgen05_gemm(
A_shared[phase % num_stages, :, :],
B_shared[phase % num_stages, :, :],
C_tmem_1,
mbar=consumed[phase % num_stages],
clear_accum=k == 0,
use_2cta=True,
)
T.tcgen05_mma_arrive(tmem_full[w & 1], arrive_2cta=True)

elif 128 <= tx < 256: # warp 4~7: epilogue
for w in T.unroll(waves):
# manual threadblock swizzle
cluster_id = block_id // 2
tile_id = num_clusters * w + cluster_id
bx_cluster = (tile_id // group_size) % m_clusters
bx = bx_cluster * 2 + cta_id
by = (tile_id % group_size) + (tile_id // group_size) // m_clusters * group_size

if bx * block_M < M and by * block_N < N:
T.mbarrier_wait_parity(tmem_full[w & 1], (w // 2) & 1)
T.sync_threads(1, 128)
if (w & 1) == 0:
T.copy(C_tmem_0, C_local)
else:
T.copy(C_tmem_1, C_local)
T.mbarrier_arrive(tmem_empty[w & 1], 0)

if use_tma_store:
for i in T.unroll(T.ceildiv(block_N, store_block_N)):
T.copy(C_local[:, i * store_block_N : (i + 1) * store_block_N], C_shared)
T.copy(C_shared, C[bx * block_M, by * block_N + i * store_block_N])
else:
T.copy(C_local, C_local_cast)
T.copy(C_local_cast, C[bx * block_M, by * block_N])

return C


def main():
M, N, K = 8192, 8192, 8192
block_M, block_N, block_K = 128, 256, 64
store_block_N = 128
store_block_N = 64
in_dtype, out_dtype, accum_dtype = T.bfloat16, T.bfloat16, T.float
num_stages = 4
enable_2cta_tcgen5mma = True
num_stages = 6 if enable_2cta_tcgen5mma else 4 # Each cta only needs to load half of B, enabling larger stages
kernel = gemm_persistent_2cta if enable_2cta_tcgen5mma else gemm_persistent

a = torch.randn(M, K, device="cuda", dtype=torch.bfloat16)
b = torch.randn(K, N, device="cuda", dtype=torch.bfloat16)
print(gemm.get_kernel_source(a, b, block_M, block_N, store_block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages))
c = gemm(a, b, block_M, block_N, store_block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages)
print(kernel.get_kernel_source(a, b, block_M, block_N, store_block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages))
c = kernel(a, b, block_M, block_N, store_block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages)

ref_c = (a.to(torch.float) @ b.to(torch.float)).to(torch.bfloat16)
torch.testing.assert_close(c, ref_c, rtol=1e-2, atol=1e-2)
print("All checks passed. ✅")

tl_latency = do_bench(
lambda: gemm(a, b, block_M, block_N, store_block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages), backend="cupti"
lambda: kernel(a, b, block_M, block_N, store_block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages), backend="cupti"
)
torch_latency = do_bench(lambda: a @ b, backend="cupti")
print(f"Tilelang latency: {tl_latency} ms")
Expand Down
2 changes: 1 addition & 1 deletion maint/gemm_v2/correctness_evaluation_tcgen05.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,7 @@ def main(
for k in T.Pipelined(T.ceildiv(K, block_K), num_stages=num_stages):
T.copy(A[by * block_M, k * block_K], A_shared)
T.copy(B[bx * block_N, k * block_K], B_shared)
T.gemm(A_shared, B_shared, C_tmem, trans_A, trans_B, mbar=mbar, wg_wait=-1, clear_accum=k == 0)
T.tcgen05_gemm(A_shared, B_shared, C_tmem, trans_A, trans_B, mbar=mbar, clear_accum=k == 0)
T.mbarrier_wait_parity(mbar, k % 2)

T.copy(C_tmem, C_local)
Expand Down
Loading
Loading