submission 490544
jiab_85281 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1069 lines, June 9 Researcher Reciprocity License v1.0.
tmp.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-490544?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:1ed2f2dcd5f61443bdbd357ea612e234d2ca70f595a559948b820ecf364fd3a6
license declaredunknown
license concludedunknown
authorsjiab_85281
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
__global__ __cluster_dims__(2, 1, 1) __launch_bounds__(TB_SIZE)fused-epilogue
constexpr int EP_STRIDE = 136; // padded stride for chunk32 epiloguembarrier
__device__ inline void mbarrier_init(int mbar_addr, int count) {shared-memory
int smem_size_bytes;tcgen05
asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(CTA_GROUP));tile-k = 256
constexpr int BLOCK_K = 256;tile-m = 128
constexpr int BLOCK_M = 128;tile-n = 128
constexpr int BLOCK_N = 128;tma
CUtensorMap A_full[MAX_GROUPS]; // orig B data (MMA A operand), box_h=128Kernel source
tmp.py1069 lines
#!POPCORN gpu NVIDIA
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
# CTA-copy no-ep variant aligned to cta_full control flow.
cuda_src = """
#include <cudaTypedefs.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <torch/library.h>
#include <ATen/core/Tensor.h>
#include <cstdint>
// ============================================================================
// Constants
// ============================================================================
constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;
constexpr uint64_t EVICT_NORMAL = 0x1000000000000000ULL;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000ULL;
constexpr int BLOCK_M = 128;
constexpr int BLOCK_N = 128;
constexpr int BLOCK_K = 256;
constexpr int A_SIZE = BLOCK_M * BLOCK_K / 2; // 16384
constexpr int B_SIZE = BLOCK_N * BLOCK_K / 2; // 16384
constexpr int B_HALF = BLOCK_N * BLOCK_K / 4; // 8192 (per CTA in cta_group::2)
constexpr int SFA_SIZE = 128 * BLOCK_K / 16; // 2048
constexpr int SFB_SIZE = 128 * BLOCK_K / 16; // 2048
constexpr int SF_STAGE = SFA_SIZE + SFB_SIZE; // 4096
constexpr int EP_STRIDE = 136; // padded stride for chunk32 epilogue
constexpr int EP_SMEM_BYTES = 32 * EP_STRIDE * (int)sizeof(half); // 8704B scratch
constexpr int NUM_EP_WARPS = 4;
constexpr int NUM_WARPS = NUM_EP_WARPS + 2; // 6
constexpr int TB_SIZE = NUM_WARPS * WARP_SIZE; // 192
constexpr int TMEM_COLS = 512;
constexpr int MAX_LAUNCH_CTAS = 148;
constexpr int MAX_CLUSTERS = MAX_LAUNCH_CTAS / 2;
constexpr int MAX_TILE_LUT = 1024;
constexpr int D_TMEM0 = 0;
constexpr int D_TMEM1 = BLOCK_N; // 128
constexpr int SFA_TMEM = 2 * BLOCK_N; // 256
constexpr int SFB_TMEM = SFA_TMEM + 4 * (BLOCK_K / MMA_K); // 272
constexpr int MAX_GROUPS = 8;
constexpr int MAX_NS = 10; // maximum pipeline stages
// ============================================================================
// Device structures
// ============================================================================
struct GroupInfo {
half* c_ptr;
int M, N, K;
};
// Fat LUT: all per-tile info pre-computed host-side. SoA layout with uint16 arrays.
// No fused-M: each LUT entry is a single (coord_n, coord_m) tile.
struct KernelParams {
GroupInfo groups[MAX_GROUPS];
int num_groups;
int total_tiles;
int launch_ctas;
int smem_size_bytes;
int ns; // pipeline stages (smem-driven; independent of num_k divisibility)
int main_stg; // max stage size for smem budgeting / NS selection
int32_t lut_worker_start[MAX_CLUSTERS];
int32_t lut_worker_count[MAX_CLUSTERS];
// Fat LUT arrays — pre-computed per tile.
uint16_t lut_gidx[MAX_TILE_LUT]; // group index
uint16_t lut_coord_m[MAX_TILE_LUT]; // M-tile coord (for SFB tmap)
uint16_t lut_off_n_cta_r0[MAX_TILE_LUT]; // N offset for cta_rank=0
uint16_t lut_off_n_cta_r1[MAX_TILE_LUT]; // N offset for cta_rank=1
uint16_t lut_coord_n_cta_r0[MAX_TILE_LUT]; // N-tile coord for cta_rank=0 (SFA tmap)
uint16_t lut_coord_n_cta_r1[MAX_TILE_LUT]; // N-tile coord for cta_rank=1 (SFA tmap)
uint16_t lut_off_m_b_r0[MAX_TILE_LUT]; // B operand y-offset for cta_rank=0
uint16_t lut_off_m_b_r1[MAX_TILE_LUT]; // B operand y-offset for cta_rank=1
uint16_t lut_expect_bytes_r0[MAX_TILE_LUT]; // TMA expect_tx total bytes for cta_rank=0
uint16_t lut_expect_bytes_r1[MAX_TILE_LUT]; // TMA expect_tx total bytes for cta_rank=1
uint16_t lut_mma_n_num_k[MAX_TILE_LUT]; // [7:0]=mma_n, [15:8]=num_k (K/BLOCK_K)
// Packed tmap selection: bits [1:0]=a_tmap_r0, [3:2]=a_tmap_r1, [5:4]=b_tmap_r0, [7:6]=b_tmap_r1
uint8_t lut_tmap_sel[MAX_TILE_LUT];
};
enum : int {
SCHED_BASE = 0,
SCHED_REV = 1,
SCHED_F2_G2 = 2,
SCHED_F1_G1 = 3,
};
enum : int {
PROFILE_GENERIC_G8 = 0,
PROFILE_GENERIC_G2 = 1,
PROFILE_BENCH1 = 2,
PROFILE_BENCH2 = 3,
PROFILE_BENCH3 = 4,
PROFILE_BENCH4 = 5,
};
constexpr int V_FORCE_SCHEDULE = -1;
constexpr int V_BENCH1_CLUSTERS = -1;
constexpr int V_BENCH2_CLUSTERS = -1;
constexpr int V_GENERIC_CLUSTERS = -1;
struct TmapParamPackG8 {
CUtensorMap A_full[MAX_GROUPS]; // orig B data (MMA A operand), box_h=128
CUtensorMap A_tail[MAX_GROUPS]; // orig B data, N-tail
CUtensorMap B_full[MAX_GROUPS]; // orig A data (MMA B operand), box_h=64 per CTA
CUtensorMap B_tail0[MAX_GROUPS]; // orig A data, M-tail for CTA0
CUtensorMap B_tail1[MAX_GROUPS]; // orig A data, M-tail for CTA1
CUtensorMap SFA[MAX_GROUPS]; // orig SFB (scale for MMA A)
CUtensorMap SFB[MAX_GROUPS]; // orig SFA (scale for MMA B)
};
// ============================================================================
// Inline PTX helpers
// ============================================================================
__device__ inline constexpr uint64_t desc_encode(uint64_t x) {
return (x & 0x3'FFFFULL) >> 4ULL;
}
__device__ inline uint32_t elect_sync() {
uint32_t pred = 0;
asm volatile(
"{\\n\\t"
".reg .pred %%px;\\n\\t"
"elect.sync _|%%px, %1;\\n\\t"
"@%%px mov.s32 %0, 1;\\n\\t"
"}"
: "+r"(pred) : "r"(0xFFFFFFFF));
return pred;
}
__device__ inline uint32_t get_cluster_ctarank() {
uint32_t rank = 0;
asm volatile("mov.u32 %0, %%cluster_ctarank;" : "=r"(rank));
return rank;
}
__device__ inline void mbarrier_init(int mbar_addr, int count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
}
__device__ void mbarrier_wait(int mbar_addr, int phase) {
uint32_t ticks = 0x10000;
asm volatile(
"{\\n\\t"
".reg .pred P1;\\n\\t"
"LAB_WAIT:\\n\\t"
"mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\\n\\t"
"@!P1 bra.uni LAB_WAIT;\\n\\t"
"}"
:: "r"(mbar_addr), "r"(phase), "r"(ticks));
}
__device__ inline void mbarrier_arrive_expect_tx(int mbar_addr, int size) {
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
:: "r"(mbar_addr), "r"(size) : "memory");
}
template <int CTA_GROUP = 2>
__device__ inline void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {
asm volatile("cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%7.L2::cache_hint "
"[%0], [%1, {%2, %3, %4}], [%5], %6;"
:: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy), "n"(CTA_GROUP)
: "memory");
}
template <int CTA_GROUP = 2>
__device__ inline void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(CTA_GROUP));
}
template <int CTA_GROUP = 2>
__device__ inline void tcgen05_mma_nvfp4(int d_tmem, uint64_t a_desc, uint64_t b_desc, uint32_t i_desc,
int scale_A_tmem, int scale_B_tmem, int enable_input_d) {
asm volatile(
"{\\n\\t"
".reg .pred p;\\n\\t"
"setp.ne.b32 p, %6, 0;\\n\\t"
"tcgen05.mma.cta_group::%7.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse [%0], %1, %2, %3, [%4], [%5], p;\\n\\t"
"}"
:: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
"r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d), "n"(CTA_GROUP));
}
template <int CTA_GROUP = 2>
__device__ inline void tcgen05_commit(int mbar_addr) {
asm volatile("tcgen05.commit.cta_group::%1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mbar_addr), "n"(CTA_GROUP) : "memory");
}
template <int CTA_GROUP = 2>
__device__ inline void tcgen05_commit_mcast(int mbar_addr, uint16_t cta_mask) {
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:: "r"(mbar_addr), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
}
__device__ inline void tcgen05_ld_16x256bx2(float *tmp, int row, int col) {
asm volatile("tcgen05.ld.sync.aligned.16x256b.x2.b32 "
"{ %0, %1, %2, %3, %4, %5, %6, %7 }, [%8];"
: "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]),
"=f"(tmp[4]), "=f"(tmp[5]), "=f"(tmp[6]), "=f"(tmp[7])
: "r"((row << 16) | col));
}
// Chunk32 transposed epilogue (copied from cg2_full style).
__device__ inline void do_epilogue_transposed_chunk32(
int warp_id, int lane_id, int cta_rank,
int done_mbar, int done_phase, int d_tmem_base,
half* __restrict__ smem_ep, half* __restrict__ c_ptr,
int M, int N, int off_m, int off_n, int mma_n
) {
mbarrier_wait(done_mbar, done_phase);
asm volatile("tcgen05.fence::after_thread_sync;");
const int col_lane = (lane_id % 4) * 2;
const int row_lane = lane_id / 4;
const int tid_ep = warp_id * WARP_SIZE + lane_id;
const int residue_n = N - off_n;
const int num_chunks = (mma_n + 31) / 32;
#pragma unroll 1
for (int chunk = 0; chunk < num_chunks; chunk++) {
#pragma unroll
for (int mc = 0; mc < 2; mc++) {
#pragma unroll
for (int m = 0; m < 2; m++) {
const int tm = cta_rank * BLOCK_M + warp_id * 32 + m * 16;
float vals[8];
tcgen05_ld_16x256bx2(vals, tm, d_tmem_base + chunk * 32 + mc * 16);
asm volatile("tcgen05.wait::ld.sync.aligned;");
const int n0 = warp_id * 32 + m * 16 + row_lane;
const int n1 = n0 + 8;
const int m_off = mc * 16;
smem_ep[(col_lane + m_off) * EP_STRIDE + n0] = __float2half_rn(vals[0]);
smem_ep[(col_lane + 1 + m_off) * EP_STRIDE + n0] = __float2half_rn(vals[1]);
smem_ep[(col_lane + m_off) * EP_STRIDE + n1] = __float2half_rn(vals[2]);
smem_ep[(col_lane + 1 + m_off) * EP_STRIDE + n1] = __float2half_rn(vals[3]);
smem_ep[(col_lane + 8 + m_off) * EP_STRIDE + n0] = __float2half_rn(vals[4]);
smem_ep[(col_lane + 9 + m_off) * EP_STRIDE + n0] = __float2half_rn(vals[5]);
smem_ep[(col_lane + 8 + m_off) * EP_STRIDE + n1] = __float2half_rn(vals[6]);
smem_ep[(col_lane + 9 + m_off) * EP_STRIDE + n1] = __float2half_rn(vals[7]);
}
}
asm volatile("bar.sync 15, %0;" :: "r"(NUM_EP_WARPS * WARP_SIZE));
const int n_group = tid_ep % 8;
const int n_start = n_group * 16;
#pragma unroll
for (int pass = 0; pass < 2; pass++) {
const int m_row = pass * 16 + tid_ep / 8;
const int m_local = chunk * 32 + m_row;
const int m_global = off_m + m_local;
if (m_local < mma_n && m_global < M && n_start < residue_n) {
half* dst_base = &c_ptr[m_global * N + off_n + n_start];
const half* src_base = &smem_ep[m_row * EP_STRIDE + n_start];
if (n_start + 16 <= residue_n) {
*reinterpret_cast<int4*>(dst_base) = *reinterpret_cast<const int4*>(src_base);
*reinterpret_cast<int4*>(dst_base + 8) = *reinterpret_cast<const int4*>(src_base + 8);
} else {
#pragma unroll
for (int i = 0; i < 16; i++) {
if (n_start + i < residue_n) dst_base[i] = src_base[i];
}
}
}
}
asm volatile("bar.sync 15, %0;" :: "r"(NUM_EP_WARPS * WARP_SIZE));
}
}
__device__ inline void do_epilogue_transposed(
int warp_id, int lane_id, int cta_rank,
int done_mbar, int done_phase, int d_tmem_base,
half* __restrict__ smem_ep, half* __restrict__ c_ptr,
int M, int N, int off_m, int off_n, int mma_n
) {
do_epilogue_transposed_chunk32(
warp_id, lane_id, cta_rank, done_mbar, done_phase, d_tmem_base,
smem_ep, c_ptr, M, N, off_m, off_n, mma_n);
}
// ============================================================================
// TensorMap Initialization
// ============================================================================
void check_cu(CUresult err) {
if (err == CUDA_SUCCESS) return;
const char *msg;
if (cuGetErrorString(err, &msg) != CUDA_SUCCESS) msg = "unknown";
TORCH_CHECK(false, "cuTensorMapEncodeTiled error: ", msg);
}
void init_AB_tmap(CUtensorMap *tmap, const char *ptr, uint64_t height, uint64_t width,
uint32_t box_h, uint32_t box_w, CUtensorMapL2promotion l2_promotion) {
constexpr uint32_t rank = 3;
uint64_t globalDim[rank] = {256, height, width / 256};
uint64_t globalStrides[rank - 1] = {width / 2, 128};
uint32_t boxDim[rank] = {256, box_h, box_w / 256};
uint32_t elementStrides[rank] = {1, 1, 1};
check_cu(cuTensorMapEncodeTiled(tmap, CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B, rank, (void *)ptr,
globalDim, globalStrides, boxDim, elementStrides,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
l2_promotion, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE));
}
void init_SF_tmap(CUtensorMap *tmap, const char *ptr, uint64_t mn, uint64_t K,
CUtensorMapL2promotion l2_promotion) {
constexpr uint32_t rank = 3;
const uint64_t k_blocks = K / 64;
const uint64_t mn_blocks = (mn + 127) / 128;
const uint32_t tile_k_blocks = BLOCK_K / 64;
constexpr uint64_t SF_BLOCK_BYTES = 512;
constexpr uint64_t X_ELEMS = SF_BLOCK_BYTES / sizeof(uint16_t);
uint64_t globalDim[rank] = {X_ELEMS, mn_blocks, k_blocks};
uint64_t globalStrides[rank-1] = {k_blocks * SF_BLOCK_BYTES, SF_BLOCK_BYTES};
uint32_t boxDim[rank] = {(uint32_t)X_ELEMS, 1, tile_k_blocks};
uint32_t elementStrides[rank] = {1, 1, 1};
check_cu(cuTensorMapEncodeTiled(tmap, CU_TENSOR_MAP_DATA_TYPE_UINT16, rank, (void *)ptr,
globalDim, globalStrides, boxDim, elementStrides,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_NONE,
l2_promotion, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE));
}
// ============================================================================
// Host-side fat LUT builder
// ============================================================================
struct TileLutTmp {
uint16_t gidx;
uint16_t coord_m;
uint16_t off_n_cta_r0, off_n_cta_r1;
uint16_t coord_n_cta_r0, coord_n_cta_r1;
uint16_t off_m_b_r0, off_m_b_r1;
uint16_t expect_bytes_r0, expect_bytes_r1;
uint16_t mma_n_num_k; // [7:0]=mma_n, [15:8]=num_k
uint8_t tmap_sel; // packed tmap indices
};
inline void build_tile_lut(KernelParams& params, int schedule) {
TileLutTmp tmp[MAX_TILE_LUT];
int total = 0;
for (int g = 0; g < params.num_groups; g++) {
const GroupInfo& gi = params.groups[g];
const int M = gi.M, N = gi.N, K = gi.K;
const int n_tiles = (N + 255) / 256;
const int m_tiles = (M + BLOCK_N - 1) / BLOCK_N;
const int m_rem = M % BLOCK_N;
const int mma_n_tail = (m_rem == 0) ? BLOCK_N : ((m_rem + 15) & ~15);
for (int cn = 0; cn < n_tiles; cn++) {
for (int cm = 0; cm < m_tiles; cm++) {
TORCH_CHECK(total < MAX_TILE_LUT, "tile LUT overflow");
const bool is_m_tail = (cm == m_tiles - 1) && (mma_n_tail != BLOCK_N);
const int mma_n = is_m_tail ? mma_n_tail : BLOCK_N;
const int b_half_rows = mma_n / 2;
const int off_m_base = cm * BLOCK_N;
// Per cta_rank: compute N-axis coords
int off_n_r0 = cn * 256;
int off_n_r1 = cn * 256 + BLOCK_M;
int coord_n_cta_r0 = cn * 2;
int coord_n_cta_r1 = cn * 2 + 1;
// Handle N-tail: if cta_rank=1 would be out of bounds, alias to rank=0
if (off_n_r1 >= N) {
off_n_r1 = off_n_r0;
coord_n_cta_r1 = coord_n_cta_r0;
}
// A tmap selection (based on N residue for each rank)
int n_residue_r0 = N - off_n_r0;
int n_residue_r1 = N - off_n_r1;
bool is_a_tail_r0 = (n_residue_r0 > 0 && n_residue_r0 < BLOCK_M);
bool is_a_tail_r1 = (n_residue_r1 > 0 && n_residue_r1 < BLOCK_M);
// Both ranks in a cluster see same coord_n, but different cta offsets.
// The A tmap idx is actually the same for both ranks within same coord_n
// because A_full vs A_tail depends on the per-CTA N residue.
// We store per-rank since they can differ.
uint16_t a_tmap_r0 = is_a_tail_r0 ? 1 : 0;
uint16_t a_tmap_r1 = is_a_tail_r1 ? 1 : 0;
// For simplicity, store worst case (if either is tail, both get tail idx)
// Actually no - each CTA independently selects its A tmap. Store per-rank.
// But our LUT only has one a_tmap_idx field. Let's use the cta_rank to select.
// Actually, a_tail only matters at the N boundary. For cn < n_tiles-1, both are full.
// For cn == n_tiles-1, rank0 might be tail, rank1 might be OOB (aliased to rank0).
// So if rank1 is aliased to rank0, they share the same tail status.
// Let's just store per-rank.
int a_bytes_r0 = is_a_tail_r0 ? (n_residue_r0 * BLOCK_K / 2) : A_SIZE;
int a_bytes_r1 = is_a_tail_r1 ? (n_residue_r1 * BLOCK_K / 2) : A_SIZE;
// B operand offsets per rank
int off_m_b_r0 = off_m_base;
int off_m_b_r1 = off_m_base + b_half_rows;
int b_rows_r0 = b_half_rows;
int b_rows_r1 = b_half_rows;
// B tmap selection
uint16_t b_tmap_r0 = 0; // B_full
uint16_t b_tmap_r1 = 0; // B_full
if (is_m_tail) {
const int m_residue = M - off_m_base;
b_rows_r0 = (m_residue < b_half_rows) ? m_residue : b_half_rows;
b_rows_r1 = m_residue - b_half_rows;
if (b_rows_r1 < 0) b_rows_r1 = 0;
if (b_rows_r1 > b_half_rows) b_rows_r1 = b_half_rows;
b_tmap_r0 = 1; // B_tail0
b_tmap_r1 = 2; // B_tail1
}
if (b_rows_r0 < 1) b_rows_r0 = 1;
if (b_rows_r1 < 1) {
b_rows_r1 = 1;
off_m_b_r1 = off_m_base; // safe fallback
}
int b_bytes_r0 = b_rows_r0 * BLOCK_K / 2;
int b_bytes_r1 = b_rows_r1 * BLOCK_K / 2;
int expect_r0 = a_bytes_r0 + b_bytes_r0 + SFA_SIZE + SFB_SIZE;
int expect_r1 = a_bytes_r1 + b_bytes_r1 + SFA_SIZE + SFB_SIZE;
// Pack tmap indices: [1:0]=a_r0, [3:2]=a_r1, [5:4]=b_r0, [7:6]=b_r1
uint8_t tmap_sel = (a_tmap_r0 & 3) | ((a_tmap_r1 & 3) << 2)
| ((b_tmap_r0 & 3) << 4) | ((b_tmap_r1 & 3) << 6);
tmp[total++] = {
(uint16_t)g,
(uint16_t)cm,
(uint16_t)off_n_r0, (uint16_t)off_n_r1,
(uint16_t)coord_n_cta_r0, (uint16_t)coord_n_cta_r1,
(uint16_t)off_m_b_r0, (uint16_t)off_m_b_r1,
(uint16_t)expect_r0, (uint16_t)expect_r1,
(uint16_t)((mma_n & 0xFF) | ((K / BLOCK_K) << 8)),
tmap_sel,
};
}
}
}
params.total_tiles = total;
TORCH_CHECK(params.total_tiles <= MAX_TILE_LUT, "total_tiles exceeds LUT capacity");
const int num_clusters = params.launch_ctas / 2;
TORCH_CHECK(num_clusters > 0 && num_clusters <= MAX_CLUSTERS, "num_clusters out of range");
int cursor = 0;
for (int cid = 0; cid < num_clusters; cid++) {
int logical_cid = cid;
if (schedule == SCHED_F2_G2) logical_cid = (cid * 17) % num_clusters;
else if (schedule == SCHED_F1_G1) logical_cid = num_clusters - 1 - cid;
const int my_count = (params.total_tiles - logical_cid + num_clusters - 1) / num_clusters;
params.lut_worker_start[cid] = cursor;
params.lut_worker_count[cid] = my_count;
for (int tile_iter = 0; tile_iter < my_count; tile_iter++) {
int k = tile_iter;
if (schedule == SCHED_REV || schedule == SCHED_F1_G1) {
k = my_count - 1 - tile_iter;
} else if (schedule == SCHED_F2_G2) {
const int h = (my_count + 1) >> 1;
k = (tile_iter < h) ? (tile_iter << 1) : (((tile_iter - h) << 1) + 1);
}
const int tile_id = logical_cid + k * num_clusters;
TORCH_CHECK(tile_id >= 0 && tile_id < params.total_tiles, "tile_id out of range");
const TileLutTmp& t = tmp[tile_id];
params.lut_gidx[cursor] = t.gidx;
params.lut_coord_m[cursor] = t.coord_m;
params.lut_off_n_cta_r0[cursor] = t.off_n_cta_r0;
params.lut_off_n_cta_r1[cursor] = t.off_n_cta_r1;
params.lut_coord_n_cta_r0[cursor] = t.coord_n_cta_r0;
params.lut_coord_n_cta_r1[cursor] = t.coord_n_cta_r1;
params.lut_off_m_b_r0[cursor] = t.off_m_b_r0;
params.lut_off_m_b_r1[cursor] = t.off_m_b_r1;
params.lut_expect_bytes_r0[cursor] = t.expect_bytes_r0;
params.lut_expect_bytes_r1[cursor] = t.expect_bytes_r1;
params.lut_mma_n_num_k[cursor] = t.mma_n_num_k;
params.lut_tmap_sel[cursor] = t.tmap_sel;
cursor++;
}
}
for (int cid = num_clusters; cid < MAX_CLUSTERS; cid++) {
params.lut_worker_start[cid] = 0;
params.lut_worker_count[cid] = 0;
}
TORCH_CHECK(cursor == params.total_tiles, "tile LUT size mismatch");
}
// ============================================================================
// Kernel — persistent mbars, fat LUT, NS as template param
// ============================================================================
template <int SCHEDULE_ID, int PROFILE_ID>
__global__ __cluster_dims__(2, 1, 1) __launch_bounds__(TB_SIZE)
void grouped_gemm_kernel(
const __grid_constant__ KernelParams params,
const __grid_constant__ TmapParamPackG8 tmap_pack_g8
) {
const int tid = threadIdx.x;
const int warp_id = tid / WARP_SIZE;
const int lane_id = tid % WARP_SIZE;
const int bid = blockIdx.x;
if (bid >= params.launch_ctas) return;
const int cta_rank = static_cast<int>(get_cluster_ctarank());
const int cluster_id = bid / 2;
const int my_count = params.lut_worker_count[cluster_id];
const int worker_start = params.lut_worker_start[cluster_id];
extern __shared__ __align__(1024) char smem_raw[];
const int smem = static_cast<int>(__cvta_generic_to_shared(smem_raw));
const int smem_main = smem;
half* smem_ep = reinterpret_cast<half*>(smem_raw + params.smem_size_bytes - EP_SMEM_BYTES);
const int NS = params.ns;
const int main_stg = params.main_stg;
const int smem_sf = smem_main + NS * main_stg;
// Mbarrier layout:
// - NS tma + NS mma
// - 2 done mbars (TMEM slot ready for EP)
// - 2 ep mbars (EP drained slot; TMEM backpressure)
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ int64_t mbars[2 * MAX_NS + 4];
__shared__ int32_t tmem_alloc_buf;
const int mbar_base = static_cast<int>(__cvta_generic_to_shared(mbars));
const int tma_mbar = mbar_base;
const int mma_mbar = tma_mbar + NS * 8;
const int done_mbar0 = mma_mbar + NS * 8;
const int done_mbar1 = done_mbar0 + 8;
const int ep_mbar0 = done_mbar1 + 8;
const int ep_mbar1 = ep_mbar0 + 8;
if (my_count <= 0) return;
// === ONE-TIME SETUP ===
// MMA warp: allocate TMEM (both CTAs must issue)
if (warp_id == NUM_WARPS - 1) {
int alloc_addr = static_cast<int>(__cvta_generic_to_shared(&tmem_alloc_buf));
asm volatile("tcgen05.alloc.cta_group::2.sync.aligned.shared::cta.b32 [%0], %1;"
:: "r"(alloc_addr), "r"(TMEM_COLS));
}
// Warp 0: init all mbarriers ONCE for the entire kernel lifetime
if (warp_id == 0 && elect_sync()) {
for (int i = 0; i < NS; i++) {
mbarrier_init(tma_mbar + i * 8, 2); // 2 CTAs arrive
mbarrier_init(mma_mbar + i * 8, 1); // 1 MMA warp (CTA0) arrives
}
mbarrier_init(done_mbar0, 1);
mbarrier_init(done_mbar1, 1);
mbarrier_init(ep_mbar0, 2);
mbarrier_init(ep_mbar1, 2);
asm volatile("fence.mbarrier_init.release.cluster;");
}
// Single cluster barrier for the entire kernel
asm volatile("barrier.cluster.arrive.relaxed.aligned;");
asm volatile("barrier.cluster.wait.acquire.aligned;");
constexpr uint64_t cache_A =
(PROFILE_ID == PROFILE_BENCH3 || PROFILE_ID == PROFILE_BENCH4) ? EVICT_FIRST :
((PROFILE_ID == PROFILE_GENERIC_G2) ? EVICT_NORMAL : 0ULL);
constexpr uint64_t cache_B =
(PROFILE_ID == PROFILE_BENCH3 || PROFILE_ID == PROFILE_GENERIC_G2) ? EVICT_FIRST :
((PROFILE_ID == PROFILE_BENCH4) ? EVICT_NORMAL : 0ULL);
constexpr uint64_t cache_SF = cache_B;
constexpr uint16_t cta_mask = 0x3;
constexpr int SF_K_PER_BLOCK_L = BLOCK_K / 64;
constexpr uint32_t MMA_M_2CTA = BLOCK_M * 2;
auto make_desc_AB = [](int addr) -> uint64_t {
return desc_encode(addr) | (desc_encode(8 * 128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
// Helper to select A tmap based on index
auto get_a_tmap = [&](int gidx, int a_idx) -> const void* {
if (a_idx == 0) return static_cast<const void*>(&tmap_pack_g8.A_full[gidx]);
return static_cast<const void*>(&tmap_pack_g8.A_tail[gidx]);
};
// Helper to select B tmap based on index
auto get_b_tmap = [&](int gidx, int b_idx) -> const void* {
if (b_idx == 0) return static_cast<const void*>(&tmap_pack_g8.B_full[gidx]);
if (b_idx == 1) return static_cast<const void*>(&tmap_pack_g8.B_tail0[gidx]);
return static_cast<const void*>(&tmap_pack_g8.B_tail1[gidx]);
};
// === TMA WARP ===
if (warp_id == NUM_WARPS - 2 && elect_sync()) {
int tma_stage = 0;
int mma_wait_phase = 1;
int total_produced = 0;
for (int tile = 0; tile < my_count; tile++) {
const int lut_idx = worker_start + tile;
const int gidx = static_cast<int>(params.lut_gidx[lut_idx]);
const int tile_num_k = static_cast<int>(params.lut_mma_n_num_k[lut_idx] >> 8);
const int off_n_cta = cta_rank == 0
? static_cast<int>(params.lut_off_n_cta_r0[lut_idx])
: static_cast<int>(params.lut_off_n_cta_r1[lut_idx]);
const int coord_n_cta = cta_rank == 0
? static_cast<int>(params.lut_coord_n_cta_r0[lut_idx])
: static_cast<int>(params.lut_coord_n_cta_r1[lut_idx]);
const int off_m_b = cta_rank == 0
? static_cast<int>(params.lut_off_m_b_r0[lut_idx])
: static_cast<int>(params.lut_off_m_b_r1[lut_idx]);
const int coord_m = static_cast<int>(params.lut_coord_m[lut_idx]);
const int tma_expect_bytes = cta_rank == 0
? static_cast<int>(params.lut_expect_bytes_r0[lut_idx])
: static_cast<int>(params.lut_expect_bytes_r1[lut_idx]);
const uint8_t tmap_sel = params.lut_tmap_sel[lut_idx];
const int a_tmap_idx = cta_rank == 0 ? (tmap_sel & 3) : ((tmap_sel >> 2) & 3);
const int b_tmap_idx = cta_rank == 0 ? ((tmap_sel >> 4) & 3) : ((tmap_sel >> 6) & 3);
const void* A_tmap = get_a_tmap(gidx, a_tmap_idx);
const void* B_tmap = get_b_tmap(gidx, b_tmap_idx);
const void* SFA_tmap = static_cast<const void*>(&tmap_pack_g8.SFA[gidx]);
const void* SFB_tmap = static_cast<const void*>(&tmap_pack_g8.SFB[gidx]);
#pragma unroll 1
for (int ik = 0; ik < tile_num_k; ik++) {
if (total_produced >= NS) {
mbarrier_wait(mma_mbar + tma_stage * 8, mma_wait_phase);
}
const int mbar_addr = (tma_mbar + tma_stage * 8) & 0xFEFFFFFF;
int A_s = smem_main + tma_stage * main_stg;
int SFA_s = smem_sf + tma_stage * SF_STAGE;
tma_3d_gmem2smem(A_s + A_SIZE, B_tmap, 0, off_m_b, ik, mbar_addr, cache_B);
tma_3d_gmem2smem(A_s, A_tmap, 0, off_n_cta, ik, mbar_addr, cache_A);
const int z_sf = ik * SF_K_PER_BLOCK_L;
tma_3d_gmem2smem(SFA_s, SFA_tmap, 0, coord_n_cta, z_sf, mbar_addr, cache_SF);
tma_3d_gmem2smem(SFA_s + SFA_SIZE, SFB_tmap, 0, coord_m, z_sf, mbar_addr, cache_SF);
mbarrier_arrive_expect_tx(mbar_addr, tma_expect_bytes);
total_produced++;
tma_stage++;
if (tma_stage == NS) {
tma_stage = 0;
mma_wait_phase ^= 1;
}
}
}
}
// === MMA WARP (CTA0 only) ===
if (cta_rank == 0 && warp_id == NUM_WARPS - 1 && elect_sync()) {
int mma_stage = 0;
int tma_wait_phase = 0;
for (int tile = 0; tile < my_count; tile++) {
const int slot = tile & 1;
const int d_tmem_base = (slot == 0) ? D_TMEM0 : D_TMEM1;
const int done_mbar = (slot == 0) ? done_mbar0 : done_mbar1;
// Do not reuse a TMEM slot before EP drains it on both CTAs.
if (tile >= 2) {
const int ep_wait_phase = (((tile >> 1) - 1) & 1);
mbarrier_wait((slot == 0) ? ep_mbar0 : ep_mbar1, ep_wait_phase);
}
const int lut_idx = worker_start + tile;
const int mma_n = static_cast<int>(params.lut_mma_n_num_k[lut_idx] & 0xFF);
const int tile_num_k = static_cast<int>(params.lut_mma_n_num_k[lut_idx] >> 8);
const uint32_t i_desc = (1U << 7U) | (1U << 10U) |
(((uint32_t)mma_n >> 3U) << 17U) | (((uint32_t)MMA_M_2CTA >> 7U) << 27U);
#pragma unroll 1
for (int ik = 0; ik < tile_num_k; ik++) {
mbarrier_wait(tma_mbar + mma_stage * 8, tma_wait_phase);
int A_s = smem_main + mma_stage * main_stg;
int B_s = A_s + A_SIZE;
int SFA_s = smem_sf + mma_stage * SF_STAGE;
int SFB_s = SFA_s + SFA_SIZE;
constexpr uint64_t sf_base = desc_encode(0) | (desc_encode(8*16) << 32ULL) | (1ULL << 46ULL);
uint64_t sfa_desc = sf_base + ((uint64_t)SFA_s >> 4ULL);
uint64_t sfb_desc = sf_base + ((uint64_t)SFB_s >> 4ULL);
#pragma unroll
for (int kk = 0; kk < BLOCK_K / MMA_K; kk++) {
tcgen05_cp_nvfp4(SFA_TMEM + kk * 4, sfa_desc + (uint64_t)kk * (512ULL >> 4ULL));
tcgen05_cp_nvfp4(SFB_TMEM + kk * 4, sfb_desc + (uint64_t)kk * (512ULL >> 4ULL));
}
#pragma unroll
for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
uint64_t a_desc = make_desc_AB(A_s + k2 * 32);
uint64_t b_desc = make_desc_AB(B_s + k2 * 32);
// Reset accumulator at start of each tile (enable_d=0 clears accum)
int enable_d = (ik == 0 && k2 == 0) ? 0 : 1;
tcgen05_mma_nvfp4(d_tmem_base, a_desc, b_desc, i_desc,
SFA_TMEM + k2 * 4, SFB_TMEM + k2 * 4, enable_d);
}
tcgen05_commit_mcast(mma_mbar + mma_stage * 8, cta_mask);
mma_stage++;
if (mma_stage == NS) {
mma_stage = 0;
tma_wait_phase ^= 1;
}
}
tcgen05_commit_mcast(done_mbar, cta_mask);
}
}
// === EP WARPS (both CTAs) ===
if (warp_id < NUM_EP_WARPS) {
for (int tile = 0; tile < my_count; tile++) {
const int slot = tile & 1;
const int done_mbar = (slot == 0) ? done_mbar0 : done_mbar1;
const int done_phase = (tile >> 1) & 1;
const int d_tmem_base = (slot == 0) ? D_TMEM0 : D_TMEM1;
const int lut_idx = worker_start + tile;
const int gidx = static_cast<int>(params.lut_gidx[lut_idx]);
const GroupInfo& gi = params.groups[gidx];
const int M = gi.M;
const int N = gi.N;
const int off_m = static_cast<int>(params.lut_coord_m[lut_idx]) * BLOCK_N;
const int mma_n = static_cast<int>(params.lut_mma_n_num_k[lut_idx] & 0xFF);
const int off_n_r0 = static_cast<int>(params.lut_off_n_cta_r0[lut_idx]);
const int off_n_r1 = static_cast<int>(params.lut_off_n_cta_r1[lut_idx]);
const bool rank1_aliased = (off_n_r1 == off_n_r0);
const int off_n = (cta_rank == 0) ? off_n_r0 : off_n_r1;
if (!(cta_rank == 1 && rank1_aliased)) {
do_epilogue_transposed(
warp_id, lane_id, cta_rank,
done_mbar, done_phase, d_tmem_base,
smem_ep, gi.c_ptr, M, N, off_m, off_n, mma_n);
} else {
mbarrier_wait(done_mbar, done_phase);
}
if (warp_id == 0 && elect_sync()) {
tcgen05_commit_mcast((slot == 0) ? ep_mbar0 : ep_mbar1, cta_mask);
}
}
}
// === CLEANUP ===
asm volatile("barrier.cluster.arrive.relaxed.aligned;");
asm volatile("barrier.cluster.wait.acquire.aligned;");
if (warp_id == 0)
asm volatile("tcgen05.dealloc.cta_group::2.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));
}
template <int SCHEDULE_ID, int PROFILE_ID>
inline void launch_grouped_kernel(const KernelParams& params, const TmapParamPackG8& tmap_pack_g8, int smem_size) {
auto kernel = grouped_gemm_kernel<SCHEDULE_ID, PROFILE_ID>;
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, cudaSharedmemCarveoutMaxShared);
cudaFuncSetAttribute(kernel, cudaFuncAttributeNonPortableClusterSizeAllowed, 1);
cudaLaunchConfig_t launch_config = {};
launch_config.gridDim = params.launch_ctas;
launch_config.blockDim = TB_SIZE;
launch_config.dynamicSmemBytes = smem_size;
cudaLaunchAttribute cluster_attr = {};
cluster_attr.id = cudaLaunchAttributeClusterDimension;
cluster_attr.val.clusterDim.x = 2;
cluster_attr.val.clusterDim.y = 1;
cluster_attr.val.clusterDim.z = 1;
launch_config.attrs = &cluster_attr;
launch_config.numAttrs = 1;
cudaLaunchKernelEx(&launch_config, kernel, params, tmap_pack_g8);
}
// ============================================================================
// Host launch
// ============================================================================
void grouped_gemm_impl(
at::TensorList A_list,
at::TensorList B_list,
at::TensorList C_list,
at::TensorList SFA_list,
at::TensorList SFB_list
) {
int G = A_list.size();
TORCH_CHECK(G <= MAX_GROUPS, "num groups exceeds MAX_GROUPS");
if (G == 0) return;
KernelParams params = {};
params.num_groups = G;
static int smem_size = 0;
static int smem_avail = 0;
if (!smem_size) {
int dev; cudaGetDevice(&dev);
int smem_max;
cudaDeviceGetAttribute(&smem_max, cudaDevAttrMaxSharedMemoryPerBlockOptin, dev);
smem_size = smem_max - 1024;
smem_avail = smem_size - EP_SMEM_BYTES;
TORCH_CHECK(smem_avail > 0, "Insufficient shared memory for EP scratch");
}
params.smem_size_bytes = smem_size;
// Choose NS: largest value <= MAX_NS that fits in smem (independent of num_k divisibility).
// For stage sizing, use the worst case (full tiles): main_stg = A_SIZE + B_HALF
// First pass: find max B payload across all tiles to determine stage size
int max_b_half = B_HALF; // full tile
// Actually all tiles fit B_HALF at most, tail tiles use less.
// Use full B_HALF for stage size (wastes some smem on tail tiles but keeps layout uniform).
int main_stg = A_SIZE + max_b_half;
params.main_stg = main_stg;
// Find best NS: largest that fits in smem.
int best_ns = 1;
for (int ns = MAX_NS; ns >= 1; ns--) {
int total_smem = ns * main_stg + ns * SF_STAGE;
if (total_smem <= smem_avail) {
best_ns = ns;
break;
}
}
params.ns = best_ns;
int raw_total_tiles = 0;
for (int g = 0; g < G; g++) {
int Mi = A_list[g].size(0);
int Ki = A_list[g].size(1) * 2;
int Ni = B_list[g].size(0);
int nt = (Ni + 255) / 256;
int mt = (Mi + BLOCK_N - 1) / BLOCK_N;
params.groups[g] = {(half *)C_list[g].data_ptr(), Mi, Ni, Ki};
raw_total_tiles += mt * nt;
}
const int cap_clusters = MAX_CLUSTERS;
int num_clusters = raw_total_tiles < cap_clusters ? raw_total_tiles : cap_clusters;
const bool is_bench1 = (params.num_groups == 8 && raw_total_tiles == 176);
const bool is_bench2 = (params.num_groups == 8 && raw_total_tiles == 364);
if (is_bench1 && V_BENCH1_CLUSTERS > 0) num_clusters = V_BENCH1_CLUSTERS;
if (is_bench2 && V_BENCH2_CLUSTERS > 0) num_clusters = V_BENCH2_CLUSTERS;
if (!is_bench1 && !is_bench2 && V_GENERIC_CLUSTERS > 0) num_clusters = V_GENERIC_CLUSTERS;
if (num_clusters > MAX_CLUSTERS) num_clusters = MAX_CLUSTERS;
if (num_clusters > raw_total_tiles) num_clusters = raw_total_tiles;
if (num_clusters < 1) num_clusters = 1;
params.launch_ctas = num_clusters * 2;
int bench_profile = PROFILE_GENERIC_G8;
if (params.num_groups == 8 && raw_total_tiles == 176) {
bench_profile = PROFILE_BENCH1;
} else if (params.num_groups == 8 && raw_total_tiles == 364) {
bench_profile = PROFILE_BENCH2;
} else if (params.num_groups == 2 && raw_total_tiles == 60) {
bench_profile = PROFILE_BENCH3;
} else if (params.num_groups == 2 && raw_total_tiles == 64) {
bench_profile = PROFILE_BENCH4;
} else if (params.num_groups == 2) {
bench_profile = PROFILE_GENERIC_G2;
}
CUtensorMapL2promotion ab_l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_L2_256B;
CUtensorMapL2promotion sf_l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_L2_128B;
switch (bench_profile) {
case PROFILE_BENCH1:
ab_l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_L2_256B;
sf_l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_L2_256B;
break;
case PROFILE_BENCH2:
ab_l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_L2_256B;
sf_l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_NONE;
break;
case PROFILE_BENCH3:
case PROFILE_BENCH4:
case PROFILE_GENERIC_G2:
ab_l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_NONE;
sf_l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_NONE;
break;
case PROFILE_GENERIC_G8:
default:
ab_l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_L2_256B;
sf_l2_promotion = CU_TENSOR_MAP_L2_PROMOTION_L2_128B;
break;
}
TmapParamPackG8 tmap_pack_g8 = {};
bool uniform_nk = true;
for (int g = 1; g < G; g++) {
if (B_list[g].size(0) != B_list[0].size(0) || A_list[g].size(1) != A_list[0].size(1)) {
uniform_nk = false;
break;
}
}
int sfb_src[MAX_GROUPS];
for (int g = 0; g < MAX_GROUPS; g++) sfb_src[g] = -1;
if (uniform_nk) {
int mn_first[4] = {-1, -1, -1, -1};
for (int g = 0; g < G; g++) {
int mnb = (A_list[g].size(0) + 127) / 128;
sfb_src[g] = (mnb < 4) ? mn_first[mnb] : -1;
if (mnb < 4 && mn_first[mnb] < 0) mn_first[mnb] = g;
}
}
for (int g = 0; g < G; g++) {
int Mi = A_list[g].size(0);
int Ki = A_list[g].size(1) * 2;
int Ni = B_list[g].size(0);
int n_tail = Ni % BLOCK_M;
if (g > 0 && uniform_nk) {
tmap_pack_g8.A_full[g] = tmap_pack_g8.A_full[0];
check_cu(cuTensorMapReplaceAddress(&tmap_pack_g8.A_full[g], (void *)B_list[g].data_ptr()));
if (n_tail == 0) {
tmap_pack_g8.A_tail[g] = tmap_pack_g8.A_full[g];
} else {
tmap_pack_g8.A_tail[g] = tmap_pack_g8.A_tail[0];
check_cu(cuTensorMapReplaceAddress(&tmap_pack_g8.A_tail[g], (void *)B_list[g].data_ptr()));
}
} else {
init_AB_tmap(&tmap_pack_g8.A_full[g], (const char *)B_list[g].data_ptr(), Ni, Ki, BLOCK_M, BLOCK_K, ab_l2_promotion);
if (n_tail == 0) {
tmap_pack_g8.A_tail[g] = tmap_pack_g8.A_full[g];
} else {
init_AB_tmap(&tmap_pack_g8.A_tail[g], (const char *)B_list[g].data_ptr(), Ni, Ki, n_tail, BLOCK_K, ab_l2_promotion);
}
}
int m_tail = Mi % BLOCK_N;
int mma_n_tail = (m_tail == 0) ? BLOCK_N : ((m_tail + 15) & ~15);
int b_full_box_h = BLOCK_N / 2;
int b_tail_half = mma_n_tail / 2;
int b_tail0_box_h = m_tail < b_tail_half ? m_tail : b_tail_half;
int b_tail1_box_h = m_tail - b_tail_half;
if (b_tail1_box_h < 0) b_tail1_box_h = 0;
if (b_tail1_box_h > b_tail_half) b_tail1_box_h = b_tail_half;
init_AB_tmap(&tmap_pack_g8.B_full[g], (const char *)A_list[g].data_ptr(), Mi, Ki, b_full_box_h, BLOCK_K, ab_l2_promotion);
if (m_tail == 0) {
tmap_pack_g8.B_tail0[g] = tmap_pack_g8.B_full[g];
tmap_pack_g8.B_tail1[g] = tmap_pack_g8.B_full[g];
} else {
init_AB_tmap(&tmap_pack_g8.B_tail0[g], (const char *)A_list[g].data_ptr(), Mi, Ki, b_tail0_box_h, BLOCK_K, ab_l2_promotion);
const int b_tail1_box_safe = b_tail1_box_h > 0 ? b_tail1_box_h : 1;
init_AB_tmap(&tmap_pack_g8.B_tail1[g], (const char *)A_list[g].data_ptr(), Mi, Ki, b_tail1_box_safe, BLOCK_K, ab_l2_promotion);
}
if (g > 0 && uniform_nk) {
tmap_pack_g8.SFA[g] = tmap_pack_g8.SFA[0];
check_cu(cuTensorMapReplaceAddress(&tmap_pack_g8.SFA[g], (void *)SFB_list[g].data_ptr()));
} else {
init_SF_tmap(&tmap_pack_g8.SFA[g], (const char *)SFB_list[g].data_ptr(), Ni, Ki, sf_l2_promotion);
}
if (uniform_nk && sfb_src[g] >= 0) {
tmap_pack_g8.SFB[g] = tmap_pack_g8.SFB[sfb_src[g]];
check_cu(cuTensorMapReplaceAddress(&tmap_pack_g8.SFB[g], (void *)SFA_list[g].data_ptr()));
} else {
init_SF_tmap(&tmap_pack_g8.SFB[g], (const char *)SFA_list[g].data_ptr(), Mi, Ki, sf_l2_promotion);
}
}
int schedule = SCHED_REV;
if (params.num_groups == 8 && raw_total_tiles == 176) {
schedule = SCHED_F2_G2;
} else if (params.num_groups == 8 && raw_total_tiles == 364) {
schedule = SCHED_F1_G1;
} else if (params.num_groups == 2 && raw_total_tiles == 64) {
schedule = SCHED_REV;
} else if (params.num_groups == 2) {
schedule = SCHED_BASE;
}
if (V_FORCE_SCHEDULE >= 0) schedule = V_FORCE_SCHEDULE;
build_tile_lut(params, schedule);
#define LAUNCH_FOR_PROFILE(PROFILE_ID) \
switch (schedule) { \
case SCHED_BASE: \
launch_grouped_kernel<SCHED_BASE, PROFILE_ID>(params, tmap_pack_g8, smem_size); \
break; \
case SCHED_F2_G2: \
launch_grouped_kernel<SCHED_F2_G2, PROFILE_ID>(params, tmap_pack_g8, smem_size); \
break; \
case SCHED_F1_G1: \
launch_grouped_kernel<SCHED_F1_G1, PROFILE_ID>(params, tmap_pack_g8, smem_size); \
break; \
case SCHED_REV: \
default: \
launch_grouped_kernel<SCHED_REV, PROFILE_ID>(params, tmap_pack_g8, smem_size); \
break; \
}
switch (bench_profile) {
case PROFILE_BENCH1:
LAUNCH_FOR_PROFILE(PROFILE_BENCH1);
break;
case PROFILE_BENCH2:
LAUNCH_FOR_PROFILE(PROFILE_BENCH2);
break;
case PROFILE_BENCH3:
LAUNCH_FOR_PROFILE(PROFILE_BENCH3);
break;
case PROFILE_BENCH4:
LAUNCH_FOR_PROFILE(PROFILE_BENCH4);
break;
case PROFILE_GENERIC_G2:
LAUNCH_FOR_PROFILE(PROFILE_GENERIC_G2);
break;
case PROFILE_GENERIC_G8:
default:
LAUNCH_FOR_PROFILE(PROFILE_GENERIC_G8);
break;
}
#undef LAUNCH_FOR_PROFILE
}
TORCH_LIBRARY(gg_cg2_cta_noep_nosmem, m) {
m.def("run(Tensor[] A, Tensor[] B, Tensor[] C, Tensor[] SFA, Tensor[] SFB) -> ()");
m.impl("run", &grouped_gemm_impl);
}
"""
load_inline(
"grouped_gemm_cg2_cta_noep_nosmem",
cpp_sources="",
cuda_sources=cuda_src,
is_python_module=False,
no_implicit_headers=True,
extra_cuda_cflags=[
"-O3", "-gencode=arch=compute_100a,code=sm_100a",
"--use_fast_math", "--expt-relaxed-constexpr",
"--relocatable-device-code=false", "-lineinfo",
],
extra_ldflags=["-lcuda"],
)
_run = torch.ops.gg_cg2_cta_noep_nosmem.run
def custom_kernel(data: input_t) -> output_t:
abc, _, sf_reordered, _ = data
a, b, c = zip(*abc)
sfa, sfb = zip(*sf_reordered)
_run(list(a), list(b), list(c), list(sfa), list(sfb))
return list(c)
scrolls · 1069 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 487098.
⋯ 3 unchanged linesfrom task import input_t, output_tfrom torch.utils.cpp_extension import load_inline+ # CTA-copy no-ep variant aligned to cta_full control flow.cuda_src = """#include <cudaTypedefs.h>#include <cuda_fp16.h>⋯ 16 unchanged linesconstexpr int BLOCK_K = 256;constexpr int A_SIZE = BLOCK_M * BLOCK_K / 2; // 16384constexpr int B_SIZE = BLOCK_N * BLOCK_K / 2; // 16384+ constexpr int B_HALF = BLOCK_N * BLOCK_K / 4; // 8192 (per CTA in cta_group::2)constexpr int SFA_SIZE = 128 * BLOCK_K / 16; // 2048constexpr int SFB_SIZE = 128 * BLOCK_K / 16; // 2048- constexpr int MAIN_STAGE = A_SIZE + B_SIZE; // 32768constexpr int SF_STAGE = SFA_SIZE + SFB_SIZE; // 4096- constexpr int TMAP_SMEM = 4 * 128; // 512+ constexpr int EP_STRIDE = 136; // padded stride for chunk32 epilogue+ constexpr int EP_SMEM_BYTES = 32 * EP_STRIDE * (int)sizeof(half); // 8704B scratch- constexpr int NUM_STAGES = 6;- constexpr int SMEM_SIZE = TMAP_SMEM + MAIN_STAGE * NUM_STAGES + SF_STAGE * NUM_STAGES;- constexpr int NUM_MBAR = NUM_STAGES * 2 + 2; // tma + mma + 2xdone-constexpr int NUM_EP_WARPS = 4;constexpr int NUM_WARPS = NUM_EP_WARPS + 2; // 6constexpr int TB_SIZE = NUM_WARPS * WARP_SIZE; // 192constexpr int TMEM_COLS = 512;constexpr int MAX_LAUNCH_CTAS = 148;+ constexpr int MAX_CLUSTERS = MAX_LAUNCH_CTAS / 2;+ constexpr int MAX_TILE_LUT = 1024;constexpr int D_TMEM0 = 0;constexpr int D_TMEM1 = BLOCK_N; // 128constexpr int SFA_TMEM = 2 * BLOCK_N; // 256constexpr int SFB_TMEM = SFA_TMEM + 4 * (BLOCK_K / MMA_K); // 272- constexpr uint32_t I_DESC = (1U << 7U) | (1U << 10U) |- ((uint32_t)BLOCK_N >> 3U << 17U) | ((uint32_t)BLOCK_M >> 7U << 27U);-constexpr int MAX_GROUPS = 8;- constexpr int TMAPS_PER_GROUP = 5;- constexpr int TMAP_A_FULL = 0;- constexpr int TMAP_A_TAIL = 1;- constexpr int TMAP_B = 2;- constexpr int TMAP_SFA = 3;- constexpr int TMAP_SFB = 4;+ constexpr int MAX_NS = 10; // maximum pipeline stages// ============================================================================// Device structures⋯ 1 unchanged linesstruct GroupInfo {half* c_ptr;int M, N, K;- int tile_offset;- int m_tiles, n_tiles;};+ // Fat LUT: all per-tile info pre-computed host-side. SoA layout with uint16 arrays.+ // No fused-M: each LUT entry is a single (coord_n, coord_m) tile.struct KernelParams {GroupInfo groups[MAX_GROUPS];int num_groups;int total_tiles;int launch_ctas;+ int smem_size_bytes;+ int ns; // pipeline stages (smem-driven; independent of num_k divisibility)+ int main_stg; // max stage size for smem budgeting / NS selection++ int32_t lut_worker_start[MAX_CLUSTERS];+ int32_t lut_worker_count[MAX_CLUSTERS];++ // Fat LUT arrays — pre-computed per tile.+ uint16_t lut_gidx[MAX_TILE_LUT]; // group index+ uint16_t lut_coord_m[MAX_TILE_LUT]; // M-tile coord (for SFB tmap)+ uint16_t lut_off_n_cta_r0[MAX_TILE_LUT]; // N offset for cta_rank=0+ uint16_t lut_off_n_cta_r1[MAX_TILE_LUT]; // N offset for cta_rank=1+ uint16_t lut_coord_n_cta_r0[MAX_TILE_LUT]; // N-tile coord for cta_rank=0 (SFA tmap)+ uint16_t lut_coord_n_cta_r1[MAX_TILE_LUT]; // N-tile coord for cta_rank=1 (SFA tmap)+ uint16_t lut_off_m_b_r0[MAX_TILE_LUT]; // B operand y-offset for cta_rank=0+ uint16_t lut_off_m_b_r1[MAX_TILE_LUT]; // B operand y-offset for cta_rank=1+ uint16_t lut_expect_bytes_r0[MAX_TILE_LUT]; // TMA expect_tx total bytes for cta_rank=0+ uint16_t lut_expect_bytes_r1[MAX_TILE_LUT]; // TMA expect_tx total bytes for cta_rank=1+ uint16_t lut_mma_n_num_k[MAX_TILE_LUT]; // [7:0]=mma_n, [15:8]=num_k (K/BLOCK_K)+ // Packed tmap selection: bits [1:0]=a_tmap_r0, [3:2]=a_tmap_r1, [5:4]=b_tmap_r0, [7:6]=b_tmap_r1+ uint8_t lut_tmap_sel[MAX_TILE_LUT];};enum : int {- SCHED_BASE = 0, // f0_g0- SCHED_REV = 1, // f1_g0- SCHED_F2_G2 = 2, // f2_g2- SCHED_F1_G1 = 3, // f1_g1+ SCHED_BASE = 0,+ SCHED_REV = 1,+ SCHED_F2_G2 = 2,+ SCHED_F1_G1 = 3,};enum : int {PROFILE_GENERIC_G8 = 0,PROFILE_GENERIC_G2 = 1,- PROFILE_BENCH1 = 2, // g=8, total_tiles=352- PROFILE_BENCH2 = 3, // g=8, total_tiles=728- PROFILE_BENCH3 = 4, // g=2, total_tiles=120- PROFILE_BENCH4 = 5, // g=2, total_tiles=128+ PROFILE_BENCH1 = 2,+ PROFILE_BENCH2 = 3,+ PROFILE_BENCH3 = 4,+ PROFILE_BENCH4 = 5,};+ constexpr int V_FORCE_SCHEDULE = -1;+ constexpr int V_BENCH1_CLUSTERS = -1;+ constexpr int V_BENCH2_CLUSTERS = -1;+ constexpr int V_GENERIC_CLUSTERS = -1;+struct TmapParamPackG8 {- CUtensorMap A_full[MAX_GROUPS];- CUtensorMap A_tail[MAX_GROUPS];- CUtensorMap B[MAX_GROUPS];- CUtensorMap SFA[MAX_GROUPS];- CUtensorMap SFB[MAX_GROUPS];+ CUtensorMap A_full[MAX_GROUPS]; // orig B data (MMA A operand), box_h=128+ CUtensorMap A_tail[MAX_GROUPS]; // orig B data, N-tail+ CUtensorMap B_full[MAX_GROUPS]; // orig A data (MMA B operand), box_h=64 per CTA+ CUtensorMap B_tail0[MAX_GROUPS]; // orig A data, M-tail for CTA0+ CUtensorMap B_tail1[MAX_GROUPS]; // orig A data, M-tail for CTA1+ CUtensorMap SFA[MAX_GROUPS]; // orig SFB (scale for MMA A)+ CUtensorMap SFB[MAX_GROUPS]; // orig SFA (scale for MMA B)};// ============================================================================⋯ 15 unchanged linesreturn pred;}+ __device__ inline uint32_t get_cluster_ctarank() {+ uint32_t rank = 0;+ asm volatile("mov.u32 %0, %%cluster_ctarank;" : "=r"(rank));+ return rank;+ }+__device__ inline void mbarrier_init(int mbar_addr, int count) {asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));}__device__ void mbarrier_wait(int mbar_addr, int phase) {- uint32_t ticks = 0x989680;+ uint32_t ticks = 0x10000;asm volatile("{\\n\\t"".reg .pred P1;\\n\\t""LAB_WAIT:\\n\\t""mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\\n\\t"- "@P1 bra.uni DONE;\\n\\t"- "bra.uni LAB_WAIT;\\n\\t"- "DONE:\\n\\t"+ "@!P1 bra.uni LAB_WAIT;\\n\\t""}":: "r"(mbar_addr), "r"(phase), "r"(ticks));}__device__ inline void mbarrier_arrive_expect_tx(int mbar_addr, int size) {- asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"+ asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;":: "r"(mbar_addr), "r"(size) : "memory");}- __device__ inline void tma_load_1d(int dst, const void *tmap_ptr, int x, int mbar_addr, uint64_t cache_policy) {- uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(tmap_ptr);- asm volatile("cp.async.bulk.tensor.1d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "- "[%0], [%1, {%3}], [%2], %4;"- :: "r"(dst), "l"(gmem_int_desc), "r"(mbar_addr), "r"(x), "l"(cache_policy) : "memory");+ template <int CTA_GROUP = 2>+ __device__ inline void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {+ asm volatile("cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%7.L2::cache_hint "+ "[%0], [%1, {%2, %3, %4}], [%5], %6;"+ :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy), "n"(CTA_GROUP)+ : "memory");}- __device__ inline void tma_load_3d(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr, uint64_t cache_policy) {- uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(tmap_ptr);- asm volatile("cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::1.L2::cache_hint "- "[%0], [%1, {%3, %4, %5}], [%2], %6;"- :: "r"(dst), "l"(gmem_int_desc), "r"(mbar_addr),- "r"(x), "r"(y), "r"(z), "l"(cache_policy) : "memory");- }-+ template <int CTA_GROUP = 2>__device__ inline void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {- asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));+ asm volatile("tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc), "n"(CTA_GROUP));}- __device__ inline void tcgen05_mma_nvfp4(uint64_t a_desc, uint64_t b_desc, uint32_t i_desc,- int scale_A_tmem, int scale_B_tmem, int enable_input_d, int d_tmem) {+ template <int CTA_GROUP = 2>+ __device__ inline void tcgen05_mma_nvfp4(int d_tmem, uint64_t a_desc, uint64_t b_desc, uint32_t i_desc,+ int scale_A_tmem, int scale_B_tmem, int enable_input_d) {asm volatile("{\\n\\t"".reg .pred p;\\n\\t""setp.ne.b32 p, %6, 0;\\n\\t"- "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], p;\\n\\t"+ "tcgen05.mma.cta_group::%7.kind::mxf4nvf4.block_scale.block16.collector::a::lastuse [%0], %1, %2, %3, [%4], [%5], p;\\n\\t""}":: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),- "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d));+ "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d), "n"(CTA_GROUP));}+ template <int CTA_GROUP = 2>__device__ inline void tcgen05_commit(int mbar_addr) {- asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"- :: "r"(mbar_addr) : "memory");+ asm volatile("tcgen05.commit.cta_group::%1.mbarrier::arrive::one.shared::cluster.b64 [%0];"+ :: "r"(mbar_addr), "n"(CTA_GROUP) : "memory");}- static constexpr char SHAPE_16x256b[] = ".16x256b";- static constexpr char NUM_x2[] = ".x2";+ template <int CTA_GROUP = 2>+ __device__ inline void tcgen05_commit_mcast(int mbar_addr, uint16_t cta_mask) {+ asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"+ :: "r"(mbar_addr), "h"(cta_mask), "n"(CTA_GROUP) : "memory");+ }__device__ inline void tcgen05_ld_16x256bx2(float *tmp, int row, int col) {asm volatile("tcgen05.ld.sync.aligned.16x256b.x2.b32 "⋯ 3 unchanged lines: "r"((row << 16) | col));}- __device__ inline void fence_proxy_tensormap(const void *smem_ptr) {- uint64_t addr = reinterpret_cast<uint64_t>(smem_ptr);- asm volatile("fence.proxy.tensormap::generic.acquire.gpu [%0], 128;" :: "l"(addr));- }-- __device__ inline void fence_proxy_tensormap_release_gpu() {- asm volatile("fence.proxy.tensormap::generic.release.gpu;" ::: "memory");- }-- __device__ inline void tmap_replace_global_address(CUtensorMap *tmap_ptr, uint64_t new_addr) {- asm volatile("tensormap.replace.tile.global_address.global.b1024.b64 [%0], %1;"- :: "l"(tmap_ptr), "l"(new_addr) : "memory");- }-- __device__ inline void do_epilogue(int warp_id, int lane_id, int done_mbar, int d_tmem_base,- half* c_ptr, int M, int N, int off_m, int off_n) {- mbarrier_wait(done_mbar, 0);+ // Chunk32 transposed epilogue (copied from cg2_full style).+ __device__ inline void do_epilogue_transposed_chunk32(+ int warp_id, int lane_id, int cta_rank,+ int done_mbar, int done_phase, int d_tmem_base,+ half* __restrict__ smem_ep, half* __restrict__ c_ptr,+ int M, int N, int off_m, int off_n, int mma_n+ ) {+ mbarrier_wait(done_mbar, done_phase);asm volatile("tcgen05.fence::after_thread_sync;");const int col_lane = (lane_id % 4) * 2;const int row_lane = lane_id / 4;- const int residue_m = M - off_m;+ const int tid_ep = warp_id * WARP_SIZE + lane_id;+ const int residue_n = N - off_n;+ const int num_chunks = (mma_n + 31) / 32;- #pragma unroll- for (int m = 0; m < 2; m++) {- const int tm = warp_id * 32 + m * 16;- const int out_row0 = off_m + tm + row_lane;- const int out_row1 = out_row0 + 8;+ #pragma unroll 1+ for (int chunk = 0; chunk < num_chunks; chunk++) {+ #pragma unroll+ for (int mc = 0; mc < 2; mc++) {+ #pragma unroll+ for (int m = 0; m < 2; m++) {+ const int tm = cta_rank * BLOCK_M + warp_id * 32 + m * 16;+ float vals[8];+ tcgen05_ld_16x256bx2(vals, tm, d_tmem_base + chunk * 32 + mc * 16);+ asm volatile("tcgen05.wait::ld.sync.aligned;");+ const int n0 = warp_id * 32 + m * 16 + row_lane;+ const int n1 = n0 + 8;+ const int m_off = mc * 16;++ smem_ep[(col_lane + m_off) * EP_STRIDE + n0] = __float2half_rn(vals[0]);+ smem_ep[(col_lane + 1 + m_off) * EP_STRIDE + n0] = __float2half_rn(vals[1]);+ smem_ep[(col_lane + m_off) * EP_STRIDE + n1] = __float2half_rn(vals[2]);+ smem_ep[(col_lane + 1 + m_off) * EP_STRIDE + n1] = __float2half_rn(vals[3]);+ smem_ep[(col_lane + 8 + m_off) * EP_STRIDE + n0] = __float2half_rn(vals[4]);+ smem_ep[(col_lane + 9 + m_off) * EP_STRIDE + n0] = __float2half_rn(vals[5]);+ smem_ep[(col_lane + 8 + m_off) * EP_STRIDE + n1] = __float2half_rn(vals[6]);+ smem_ep[(col_lane + 9 + m_off) * EP_STRIDE + n1] = __float2half_rn(vals[7]);+ }+ }++ asm volatile("bar.sync 15, %0;" :: "r"(NUM_EP_WARPS * WARP_SIZE));++ const int n_group = tid_ep % 8;+ const int n_start = n_group * 16;#pragma unroll- for (int chunk = 0; chunk < BLOCK_N / 16; chunk++) {- float vals[8];- tcgen05_ld_16x256bx2(vals, tm, d_tmem_base + chunk * 16);- asm volatile("tcgen05.wait::ld.sync.aligned;");+ for (int pass = 0; pass < 2; pass++) {+ const int m_row = pass * 16 + tid_ep / 8;+ const int m_local = chunk * 32 + m_row;+ const int m_global = off_m + m_local;- // repeat 0: cols [chunk*16 .. chunk*16+7], repeat 1: cols [chunk*16+8 .. chunk*16+15]- const int out_col0 = off_n + chunk * 16 + col_lane;- const int out_col1 = off_n + chunk * 16 + 8 + col_lane;+ if (m_local < mma_n && m_global < M && n_start < residue_n) {+ half* dst_base = &c_ptr[m_global * N + off_n + n_start];+ const half* src_base = &smem_ep[m_row * EP_STRIDE + n_start];- if (tm + row_lane < residue_m) {- reinterpret_cast<half2 *>(c_ptr + out_row0 * N + out_col0)[0] =- __float22half2_rn({vals[0], vals[1]});- reinterpret_cast<half2 *>(c_ptr + out_row0 * N + out_col1)[0] =- __float22half2_rn({vals[4], vals[5]});+ if (n_start + 16 <= residue_n) {+ *reinterpret_cast<int4*>(dst_base) = *reinterpret_cast<const int4*>(src_base);+ *reinterpret_cast<int4*>(dst_base + 8) = *reinterpret_cast<const int4*>(src_base + 8);+ } else {+ #pragma unroll+ for (int i = 0; i < 16; i++) {+ if (n_start + i < residue_n) dst_base[i] = src_base[i];+ }+ }}- if (tm + row_lane + 8 < residue_m) {- reinterpret_cast<half2 *>(c_ptr + out_row1 * N + out_col0)[0] =- __float22half2_rn({vals[2], vals[3]});- reinterpret_cast<half2 *>(c_ptr + out_row1 * N + out_col1)[0] =- __float22half2_rn({vals[6], vals[7]});- }}++ asm volatile("bar.sync 15, %0;" :: "r"(NUM_EP_WARPS * WARP_SIZE));}}+ __device__ inline void do_epilogue_transposed(+ int warp_id, int lane_id, int cta_rank,+ int done_mbar, int done_phase, int d_tmem_base,+ half* __restrict__ smem_ep, half* __restrict__ c_ptr,+ int M, int N, int off_m, int off_n, int mma_n+ ) {+ do_epilogue_transposed_chunk32(+ warp_id, lane_id, cta_rank, done_mbar, done_phase, d_tmem_base,+ smem_ep, c_ptr, M, N, off_m, off_n, mma_n);+ }+// ============================================================================// TensorMap Initialization// ============================================================================⋯ 17 unchanged linesl2_promotion, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE));}- // SF reordered tensors have logical shape [32, 4, rest_m, 4, rest_k, L] but are- // a permuted view of a contiguous [L, rest_m, rest_k, 32, 4, 4] allocation.- // Physical memory is thus [rest_m][rest_k][512 bytes], i.e. each 512-byte SF tile- // (covering 128 M-rows x 1 MMA_K=64 step) is already contiguous.- // We encode this as a 3D TMA: dim0 = 256 uint16 (=512B block), dim1 = mn_blocks, dim2 = k_blocks.void init_SF_tmap(CUtensorMap *tmap, const char *ptr, uint64_t mn, uint64_t K,CUtensorMapL2promotion l2_promotion) {constexpr uint32_t rank = 3;const uint64_t k_blocks = K / 64;const uint64_t mn_blocks = (mn + 127) / 128;- const uint32_t tile_k_blocks = BLOCK_K / 64; // 4+ const uint32_t tile_k_blocks = BLOCK_K / 64;constexpr uint64_t SF_BLOCK_BYTES = 512;- constexpr uint64_t X_ELEMS = SF_BLOCK_BYTES / sizeof(uint16_t); // 256+ constexpr uint64_t X_ELEMS = SF_BLOCK_BYTES / sizeof(uint16_t);uint64_t globalDim[rank] = {X_ELEMS, mn_blocks, k_blocks};uint64_t globalStrides[rank-1] = {k_blocks * SF_BLOCK_BYTES, SF_BLOCK_BYTES};uint32_t boxDim[rank] = {(uint32_t)X_ELEMS, 1, tile_k_blocks};⋯ 5 unchanged lines}// ============================================================================- // Kernel+ // Host-side fat LUT builder// ============================================================================+ struct TileLutTmp {+ uint16_t gidx;+ uint16_t coord_m;+ uint16_t off_n_cta_r0, off_n_cta_r1;+ uint16_t coord_n_cta_r0, coord_n_cta_r1;+ uint16_t off_m_b_r0, off_m_b_r1;+ uint16_t expect_bytes_r0, expect_bytes_r1;+ uint16_t mma_n_num_k; // [7:0]=mma_n, [15:8]=num_k+ uint8_t tmap_sel; // packed tmap indices+ };++ inline void build_tile_lut(KernelParams& params, int schedule) {+ TileLutTmp tmp[MAX_TILE_LUT];+ int total = 0;++ for (int g = 0; g < params.num_groups; g++) {+ const GroupInfo& gi = params.groups[g];+ const int M = gi.M, N = gi.N, K = gi.K;+ const int n_tiles = (N + 255) / 256;+ const int m_tiles = (M + BLOCK_N - 1) / BLOCK_N;+ const int m_rem = M % BLOCK_N;+ const int mma_n_tail = (m_rem == 0) ? BLOCK_N : ((m_rem + 15) & ~15);++ for (int cn = 0; cn < n_tiles; cn++) {+ for (int cm = 0; cm < m_tiles; cm++) {+ TORCH_CHECK(total < MAX_TILE_LUT, "tile LUT overflow");++ const bool is_m_tail = (cm == m_tiles - 1) && (mma_n_tail != BLOCK_N);+ const int mma_n = is_m_tail ? mma_n_tail : BLOCK_N;+ const int b_half_rows = mma_n / 2;+ const int off_m_base = cm * BLOCK_N;++ // Per cta_rank: compute N-axis coords+ int off_n_r0 = cn * 256;+ int off_n_r1 = cn * 256 + BLOCK_M;+ int coord_n_cta_r0 = cn * 2;+ int coord_n_cta_r1 = cn * 2 + 1;+ // Handle N-tail: if cta_rank=1 would be out of bounds, alias to rank=0+ if (off_n_r1 >= N) {+ off_n_r1 = off_n_r0;+ coord_n_cta_r1 = coord_n_cta_r0;+ }++ // A tmap selection (based on N residue for each rank)+ int n_residue_r0 = N - off_n_r0;+ int n_residue_r1 = N - off_n_r1;+ bool is_a_tail_r0 = (n_residue_r0 > 0 && n_residue_r0 < BLOCK_M);+ bool is_a_tail_r1 = (n_residue_r1 > 0 && n_residue_r1 < BLOCK_M);+ // Both ranks in a cluster see same coord_n, but different cta offsets.+ // The A tmap idx is actually the same for both ranks within same coord_n+ // because A_full vs A_tail depends on the per-CTA N residue.+ // We store per-rank since they can differ.+ uint16_t a_tmap_r0 = is_a_tail_r0 ? 1 : 0;+ uint16_t a_tmap_r1 = is_a_tail_r1 ? 1 : 0;+ // For simplicity, store worst case (if either is tail, both get tail idx)+ // Actually no - each CTA independently selects its A tmap. Store per-rank.+ // But our LUT only has one a_tmap_idx field. Let's use the cta_rank to select.+ // Actually, a_tail only matters at the N boundary. For cn < n_tiles-1, both are full.+ // For cn == n_tiles-1, rank0 might be tail, rank1 might be OOB (aliased to rank0).+ // So if rank1 is aliased to rank0, they share the same tail status.+ // Let's just store per-rank.++ int a_bytes_r0 = is_a_tail_r0 ? (n_residue_r0 * BLOCK_K / 2) : A_SIZE;+ int a_bytes_r1 = is_a_tail_r1 ? (n_residue_r1 * BLOCK_K / 2) : A_SIZE;++ // B operand offsets per rank+ int off_m_b_r0 = off_m_base;+ int off_m_b_r1 = off_m_base + b_half_rows;+ int b_rows_r0 = b_half_rows;+ int b_rows_r1 = b_half_rows;++ // B tmap selection+ uint16_t b_tmap_r0 = 0; // B_full+ uint16_t b_tmap_r1 = 0; // B_full++ if (is_m_tail) {+ const int m_residue = M - off_m_base;+ b_rows_r0 = (m_residue < b_half_rows) ? m_residue : b_half_rows;+ b_rows_r1 = m_residue - b_half_rows;+ if (b_rows_r1 < 0) b_rows_r1 = 0;+ if (b_rows_r1 > b_half_rows) b_rows_r1 = b_half_rows;+ b_tmap_r0 = 1; // B_tail0+ b_tmap_r1 = 2; // B_tail1+ }+ if (b_rows_r0 < 1) b_rows_r0 = 1;+ if (b_rows_r1 < 1) {+ b_rows_r1 = 1;+ off_m_b_r1 = off_m_base; // safe fallback+ }++ int b_bytes_r0 = b_rows_r0 * BLOCK_K / 2;+ int b_bytes_r1 = b_rows_r1 * BLOCK_K / 2;+ int expect_r0 = a_bytes_r0 + b_bytes_r0 + SFA_SIZE + SFB_SIZE;+ int expect_r1 = a_bytes_r1 + b_bytes_r1 + SFA_SIZE + SFB_SIZE;+ // Pack tmap indices: [1:0]=a_r0, [3:2]=a_r1, [5:4]=b_r0, [7:6]=b_r1+ uint8_t tmap_sel = (a_tmap_r0 & 3) | ((a_tmap_r1 & 3) << 2)+ | ((b_tmap_r0 & 3) << 4) | ((b_tmap_r1 & 3) << 6);++ tmp[total++] = {+ (uint16_t)g,+ (uint16_t)cm,+ (uint16_t)off_n_r0, (uint16_t)off_n_r1,+ (uint16_t)coord_n_cta_r0, (uint16_t)coord_n_cta_r1,+ (uint16_t)off_m_b_r0, (uint16_t)off_m_b_r1,+ (uint16_t)expect_r0, (uint16_t)expect_r1,+ (uint16_t)((mma_n & 0xFF) | ((K / BLOCK_K) << 8)),+ tmap_sel,+ };+ }+ }+ }++ params.total_tiles = total;+ TORCH_CHECK(params.total_tiles <= MAX_TILE_LUT, "total_tiles exceeds LUT capacity");++ const int num_clusters = params.launch_ctas / 2;+ TORCH_CHECK(num_clusters > 0 && num_clusters <= MAX_CLUSTERS, "num_clusters out of range");++ int cursor = 0;+ for (int cid = 0; cid < num_clusters; cid++) {+ int logical_cid = cid;+ if (schedule == SCHED_F2_G2) logical_cid = (cid * 17) % num_clusters;+ else if (schedule == SCHED_F1_G1) logical_cid = num_clusters - 1 - cid;++ const int my_count = (params.total_tiles - logical_cid + num_clusters - 1) / num_clusters;+ params.lut_worker_start[cid] = cursor;+ params.lut_worker_count[cid] = my_count;++ for (int tile_iter = 0; tile_iter < my_count; tile_iter++) {+ int k = tile_iter;+ if (schedule == SCHED_REV || schedule == SCHED_F1_G1) {+ k = my_count - 1 - tile_iter;+ } else if (schedule == SCHED_F2_G2) {+ const int h = (my_count + 1) >> 1;+ k = (tile_iter < h) ? (tile_iter << 1) : (((tile_iter - h) << 1) + 1);+ }+ const int tile_id = logical_cid + k * num_clusters;+ TORCH_CHECK(tile_id >= 0 && tile_id < params.total_tiles, "tile_id out of range");++ const TileLutTmp& t = tmp[tile_id];+ params.lut_gidx[cursor] = t.gidx;+ params.lut_coord_m[cursor] = t.coord_m;+ params.lut_off_n_cta_r0[cursor] = t.off_n_cta_r0;+ params.lut_off_n_cta_r1[cursor] = t.off_n_cta_r1;+ params.lut_coord_n_cta_r0[cursor] = t.coord_n_cta_r0;+ params.lut_coord_n_cta_r1[cursor] = t.coord_n_cta_r1;+ params.lut_off_m_b_r0[cursor] = t.off_m_b_r0;+ params.lut_off_m_b_r1[cursor] = t.off_m_b_r1;+ params.lut_expect_bytes_r0[cursor] = t.expect_bytes_r0;+ params.lut_expect_bytes_r1[cursor] = t.expect_bytes_r1;+ params.lut_mma_n_num_k[cursor] = t.mma_n_num_k;+ params.lut_tmap_sel[cursor] = t.tmap_sel;+ cursor++;+ }+ }++ for (int cid = num_clusters; cid < MAX_CLUSTERS; cid++) {+ params.lut_worker_start[cid] = 0;+ params.lut_worker_count[cid] = 0;+ }++ TORCH_CHECK(cursor == params.total_tiles, "tile LUT size mismatch");+ }++ // ============================================================================+ // Kernel — persistent mbars, fat LUT, NS as template param+ // ============================================================================template <int SCHEDULE_ID, int PROFILE_ID>- __global__ __launch_bounds__(TB_SIZE)+ __global__ __cluster_dims__(2, 1, 1) __launch_bounds__(TB_SIZE)void grouped_gemm_kernel(const __grid_constant__ KernelParams params,const __grid_constant__ TmapParamPackG8 tmap_pack_g8) {- struct EpMeta {- half* c_ptr;- int M, N;- int off_m, off_n;- };-const int tid = threadIdx.x;const int warp_id = tid / WARP_SIZE;const int lane_id = tid % WARP_SIZE;const int bid = blockIdx.x;if (bid >= params.launch_ctas) return;- int logical_bid = bid;- if constexpr (SCHEDULE_ID == SCHED_F2_G2) {- logical_bid = (bid * 17) % params.launch_ctas;- } else if constexpr (SCHEDULE_ID == SCHED_F1_G1) {- logical_bid = params.launch_ctas - 1 - bid;- }- const int my_count = (params.total_tiles - logical_bid + params.launch_ctas - 1) / params.launch_ctas;- if (my_count <= 0) return;- // --- SMEM setup ---+ const int cta_rank = static_cast<int>(get_cluster_ctarank());+ const int cluster_id = bid / 2;+ const int my_count = params.lut_worker_count[cluster_id];+ const int worker_start = params.lut_worker_start[cluster_id];+extern __shared__ __align__(1024) char smem_raw[];const int smem = static_cast<int>(__cvta_generic_to_shared(smem_raw));- const int smem_main = smem + TMAP_SMEM;- const int smem_sf = smem_main + MAIN_STAGE * NUM_STAGES;+ const int smem_main = smem;+ half* smem_ep = reinterpret_cast<half*>(smem_raw + params.smem_size_bytes - EP_SMEM_BYTES);+ const int NS = params.ns;+ const int main_stg = params.main_stg;+ const int smem_sf = smem_main + NS * main_stg;++ // Mbarrier layout:+ // - NS tma + NS mma+ // - 2 done mbars (TMEM slot ready for EP)+ // - 2 ep mbars (EP drained slot; TMEM backpressure)#pragma nv_diag_suppress static_var_with_dynamic_init- __shared__ int64_t mbars[NUM_MBAR];+ __shared__ int64_t mbars[2 * MAX_NS + 4];__shared__ int32_t tmem_alloc_buf;- __shared__ EpMeta ep_meta[2];const int mbar_base = static_cast<int>(__cvta_generic_to_shared(mbars));const int tma_mbar = mbar_base;- const int mma_mbar = tma_mbar + NUM_STAGES * 8;- const int done_mbar0 = mma_mbar + NUM_STAGES * 8;+ const int mma_mbar = tma_mbar + NS * 8;+ const int done_mbar0 = mma_mbar + NS * 8;const int done_mbar1 = done_mbar0 + 8;+ const int ep_mbar0 = done_mbar1 + 8;+ const int ep_mbar1 = ep_mbar0 + 8;- // Allocate TMEM once for this CTA.- if (warp_id == 1) {+ if (my_count <= 0) return;++ // === ONE-TIME SETUP ===++ // MMA warp: allocate TMEM (both CTAs must issue)+ if (warp_id == NUM_WARPS - 1) {int alloc_addr = static_cast<int>(__cvta_generic_to_shared(&tmem_alloc_buf));- asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"+ asm volatile("tcgen05.alloc.cta_group::2.sync.aligned.shared::cta.b32 [%0], %1;":: "r"(alloc_addr), "r"(TMEM_COLS));}- __syncthreads();- // --- Descriptor helpers ---- auto make_desc_AB = [](int addr) -> uint64_t {- return desc_encode(addr) | (desc_encode(8 * 128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);- };+ // Warp 0: init all mbarriers ONCE for the entire kernel lifetime+ if (warp_id == 0 && elect_sync()) {+ for (int i = 0; i < NS; i++) {+ mbarrier_init(tma_mbar + i * 8, 2); // 2 CTAs arrive+ mbarrier_init(mma_mbar + i * 8, 1); // 1 MMA warp (CTA0) arrives+ }+ mbarrier_init(done_mbar0, 1);+ mbarrier_init(done_mbar1, 1);+ mbarrier_init(ep_mbar0, 2);+ mbarrier_init(ep_mbar1, 2);+ asm volatile("fence.mbarrier_init.release.cluster;");+ }++ // Single cluster barrier for the entire kernel+ asm volatile("barrier.cluster.arrive.relaxed.aligned;");+ asm volatile("barrier.cluster.wait.acquire.aligned;");+constexpr uint64_t cache_A =(PROFILE_ID == PROFILE_BENCH3 || PROFILE_ID == PROFILE_BENCH4) ? EVICT_FIRST :((PROFILE_ID == PROFILE_GENERIC_G2) ? EVICT_NORMAL : 0ULL);⋯ 1 unchanged lines(PROFILE_ID == PROFILE_BENCH3 || PROFILE_ID == PROFILE_GENERIC_G2) ? EVICT_FIRST :((PROFILE_ID == PROFILE_BENCH4) ? EVICT_NORMAL : 0ULL);constexpr uint64_t cache_SF = cache_B;- constexpr int SF_K_PER_BLOCK = BLOCK_K / 64; // 4- for (int tile_iter = 0; tile_iter < my_count; tile_iter++) {- const int slot = tile_iter & 1;- const int prev_slot = slot ^ 1;- const int d_tmem_base = slot ? D_TMEM1 : D_TMEM0;- const int done_mbar = slot ? done_mbar1 : done_mbar0;- const int prev_d_tmem_base = prev_slot ? D_TMEM1 : D_TMEM0;- const int prev_done_mbar = prev_slot ? done_mbar1 : done_mbar0;+ constexpr uint16_t cta_mask = 0x3;+ constexpr int SF_K_PER_BLOCK_L = BLOCK_K / 64;+ constexpr uint32_t MMA_M_2CTA = BLOCK_M * 2;- int k = tile_iter;- if constexpr (SCHEDULE_ID == SCHED_REV || SCHEDULE_ID == SCHED_F1_G1) {- k = my_count - 1 - tile_iter;- } else if constexpr (SCHEDULE_ID == SCHED_F2_G2) {- const int h = (my_count + 1) >> 1;- k = (tile_iter < h) ? (tile_iter << 1) : (((tile_iter - h) << 1) + 1);- }- const int tile_id = logical_bid + k * params.launch_ctas;- int gidx = 0;- #pragma unroll- for (int g = 1; g < MAX_GROUPS; g++) {- if (g < params.num_groups && tile_id >= params.groups[g].tile_offset)- gidx = g;- }- const GroupInfo& gi = params.groups[gidx];- const int local_tile = tile_id - gi.tile_offset;- const int coord_x = local_tile % gi.m_tiles;- const int coord_y = local_tile / gi.m_tiles;- const int M = gi.M, N = gi.N, K = gi.K;- const int num_k = K / BLOCK_K;- const int off_m = coord_x * BLOCK_M;- const int off_n = coord_y * BLOCK_N;+ auto make_desc_AB = [](int addr) -> uint64_t {+ return desc_encode(addr) | (desc_encode(8 * 128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);+ };- if (warp_id == 0 && lane_id == 0) {- ep_meta[slot].c_ptr = gi.c_ptr;- ep_meta[slot].M = M;- ep_meta[slot].N = N;- ep_meta[slot].off_m = off_m;- ep_meta[slot].off_n = off_n;- }+ // Helper to select A tmap based on index+ auto get_a_tmap = [&](int gidx, int a_idx) -> const void* {+ if (a_idx == 0) return static_cast<const void*>(&tmap_pack_g8.A_full[gidx]);+ return static_cast<const void*>(&tmap_pack_g8.A_tail[gidx]);+ };- // Reset tile-local pipeline barriers before this tile starts.- if (warp_id == 0 && elect_sync()) {- #pragma unroll- for (int i = 0; i < NUM_STAGES; i++) {- mbarrier_init(tma_mbar + i * 8, 1);- mbarrier_init(mma_mbar + i * 8, 1);- }- mbarrier_init(done_mbar, 1);- asm volatile("fence.mbarrier_init.release.cluster;");- }- __syncthreads();+ // Helper to select B tmap based on index+ auto get_b_tmap = [&](int gidx, int b_idx) -> const void* {+ if (b_idx == 0) return static_cast<const void*>(&tmap_pack_g8.B_full[gidx]);+ if (b_idx == 1) return static_cast<const void*>(&tmap_pack_g8.B_tail0[gidx]);+ return static_cast<const void*>(&tmap_pack_g8.B_tail1[gidx]);+ };- const int m_tail = M % BLOCK_M;- const bool use_A_tail = (coord_x == gi.m_tiles - 1) && (m_tail != 0);- const int a_box_h = use_A_tail ? m_tail : BLOCK_M;- const int a_bytes = a_box_h * BLOCK_K / 2;- const int tma_expect_bytes = a_bytes + B_SIZE + SF_STAGE;+ // === TMA WARP ===+ if (warp_id == NUM_WARPS - 2 && elect_sync()) {+ int tma_stage = 0;+ int mma_wait_phase = 1;+ int total_produced = 0;- const void *A_tmap = static_cast<const void *>(- &(use_A_tail ? tmap_pack_g8.A_tail[gidx] : tmap_pack_g8.A_full[gidx]));- const void *B_tmap = static_cast<const void *>(&tmap_pack_g8.B[gidx]);- const void *SFA_tmap = static_cast<const void *>(&tmap_pack_g8.SFA[gidx]);- const void *SFB_tmap = static_cast<const void *>(&tmap_pack_g8.SFB[gidx]);- // if (warp_id == 0 && lane_id == 0) {- // fence_proxy_tensormap(A_tmap);- // fence_proxy_tensormap(B_tmap);- // fence_proxy_tensormap(SFA_tmap);- // fence_proxy_tensormap(SFB_tmap);- // }- //__syncthreads();+ for (int tile = 0; tile < my_count; tile++) {+ const int lut_idx = worker_start + tile;+ const int gidx = static_cast<int>(params.lut_gidx[lut_idx]);+ const int tile_num_k = static_cast<int>(params.lut_mma_n_num_k[lut_idx] >> 8);+ const int off_n_cta = cta_rank == 0+ ? static_cast<int>(params.lut_off_n_cta_r0[lut_idx])+ : static_cast<int>(params.lut_off_n_cta_r1[lut_idx]);+ const int coord_n_cta = cta_rank == 0+ ? static_cast<int>(params.lut_coord_n_cta_r0[lut_idx])+ : static_cast<int>(params.lut_coord_n_cta_r1[lut_idx]);+ const int off_m_b = cta_rank == 0+ ? static_cast<int>(params.lut_off_m_b_r0[lut_idx])+ : static_cast<int>(params.lut_off_m_b_r1[lut_idx]);+ const int coord_m = static_cast<int>(params.lut_coord_m[lut_idx]);+ const int tma_expect_bytes = cta_rank == 0+ ? static_cast<int>(params.lut_expect_bytes_r0[lut_idx])+ : static_cast<int>(params.lut_expect_bytes_r1[lut_idx]);- // TMA producer warp.- if (warp_id == NUM_WARPS - 2 && elect_sync()) {- #pragma unroll- for (int ik = 0; ik < NUM_STAGES && ik < num_k; ik++) {- int s = ik;- int A_s = smem_main + s * MAIN_STAGE;- int B_s = A_s + A_SIZE;- int SFA_s = smem_sf + s * SF_STAGE;- int SFB_s = SFA_s + SFA_SIZE;+ const uint8_t tmap_sel = params.lut_tmap_sel[lut_idx];+ const int a_tmap_idx = cta_rank == 0 ? (tmap_sel & 3) : ((tmap_sel >> 2) & 3);+ const int b_tmap_idx = cta_rank == 0 ? ((tmap_sel >> 4) & 3) : ((tmap_sel >> 6) & 3);- tma_load_3d(A_s, A_tmap, 0, off_m, ik, tma_mbar + s * 8, cache_A);- tma_load_3d(B_s, B_tmap, 0, off_n, ik, tma_mbar + s * 8, cache_B);+ const void* A_tmap = get_a_tmap(gidx, a_tmap_idx);+ const void* B_tmap = get_b_tmap(gidx, b_tmap_idx);+ const void* SFA_tmap = static_cast<const void*>(&tmap_pack_g8.SFA[gidx]);+ const void* SFB_tmap = static_cast<const void*>(&tmap_pack_g8.SFB[gidx]);- int z_sf = ik * SF_K_PER_BLOCK;- tma_load_3d(SFA_s, SFA_tmap, 0, coord_x, z_sf, tma_mbar + s * 8, cache_SF);- tma_load_3d(SFB_s, SFB_tmap, 0, coord_y, z_sf, tma_mbar + s * 8, cache_SF);+ #pragma unroll 1+ for (int ik = 0; ik < tile_num_k; ik++) {+ if (total_produced >= NS) {+ mbarrier_wait(mma_mbar + tma_stage * 8, mma_wait_phase);+ }- mbarrier_arrive_expect_tx(tma_mbar + s * 8, tma_expect_bytes);- }+ const int mbar_addr = (tma_mbar + tma_stage * 8) & 0xFEFFFFFF;+ int A_s = smem_main + tma_stage * main_stg;+ int SFA_s = smem_sf + tma_stage * SF_STAGE;- for (int ik = NUM_STAGES; ik < num_k; ik++) {- int s = ik % NUM_STAGES;- mbarrier_wait(mma_mbar + s * 8, (ik / NUM_STAGES - 1) % 2);+ tma_3d_gmem2smem(A_s + A_SIZE, B_tmap, 0, off_m_b, ik, mbar_addr, cache_B);+ tma_3d_gmem2smem(A_s, A_tmap, 0, off_n_cta, ik, mbar_addr, cache_A);+ const int z_sf = ik * SF_K_PER_BLOCK_L;+ tma_3d_gmem2smem(SFA_s, SFA_tmap, 0, coord_n_cta, z_sf, mbar_addr, cache_SF);+ tma_3d_gmem2smem(SFA_s + SFA_SIZE, SFB_tmap, 0, coord_m, z_sf, mbar_addr, cache_SF);+ mbarrier_arrive_expect_tx(mbar_addr, tma_expect_bytes);- int A_s = smem_main + s * MAIN_STAGE;- int B_s = A_s + A_SIZE;- int SFA_s = smem_sf + s * SF_STAGE;- int SFB_s = SFA_s + SFA_SIZE;+ total_produced++;+ tma_stage++;+ if (tma_stage == NS) {+ tma_stage = 0;+ mma_wait_phase ^= 1;+ }+ }+ }+ }- tma_load_3d(A_s, A_tmap, 0, off_m, ik, tma_mbar + s * 8, cache_A);- tma_load_3d(B_s, B_tmap, 0, off_n, ik, tma_mbar + s * 8, cache_B);+ // === MMA WARP (CTA0 only) ===+ if (cta_rank == 0 && warp_id == NUM_WARPS - 1 && elect_sync()) {+ int mma_stage = 0;+ int tma_wait_phase = 0;- int z_sf = ik * SF_K_PER_BLOCK;- tma_load_3d(SFA_s, SFA_tmap, 0, coord_x, z_sf, tma_mbar + s * 8, cache_SF);- tma_load_3d(SFB_s, SFB_tmap, 0, coord_y, z_sf, tma_mbar + s * 8, cache_SF);+ for (int tile = 0; tile < my_count; tile++) {+ const int slot = tile & 1;+ const int d_tmem_base = (slot == 0) ? D_TMEM0 : D_TMEM1;+ const int done_mbar = (slot == 0) ? done_mbar0 : done_mbar1;- mbarrier_arrive_expect_tx(tma_mbar + s * 8, tma_expect_bytes);+ // Do not reuse a TMEM slot before EP drains it on both CTAs.+ if (tile >= 2) {+ const int ep_wait_phase = (((tile >> 1) - 1) & 1);+ mbarrier_wait((slot == 0) ? ep_mbar0 : ep_mbar1, ep_wait_phase);}- }- // MMA consumer warp.- if (warp_id == NUM_WARPS - 1 && elect_sync()) {+ const int lut_idx = worker_start + tile;+ const int mma_n = static_cast<int>(params.lut_mma_n_num_k[lut_idx] & 0xFF);+ const int tile_num_k = static_cast<int>(params.lut_mma_n_num_k[lut_idx] >> 8);+ const uint32_t i_desc = (1U << 7U) | (1U << 10U) |+ (((uint32_t)mma_n >> 3U) << 17U) | (((uint32_t)MMA_M_2CTA >> 7U) << 27U);+#pragma unroll 1- for (int ik = 0; ik < num_k; ik++) {- int s = ik % NUM_STAGES;- mbarrier_wait(tma_mbar + s * 8, (ik / NUM_STAGES) % 2);+ for (int ik = 0; ik < tile_num_k; ik++) {+ mbarrier_wait(tma_mbar + mma_stage * 8, tma_wait_phase);- int A_s = smem_main + s * MAIN_STAGE;+ int A_s = smem_main + mma_stage * main_stg;int B_s = A_s + A_SIZE;- int SFA_s = smem_sf + s * SF_STAGE;+ int SFA_s = smem_sf + mma_stage * SF_STAGE;int SFB_s = SFA_s + SFA_SIZE;- constexpr uint64_t sf_base = desc_encode(0) | (desc_encode(8 * 16) << 32ULL) | (1ULL << 46ULL);+ constexpr uint64_t sf_base = desc_encode(0) | (desc_encode(8*16) << 32ULL) | (1ULL << 46ULL);uint64_t sfa_desc = sf_base + ((uint64_t)SFA_s >> 4ULL);uint64_t sfb_desc = sf_base + ((uint64_t)SFB_s >> 4ULL);#pragma unroll- for (int k = 0; k < BLOCK_K / MMA_K; k++) {- tcgen05_cp_nvfp4(SFA_TMEM + k * 4, sfa_desc + (uint64_t)k * (512ULL >> 4ULL));- tcgen05_cp_nvfp4(SFB_TMEM + k * 4, sfb_desc + (uint64_t)k * (512ULL >> 4ULL));+ for (int kk = 0; kk < BLOCK_K / MMA_K; kk++) {+ tcgen05_cp_nvfp4(SFA_TMEM + kk * 4, sfa_desc + (uint64_t)kk * (512ULL >> 4ULL));+ tcgen05_cp_nvfp4(SFB_TMEM + kk * 4, sfb_desc + (uint64_t)kk * (512ULL >> 4ULL));}#pragma unrollfor (int k2 = 0; k2 < 256 / MMA_K; k2++) {uint64_t a_desc = make_desc_AB(A_s + k2 * 32);uint64_t b_desc = make_desc_AB(B_s + k2 * 32);+ // Reset accumulator at start of each tile (enable_d=0 clears accum)int enable_d = (ik == 0 && k2 == 0) ? 0 : 1;- tcgen05_mma_nvfp4(a_desc, b_desc, I_DESC,- SFA_TMEM + k2 * 4, SFB_TMEM + k2 * 4, enable_d, d_tmem_base);+ tcgen05_mma_nvfp4(d_tmem_base, a_desc, b_desc, i_desc,+ SFA_TMEM + k2 * 4, SFB_TMEM + k2 * 4, enable_d);}- tcgen05_commit(mma_mbar + s * 8);+ tcgen05_commit_mcast(mma_mbar + mma_stage * 8, cta_mask);++ mma_stage++;+ if (mma_stage == NS) {+ mma_stage = 0;+ tma_wait_phase ^= 1;+ }}- tcgen05_commit(done_mbar);+ tcgen05_commit_mcast(done_mbar, cta_mask);}-- // Overlap epilogue for previous tile with compute on current tile.- if (warp_id < NUM_EP_WARPS && tile_iter > 0) {- EpMeta meta = ep_meta[prev_slot];- do_epilogue(warp_id, lane_id, prev_done_mbar, prev_d_tmem_base,- meta.c_ptr, meta.M, meta.N, meta.off_m, meta.off_n);- }- __syncthreads();}- // Drain last tile epilogue.+ // === EP WARPS (both CTAs) ===if (warp_id < NUM_EP_WARPS) {- int final_slot = (my_count - 1) & 1;- int final_done_mbar = final_slot ? done_mbar1 : done_mbar0;- int final_d_tmem_base = final_slot ? D_TMEM1 : D_TMEM0;- EpMeta meta = ep_meta[final_slot];- do_epilogue(warp_id, lane_id, final_done_mbar, final_d_tmem_base,- meta.c_ptr, meta.M, meta.N, meta.off_m, meta.off_n);+ for (int tile = 0; tile < my_count; tile++) {+ const int slot = tile & 1;+ const int done_mbar = (slot == 0) ? done_mbar0 : done_mbar1;+ const int done_phase = (tile >> 1) & 1;+ const int d_tmem_base = (slot == 0) ? D_TMEM0 : D_TMEM1;++ const int lut_idx = worker_start + tile;+ const int gidx = static_cast<int>(params.lut_gidx[lut_idx]);+ const GroupInfo& gi = params.groups[gidx];+ const int M = gi.M;+ const int N = gi.N;+ const int off_m = static_cast<int>(params.lut_coord_m[lut_idx]) * BLOCK_N;+ const int mma_n = static_cast<int>(params.lut_mma_n_num_k[lut_idx] & 0xFF);+ const int off_n_r0 = static_cast<int>(params.lut_off_n_cta_r0[lut_idx]);+ const int off_n_r1 = static_cast<int>(params.lut_off_n_cta_r1[lut_idx]);+ const bool rank1_aliased = (off_n_r1 == off_n_r0);+ const int off_n = (cta_rank == 0) ? off_n_r0 : off_n_r1;++ if (!(cta_rank == 1 && rank1_aliased)) {+ do_epilogue_transposed(+ warp_id, lane_id, cta_rank,+ done_mbar, done_phase, d_tmem_base,+ smem_ep, gi.c_ptr, M, N, off_m, off_n, mma_n);+ } else {+ mbarrier_wait(done_mbar, done_phase);+ }++ if (warp_id == 0 && elect_sync()) {+ tcgen05_commit_mcast((slot == 0) ? ep_mbar0 : ep_mbar1, cta_mask);+ }+ }}- __syncthreads();+ // === CLEANUP ===+ asm volatile("barrier.cluster.arrive.relaxed.aligned;");+ asm volatile("barrier.cluster.wait.acquire.aligned;");+if (warp_id == 0)- asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));+ asm volatile("tcgen05.dealloc.cta_group::2.sync.aligned.b32 %0, %1;" :: "r"(0), "r"(TMEM_COLS));}template <int SCHEDULE_ID, int PROFILE_ID>inline void launch_grouped_kernel(const KernelParams& params, const TmapParamPackG8& tmap_pack_g8, int smem_size) {- cudaFuncSetAttribute(grouped_gemm_kernel<SCHEDULE_ID, PROFILE_ID>,- cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);- cudaFuncSetAttribute(grouped_gemm_kernel<SCHEDULE_ID, PROFILE_ID>,- cudaFuncAttributePreferredSharedMemoryCarveout, cudaSharedmemCarveoutMaxShared);- grouped_gemm_kernel<SCHEDULE_ID, PROFILE_ID><<<params.launch_ctas, TB_SIZE, smem_size>>>(params, tmap_pack_g8);+ auto kernel = grouped_gemm_kernel<SCHEDULE_ID, PROFILE_ID>;+ cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);+ cudaFuncSetAttribute(kernel, cudaFuncAttributePreferredSharedMemoryCarveout, cudaSharedmemCarveoutMaxShared);+ cudaFuncSetAttribute(kernel, cudaFuncAttributeNonPortableClusterSizeAllowed, 1);++ cudaLaunchConfig_t launch_config = {};+ launch_config.gridDim = params.launch_ctas;+ launch_config.blockDim = TB_SIZE;+ launch_config.dynamicSmemBytes = smem_size;++ cudaLaunchAttribute cluster_attr = {};+ cluster_attr.id = cudaLaunchAttributeClusterDimension;+ cluster_attr.val.clusterDim.x = 2;+ cluster_attr.val.clusterDim.y = 1;+ cluster_attr.val.clusterDim.z = 1;+ launch_config.attrs = &cluster_attr;+ launch_config.numAttrs = 1;++ cudaLaunchKernelEx(&launch_config, kernel, params, tmap_pack_g8);}// ============================================================================⋯ 12 unchanged linesKernelParams params = {};params.num_groups = G;- int total_tiles = 0;+ static int smem_size = 0;+ static int smem_avail = 0;+ if (!smem_size) {+ int dev; cudaGetDevice(&dev);+ int smem_max;+ cudaDeviceGetAttribute(&smem_max, cudaDevAttrMaxSharedMemoryPerBlockOptin, dev);+ smem_size = smem_max - 1024;+ smem_avail = smem_size - EP_SMEM_BYTES;+ TORCH_CHECK(smem_avail > 0, "Insufficient shared memory for EP scratch");+ }+ params.smem_size_bytes = smem_size;++ // Choose NS: largest value <= MAX_NS that fits in smem (independent of num_k divisibility).+ // For stage sizing, use the worst case (full tiles): main_stg = A_SIZE + B_HALF+ // First pass: find max B payload across all tiles to determine stage size+ int max_b_half = B_HALF; // full tile+ // Actually all tiles fit B_HALF at most, tail tiles use less.+ // Use full B_HALF for stage size (wastes some smem on tail tiles but keeps layout uniform).+ int main_stg = A_SIZE + max_b_half;+ params.main_stg = main_stg;++ // Find best NS: largest that fits in smem.+ int best_ns = 1;+ for (int ns = MAX_NS; ns >= 1; ns--) {+ int total_smem = ns * main_stg + ns * SF_STAGE;+ if (total_smem <= smem_avail) {+ best_ns = ns;+ break;+ }+ }+ params.ns = best_ns;++ int raw_total_tiles = 0;for (int g = 0; g < G; g++) {int Mi = A_list[g].size(0);int Ki = A_list[g].size(1) * 2;int Ni = B_list[g].size(0);- int mt = (Mi + BLOCK_M - 1) / BLOCK_M;- int nt = (Ni + BLOCK_N - 1) / BLOCK_N;- params.groups[g] = {(half *)C_list[g].data_ptr(), Mi, Ni, Ki, total_tiles, mt, nt};- total_tiles += mt * nt;++ int nt = (Ni + 255) / 256;+ int mt = (Mi + BLOCK_N - 1) / BLOCK_N;++ params.groups[g] = {(half *)C_list[g].data_ptr(), Mi, Ni, Ki};+ raw_total_tiles += mt * nt;}- params.total_tiles = total_tiles;- // Heuristic cap: moderate tile counts often benefit from deeper per-CTA pipelines.- const int cap_ctas = (total_tiles > 128 && total_tiles <= 384) ? 128 : MAX_LAUNCH_CTAS;- params.launch_ctas = total_tiles < cap_ctas ? total_tiles : cap_ctas;- // Benchmark-specialized profile selection with generic fallback.+ const int cap_clusters = MAX_CLUSTERS;+ int num_clusters = raw_total_tiles < cap_clusters ? raw_total_tiles : cap_clusters;++ const bool is_bench1 = (params.num_groups == 8 && raw_total_tiles == 176);+ const bool is_bench2 = (params.num_groups == 8 && raw_total_tiles == 364);+ if (is_bench1 && V_BENCH1_CLUSTERS > 0) num_clusters = V_BENCH1_CLUSTERS;+ if (is_bench2 && V_BENCH2_CLUSTERS > 0) num_clusters = V_BENCH2_CLUSTERS;+ if (!is_bench1 && !is_bench2 && V_GENERIC_CLUSTERS > 0) num_clusters = V_GENERIC_CLUSTERS;+ if (num_clusters > MAX_CLUSTERS) num_clusters = MAX_CLUSTERS;+ if (num_clusters > raw_total_tiles) num_clusters = raw_total_tiles;+ if (num_clusters < 1) num_clusters = 1;++ params.launch_ctas = num_clusters * 2;+int bench_profile = PROFILE_GENERIC_G8;- if (params.num_groups == 8 && params.total_tiles == 352) {+ if (params.num_groups == 8 && raw_total_tiles == 176) {bench_profile = PROFILE_BENCH1;- } else if (params.num_groups == 8 && params.total_tiles == 728) {+ } else if (params.num_groups == 8 && raw_total_tiles == 364) {bench_profile = PROFILE_BENCH2;- } else if (params.num_groups == 2 && params.total_tiles == 120) {+ } else if (params.num_groups == 2 && raw_total_tiles == 60) {bench_profile = PROFILE_BENCH3;- } else if (params.num_groups == 2 && params.total_tiles == 128) {+ } else if (params.num_groups == 2 && raw_total_tiles == 64) {bench_profile = PROFILE_BENCH4;} else if (params.num_groups == 2) {bench_profile = PROFILE_GENERIC_G2;⋯ 25 unchanged linesTmapParamPackG8 tmap_pack_g8 = {};- // O(G) pre-scan: all benchmarks have uniform N,K across groupsbool uniform_nk = true;for (int g = 1; g < G; g++) {if (B_list[g].size(0) != B_list[0].size(0) || A_list[g].size(1) != A_list[0].size(1)) {- uniform_nk = false; break;+ uniform_nk = false;+ break;}}- // SFA reuse LUT: mn_blocks -> first group with that value (no nested loop)- int sfa_src[MAX_GROUPS];+ int sfb_src[MAX_GROUPS];+ for (int g = 0; g < MAX_GROUPS; g++) sfb_src[g] = -1;if (uniform_nk) {int mn_first[4] = {-1, -1, -1, -1};for (int g = 0; g < G; g++) {int mnb = (A_list[g].size(0) + 127) / 128;- sfa_src[g] = (mnb < 4) ? mn_first[mnb] : -1;+ sfb_src[g] = (mnb < 4) ? mn_first[mnb] : -1;if (mnb < 4 && mn_first[mnb] < 0) mn_first[mnb] = g;}}⋯ 2 unchanged linesint Mi = A_list[g].size(0);int Ki = A_list[g].size(1) * 2;int Ni = B_list[g].size(0);- int m_tail = Mi % BLOCK_M;- init_AB_tmap(&tmap_pack_g8.A_full[g], (const char *)A_list[g].data_ptr(), Mi, Ki, BLOCK_M, BLOCK_K, ab_l2_promotion);+ int n_tail = Ni % BLOCK_M;+ if (g > 0 && uniform_nk) {+ tmap_pack_g8.A_full[g] = tmap_pack_g8.A_full[0];+ check_cu(cuTensorMapReplaceAddress(&tmap_pack_g8.A_full[g], (void *)B_list[g].data_ptr()));+ if (n_tail == 0) {+ tmap_pack_g8.A_tail[g] = tmap_pack_g8.A_full[g];+ } else {+ tmap_pack_g8.A_tail[g] = tmap_pack_g8.A_tail[0];+ check_cu(cuTensorMapReplaceAddress(&tmap_pack_g8.A_tail[g], (void *)B_list[g].data_ptr()));+ }+ } else {+ init_AB_tmap(&tmap_pack_g8.A_full[g], (const char *)B_list[g].data_ptr(), Ni, Ki, BLOCK_M, BLOCK_K, ab_l2_promotion);+ if (n_tail == 0) {+ tmap_pack_g8.A_tail[g] = tmap_pack_g8.A_full[g];+ } else {+ init_AB_tmap(&tmap_pack_g8.A_tail[g], (const char *)B_list[g].data_ptr(), Ni, Ki, n_tail, BLOCK_K, ab_l2_promotion);+ }+ }++ int m_tail = Mi % BLOCK_N;+ int mma_n_tail = (m_tail == 0) ? BLOCK_N : ((m_tail + 15) & ~15);+ int b_full_box_h = BLOCK_N / 2;+ int b_tail_half = mma_n_tail / 2;+ int b_tail0_box_h = m_tail < b_tail_half ? m_tail : b_tail_half;+ int b_tail1_box_h = m_tail - b_tail_half;+ if (b_tail1_box_h < 0) b_tail1_box_h = 0;+ if (b_tail1_box_h > b_tail_half) b_tail1_box_h = b_tail_half;+ init_AB_tmap(&tmap_pack_g8.B_full[g], (const char *)A_list[g].data_ptr(), Mi, Ki, b_full_box_h, BLOCK_K, ab_l2_promotion);if (m_tail == 0) {- tmap_pack_g8.A_tail[g] = tmap_pack_g8.A_full[g];+ tmap_pack_g8.B_tail0[g] = tmap_pack_g8.B_full[g];+ tmap_pack_g8.B_tail1[g] = tmap_pack_g8.B_full[g];} else {- init_AB_tmap(&tmap_pack_g8.A_tail[g], (const char *)A_list[g].data_ptr(), Mi, Ki, m_tail, BLOCK_K, ab_l2_promotion);+ init_AB_tmap(&tmap_pack_g8.B_tail0[g], (const char *)A_list[g].data_ptr(), Mi, Ki, b_tail0_box_h, BLOCK_K, ab_l2_promotion);+ const int b_tail1_box_safe = b_tail1_box_h > 0 ? b_tail1_box_h : 1;+ init_AB_tmap(&tmap_pack_g8.B_tail1[g], (const char *)A_list[g].data_ptr(), Mi, Ki, b_tail1_box_safe, BLOCK_K, ab_l2_promotion);}if (g > 0 && uniform_nk) {- tmap_pack_g8.B[g] = tmap_pack_g8.B[0];- check_cu(cuTensorMapReplaceAddress(&tmap_pack_g8.B[g], (void *)B_list[g].data_ptr()));- tmap_pack_g8.SFB[g] = tmap_pack_g8.SFB[0];- check_cu(cuTensorMapReplaceAddress(&tmap_pack_g8.SFB[g], (void *)SFB_list[g].data_ptr()));+ tmap_pack_g8.SFA[g] = tmap_pack_g8.SFA[0];+ check_cu(cuTensorMapReplaceAddress(&tmap_pack_g8.SFA[g], (void *)SFB_list[g].data_ptr()));} else {- init_AB_tmap(&tmap_pack_g8.B[g], (const char *)B_list[g].data_ptr(), Ni, Ki, BLOCK_N, BLOCK_K, ab_l2_promotion);- init_SF_tmap(&tmap_pack_g8.SFB[g], (const char *)SFB_list[g].data_ptr(), Ni, Ki, sf_l2_promotion);+ init_SF_tmap(&tmap_pack_g8.SFA[g], (const char *)SFB_list[g].data_ptr(), Ni, Ki, sf_l2_promotion);}- if (uniform_nk && sfa_src[g] >= 0) {- tmap_pack_g8.SFA[g] = tmap_pack_g8.SFA[sfa_src[g]];- check_cu(cuTensorMapReplaceAddress(&tmap_pack_g8.SFA[g], (void *)SFA_list[g].data_ptr()));+ if (uniform_nk && sfb_src[g] >= 0) {+ tmap_pack_g8.SFB[g] = tmap_pack_g8.SFB[sfb_src[g]];+ check_cu(cuTensorMapReplaceAddress(&tmap_pack_g8.SFB[g], (void *)SFA_list[g].data_ptr()));} else {- init_SF_tmap(&tmap_pack_g8.SFA[g], (const char *)SFA_list[g].data_ptr(), Mi, Ki, sf_l2_promotion);+ init_SF_tmap(&tmap_pack_g8.SFB[g], (const char *)SFA_list[g].data_ptr(), Mi, Ki, sf_l2_promotion);}}- static int smem_size = 0;- if (!smem_size) {- int dev; cudaGetDevice(&dev);- int smem_max;- cudaDeviceGetAttribute(&smem_max, cudaDevAttrMaxSharedMemoryPerBlockOptin, dev);⋯ diff truncated
scrolls · 1201 diff lines total
Best evidence level for this revision: reported
JSON