Skip to content
Merged
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
general support for 2cta in GetTCGEN5MMAMeta
  • Loading branch information
Rachmanino committed Mar 18, 2026
commit dbc93885f7fc3691470eebfdc6b8510090ae3f97
46 changes: 34 additions & 12 deletions src/op/tcgen5_meta.h
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@ GetTCGEN5MMAMeta(int M, int N, int K, DataType ab_dtype, DataType c_dtype,
// TODO (lei) Currently not all shapes / dtypes are supported for TCGEN5MMA.
// NOTE: In case that other tcgen5mma in the same kernel must use 1-cta,
// we should disable 2cta for the current tcgen5mma.
// TODO(wt): Add more 2cta-preferred shapes
#define FAIL \
return { \
false, TCGEN5MMAMeta { 0, 0, 0, false, false } \
Expand All @@ -37,14 +38,19 @@ GetTCGEN5MMAMeta(int M, int N, int K, DataType ab_dtype, DataType c_dtype,
(c_dtype.is_float() && c_dtype.bits() == 32)) {
if (K % 16 != 0)
FAIL;
if (M == 128 && !disable_2cta) {
// todo(wt): Add more 2cta-preferred shapes
for (int atom_n = 256; atom_n >= 16; atom_n -= 16)
if (N % atom_n == 0)
SUCCESS(256, atom_n, 16, false, true);
if (!disable_2cta) {
if (M % 128 == 0) {
for (int atom_n = 256; atom_n >= 16; atom_n -= 16)
if (N % atom_n == 0)
SUCCESS(256, atom_n, 16, false, true);
} else if (M % 64 == 0) {
for (int atom_n = 256; atom_n >= 16; atom_n -= 16)
if (N % atom_n == 0)
SUCCESS(128, atom_n, 16, false, true);
}
}
if (M % 128 == 0) {
for (int atom_n = 256; atom_n >= 16; atom_n -= 16)
for (int atom_n = 256; atom_n >= 8; atom_n -= 8)
if (N % atom_n == 0)
SUCCESS(128, atom_n, 16, false, false);
FAIL;
Expand All @@ -67,13 +73,21 @@ GetTCGEN5MMAMeta(int M, int N, int K, DataType ab_dtype, DataType c_dtype,
(c_dtype.is_float16() && c_dtype.bits() == 16))) {
if (K % 32 != 0)
FAIL;
if (!disable_2cta) {
if (M % 128 == 0) {
for (int atom_n = 256; atom_n >= 16; atom_n -= 16)
if (N % atom_n == 0)
SUCCESS(256, atom_n, 32, false, true);
} else if (M % 64 == 0) {
for (int atom_n = 256; atom_n >= 16; atom_n -= 16)
if (N % atom_n == 0)
SUCCESS(128, atom_n, 32, false, true);
}
}
if (M % 128 == 0) {
for (int atom_n : ws_valid_atom_ns)
if (N % atom_n == 0)
SUCCESS(128, atom_n, 32, true, false);
for (int atom_n = 256; atom_n >= 16; atom_n -= 16)
if (N % atom_n == 0)
SUCCESS(128, atom_n, 32, false, true);
for (int atom_n = 256; atom_n >= 8; atom_n -= 8)
if (N % atom_n == 0)
SUCCESS(128, atom_n, 32, false, false);
Expand All @@ -98,13 +112,21 @@ GetTCGEN5MMAMeta(int M, int N, int K, DataType ab_dtype, DataType c_dtype,
ab_dtype.bits() == 8 && c_dtype.is_int() && c_dtype.bits() == 32) {
if (K % 32 != 0)
FAIL;
if (!disable_2cta) {
if (M % 128 == 0) {
for (int atom_n = 256; atom_n >= 32; atom_n -= 32)
if (N % atom_n == 0)
SUCCESS(256, atom_n, 32, false, true);
} else if (M % 64 == 0) {
for (int atom_n = 256; atom_n >= 32; atom_n -= 32)
if (N % atom_n == 0)
SUCCESS(128, atom_n, 32, false, true);
}
}
if (M % 128 == 0) {
for (int atom_n : ws_valid_atom_ns)
if (N % atom_n == 0)
SUCCESS(128, atom_n, 32, true, false);
for (int atom_n = 256; atom_n >= 32; atom_n -= 32)
if (N % atom_n == 0)
SUCCESS(128, atom_n, 32, false, true);
for (int atom_n = 256; atom_n >= 8; atom_n -= (atom_n > 32 ? 16 : 8))
// steps of 16 after N > 32
if (N % atom_n == 0)
Expand Down