submission 499566
mufeez-amjad · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1385 lines, June 9 Researcher Reciprocity License v1.0.
v4c.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-499566?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:3e36a0abc7cdba51fd05db2f550adbcbe6223f99ed049b54b4f40745291390bc
license declaredunknown
license concludedunknown
authorsmufeez-amjad
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
template <int BLOCK_M, int BLOCK_N, bool LOW_M_EPILOGUE>mbarrier
__device__ inline void mbarrier_init(int mbar_addr, int count) {num-warps = 8
constexpr int TMA_NUM_WARPS = 8;persistent-kernel
template <bool PERSISTENT, int BLOCK_M, int BLOCK_N, int NS, int CLUSTER_SIZE = 1>shared-memory
__device__ inline void tma_3d_gmem2smem_multicast(int dst, const void *tmap_ptr, int x, int y, int z,tcgen05
asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));tile-k = 256
constexpr int TMA_BLOCK_K = 256;tile-n = 128
constexpr int TMA_BLOCK_N = 128;tma
CUtensorMap A_tmap;vector-width = half2
half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset + col_base);Kernel source
v4c.py1385 lines
#!POPCORN leaderboard nvfp4_group_gemm
from __future__ import annotations
from functools import lru_cache
from typing import cast
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
"""
g: 8; k: [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168]; m: [80, 176, 128, 72, 64, 248, 96, 160]; n: [4096, 4096, 4096, 4096, 4096, 4096, 4096, 4096]; seed: 1111
⏱ 47.5 ± 0.00 µs
⚡ 47.4 µs 🐌 47.5 µs
g: 8; k: [2048, 2048, 2048, 2048, 2048, 2048, 2048, 2048]; m: [40, 76, 168, 72, 164, 148, 196, 160]; n: [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168]; seed: 1111
⏱ 45.0 ± 0.04 µs
⚡ 44.6 µs 🐌 45.3 µs
g: 2; k: [4096, 4096]; m: [192, 320]; n: [3072, 3072]; seed: 1111
⏱ 14.3 ± 0.01 µs
⚡ 13.9 µs 🐌 14.6 µs
g: 2; k: [1536, 1536]; m: [128, 384]; n: [4096, 4096]; seed: 1111
⏱ 10.5 ± 0.01 µs
⚡ 10.4 µs 🐌 10.7 µs
"""
CUDA_SRC = """
#include <vector>
#include <unordered_map>
#include <cstdint>
#include <cuda.h>
#include <cudaTypedefs.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>
static inline int ceil_div(int a, int b) { return (a + b - 1) / b; }
#define CUDA_CHECK(expr) \\
do { \\
cudaError_t _err = (expr); \\
TORCH_CHECK(_err == cudaSuccess, "CUDA error: ", cudaGetErrorString(_err)); \\
} while (0)
static inline uint64_t hash_combine_u64(uint64_t h, uint64_t x) {
h ^= x;
h *= 1099511628211ULL;
return h;
}
constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000;
struct WorkItem {
int problem_idx;
int tile_m;
int tile_n;
};
// Per-problem metadata.
// Align on 128 byte boundary, useful since this is read by many CTAs.
struct __align__(128) ProblemInfo {
CUtensorMap A_tmap;
CUtensorMap B_tmap;
CUtensorMap B_tmap_256;
const char* SFA_ptr;
const char* SFB_ptr;
half* C_ptr;
int M, N, K;
int64_t Cs0, Cs1;
};
// tcgen05 descriptors encode shared-memory addresses in 16-byte units.
// Mask to the HW-supported address width and drop the 16B alignment bits.
__device__ inline constexpr uint64_t desc_encode(uint64_t x) {
return (x & 0x3'FFFFULL) >> 4ULL;
}
// elect.sync: use this to have a single lane issue TMA/tcgen05 instructions while the
// whole warp stays converged.
__device__ 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;
asm volatile("mov.u32 %0, %%cluster_ctarank;" : "=r"(rank));
return rank;
}
__device__ inline void cluster_sync() {
asm volatile("barrier.cluster.arrive;" ::: "memory");
asm volatile("barrier.cluster.wait;" ::: "memory");
}
// Shared-memory mbarrier helpers.
// Used for:
// - TMA completion barrier: consumer waits for bytes to arrive in shared memory.
// - Stage reuse barrier: producer waits until MMA is done with a stage before
// overwriting it.
__device__ inline void mbarrier_init(int mbar_addr, int count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
}
// Program the expected byte count for a TMA stage and arrive.
// This must happen before issuing any cp.async.bulk.* that completes to the
// barrier.
__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;"
:: "r"(mbar_addr), "r"(size) : "memory");
}
__device__ void mbarrier_wait(int mbar_addr, int phase) {
uint32_t ticks = 0x989680;
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)
);
}
// TMA: 3D tensor-map load from global -> shared memory.
// The (x,y,z) coordinates correspond to the CUtensorMap encoding in
// init_AB_tmap_u4.
template <int CTA_GROUP = 1>
__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 int mapa_cta_to_cluster(int cta_addr, int dest_cta) {
int cluster_addr;
asm volatile("mapa.shared::cluster.u32 %0, %1, %2;" : "=r"(cluster_addr) : "r"(cta_addr), "r"(dest_cta));
return cluster_addr;
}
__device__ inline void mbarrier_arrive_cluster(int mbar_cluster_addr) {
asm volatile("mbarrier.arrive.shared::cluster.b64 _, [%0];"
:: "r"(mbar_cluster_addr) : "memory");
}
// Cluster multicast variant of 3D TMA.
// dst and mbar_addr are in shared::cluster address space.
// The multicast mask selects which CTAs in the cluster receive the data.
__device__ inline void tma_3d_gmem2smem_multicast(int dst, const void *tmap_ptr, int x, int y, int z,
int mbar_addr, uint16_t multicast_mask) {
asm volatile(
"cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.cta_group::1 "
"[%0], [%1, {%2, %3, %4}], [%5], %6;"
:: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "h"(multicast_mask)
: "memory"
);
}
// Linear bulk copy global -> shared.
// Used for scale-factor tensors (SFA/SFB).
__device__ inline void tma_gmem2smem(int dst, const void *src, int size, int mbar_addr, uint64_t cache_policy) {
asm volatile(
"cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "
"[%0], [%1], %2, [%3], %4;"
:: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy)
: "memory"
);
}
__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));
}
__device__ inline void tcgen05_commit(int mbar_addr) {
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];\\n"
:: "r"(mbar_addr) : "memory"
);
}
__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
) {
const int d_tmem = 0;
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"
"}"
:: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
"r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d)
);
}
struct SHAPE {
static constexpr char _16x256b[] = ".16x256b";
};
struct NUM {
static constexpr char x16[] = ".x16";
};
template <const char *SHAPE, const char *NUM>
__device__ inline
void tcgen05_ld_64regs(float *tmp, int row, int col) {
asm volatile("tcgen05.ld.sync.aligned%65%66.b32 "
"{ %0, %1, %2, %3, %4, %5, %6, %7, "
" %8, %9, %10, %11, %12, %13, %14, %15, "
" %16, %17, %18, %19, %20, %21, %22, %23, "
" %24, %25, %26, %27, %28, %29, %30, %31, "
" %32, %33, %34, %35, %36, %37, %38, %39, "
" %40, %41, %42, %43, %44, %45, %46, %47, "
" %48, %49, %50, %51, %52, %53, %54, %55, "
" %56, %57, %58, %59, %60, %61, %62, %63}, [%64];"
: "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
"=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
"=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
"=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31]),
"=f"(tmp[32]), "=f"(tmp[33]), "=f"(tmp[34]), "=f"(tmp[35]), "=f"(tmp[36]), "=f"(tmp[37]), "=f"(tmp[38]), "=f"(tmp[39]),
"=f"(tmp[40]), "=f"(tmp[41]), "=f"(tmp[42]), "=f"(tmp[43]), "=f"(tmp[44]), "=f"(tmp[45]), "=f"(tmp[46]), "=f"(tmp[47]),
"=f"(tmp[48]), "=f"(tmp[49]), "=f"(tmp[50]), "=f"(tmp[51]), "=f"(tmp[52]), "=f"(tmp[53]), "=f"(tmp[54]), "=f"(tmp[55]),
"=f"(tmp[56]), "=f"(tmp[57]), "=f"(tmp[58]), "=f"(tmp[59]), "=f"(tmp[60]), "=f"(tmp[61]), "=f"(tmp[62]), "=f"(tmp[63])
: "r"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}
__device__ inline void tcgen05_ld_16x256b_x16(float *tmp, int row, int col) {
tcgen05_ld_64regs<SHAPE::_16x256b, NUM::x16>(tmp, row, col);
}
__device__ __forceinline__ void tcgen05_dealloc_cols_cta1(uint32_t tmem, int count) {
asm volatile(
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\\n"
:: "r"(tmem), "r"(count)
: "memory"
);
}
constexpr int TMA_BLOCK_N = 128;
constexpr int TMA_BLOCK_K = 256;
constexpr int TMA_NUM_WARPS = 8;
constexpr int MMA_M = 128;
constexpr int LOW_M_THRESHOLD = 96;
constexpr int TMA_WARP = 4;
// Use a second (otherwise idle) warp to issue B/SFB TMA in parallel.
// This cuts producer-side latency and reduces MMA-side barrier stalls.
constexpr int TMA_WARP_B = 6;
constexpr int MMA_WARP = 5;
// Epilogue: read fp32 accumulators from TMEM (tcgen05.ld) and store fp16 C.
// Only threads with tid < BLOCK_M participate; this maps 4 warps (0..3) to the
// 128 output rows, with each warp handling a 32-row stripe.
template <int BLOCK_M, int BLOCK_N, bool LOW_M_EPILOGUE>
__device__ __forceinline__ void epilogue_store(
const ProblemInfo& prob,
int m_offset,
int n_offset,
int tid,
int warp_id,
int lane_id
) {
if (tid >= BLOCK_M) return;
const int M = prob.M;
const int N = prob.N;
half* C_ptr = prob.C_ptr;
const int64_t Cs0 = prob.Cs0;
const int64_t Cs1 = prob.Cs1;
const bool full_n = (n_offset + BLOCK_N <= N);
const bool full_m = (m_offset + BLOCK_M <= M);
const bool full_tile = full_n && full_m;
const bool contiguous = (Cs1 == 1);
const int warp_row_base = m_offset + warp_id * 32;
if (LOW_M_EPILOGUE && warp_row_base >= M) return;
int m_iters = 2;
if (LOW_M_EPILOGUE) {
const int remaining = M - warp_row_base;
m_iters = (remaining <= 16) ? 1 : 2;
}
// Lane mapping: each lane owns two columns (half2) and one of 8 rows.
const int lane_row = lane_id >> 2;
const int lane_col = (lane_id & 3) * 2;
// We load/store in 128-column halves so tcgen05.ld has a fixed shape.
constexpr int HALF_N = 128;
const int halves = BLOCK_N / HALF_N;
for (int m = 0; m < m_iters; ++m) {
for (int half_idx = 0; half_idx < halves; ++half_idx) {
float tmp[HALF_N / 2];
const int col_base = half_idx * HALF_N;
// TMEM coordinates are relative to the CTA's output tile.
tcgen05_ld_16x256b_x16(tmp, warp_id * 32 + m * 16, col_base);
asm volatile("tcgen05.wait::ld.sync.aligned;\\n");
const int row0 = warp_row_base + m * 16 + lane_row;
const int row1 = row0 + 8;
if (contiguous) {
if (full_tile) {
half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset + col_base);
half2* row1_ptr = reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset + col_base);
#pragma unroll
for (int i = 0; i < HALF_N / 8; i++) {
const int idx = i * 4;
const int col = i * 8 + lane_col;
const int h2_idx = col >> 1;
row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
}
continue;
}
const bool row0_in = row0 < M;
const bool row1_in = row1 < M;
if (full_n) {
half2* row0_ptr = row0_in ? reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset + col_base) : nullptr;
half2* row1_ptr = row1_in ? reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset + col_base) : nullptr;
#pragma unroll
for (int i = 0; i < HALF_N / 8; i++) {
const int idx = i * 4;
const int col = i * 8 + lane_col;
const int h2_idx = col >> 1;
if (row0_in) {
row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
}
if (row1_in) {
row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
}
}
} else {
#pragma unroll
for (int i = 0; i < HALF_N / 8; i++) {
const int idx = i * 4;
const int col = n_offset + col_base + i * 8 + lane_col;
if (col < N) {
const half2 h2_row0 = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
const half2 h2_row1 = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
if (row0_in) {
if (col + 1 < N) {
reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + col)[0] = h2_row0;
} else {
C_ptr[row0 * Cs0 + col] = __low2half(h2_row0);
}
}
if (row1_in) {
if (col + 1 < N) {
reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + col)[0] = h2_row1;
} else {
C_ptr[row1 * Cs0 + col] = __low2half(h2_row1);
}
}
}
}
}
} else {
const bool row0_in = row0 < M;
const bool row1_in = row1 < M;
#pragma unroll
for (int i = 0; i < HALF_N / 8; i++) {
const int idx = i * 4;
const int col = n_offset + col_base + i * 8 + lane_col;
if (col < N) {
const half2 h2_row0 = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
const half2 h2_row1 = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
const half h00 = __low2half(h2_row0);
const half h01 = __high2half(h2_row0);
const half h10 = __low2half(h2_row1);
const half h11 = __high2half(h2_row1);
if (row0_in) {
C_ptr[row0 * Cs0 + col * Cs1] = h00;
if (col + 1 < N) C_ptr[row0 * Cs0 + (col + 1) * Cs1] = h01;
}
if (row1_in) {
C_ptr[row1 * Cs0 + col * Cs1] = h10;
if (col + 1 < N) C_ptr[row1 * Cs0 + (col + 1) * Cs1] = h11;
}
}
}
}
}
}
}
template <int BLOCK_M, int BLOCK_N>
__device__ __forceinline__ void epilogue_store_fulltile_contiguous(
const ProblemInfo& prob,
int m_offset,
int n_offset,
int tid,
int warp_id,
int lane_id
) {
if (tid >= BLOCK_M) return;
half* C_ptr = prob.C_ptr;
const int64_t Cs0 = prob.Cs0;
const int warp_row_base = m_offset + warp_id * 32;
// Lane mapping: each lane owns two columns (half2) and one of 8 rows.
const int lane_row = lane_id >> 2;
const int lane_col = (lane_id & 3) * 2;
constexpr int HALF_N = 128;
const int halves = BLOCK_N / HALF_N;
#pragma unroll
for (int m = 0; m < 2; ++m) {
const int row0 = warp_row_base + m * 16 + lane_row;
const int row1 = row0 + 8;
#pragma unroll
for (int half_idx = 0; half_idx < halves; ++half_idx) {
float tmp[HALF_N / 2];
const int col_base = half_idx * HALF_N;
tcgen05_ld_16x256b_x16(tmp, warp_id * 32 + m * 16, col_base);
asm volatile("tcgen05.wait::ld.sync.aligned;\\n");
half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset + col_base);
half2* row1_ptr = reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset + col_base);
#pragma unroll
for (int i = 0; i < HALF_N / 8; i++) {
const int idx = i * 4;
const int col = i * 8 + lane_col;
const int h2_idx = col >> 1;
row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
}
}
}
}
template <bool PERSISTENT, int BLOCK_M, int BLOCK_N, int NS, int CLUSTER_SIZE = 1>
__global__ __launch_bounds__(TMA_NUM_WARPS * WARP_SIZE)
void grouped_gemm_kernel_v4(
const ProblemInfo* __restrict__ global_probs,
const WorkItem* __restrict__ work_items,
int num_items
) {
constexpr int TMA_A_SMEM_BYTES = BLOCK_M * (TMA_BLOCK_K / 2);
constexpr int TMA_B_SMEM_BYTES = BLOCK_N * (TMA_BLOCK_K / 2);
constexpr int TMA_SFA_SMEM_BYTES = MMA_M * (TMA_BLOCK_K / 16);
constexpr int TMA_SFB_SMEM_BYTES = BLOCK_N * (TMA_BLOCK_K / 16);
constexpr int STAGE_SIZE = TMA_A_SMEM_BYTES + TMA_B_SMEM_BYTES + TMA_SFA_SMEM_BYTES + TMA_SFB_SMEM_BYTES;
const int tid = threadIdx.x;
const int lane_id = tid % WARP_SIZE;
const int warp_id = tid / WARP_SIZE;
uint32_t cta_rank = 0;
if constexpr (CLUSTER_SIZE > 1) {
cta_rank = get_cluster_ctarank();
}
// Shared memory is used as a multi-stage ring buffer.
// Per stage: [A tile][B tile][SFA][SFB]. After all stages we place mbarriers.
extern __shared__ __align__(1024) char smem_ptr[];
const int smem_base = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
constexpr int B_off = TMA_A_SMEM_BYTES;
constexpr int SFA_off = B_off + TMA_B_SMEM_BYTES;
constexpr int SFB_off = SFA_off + TMA_SFA_SMEM_BYTES;
const int mbar_base = smem_base + STAGE_SIZE * NS;
// TMEM allocation is in columns. We need 2 columns per output column because
// accumulators are fp32.
constexpr int TMEM_COLS = BLOCK_N * 2;
constexpr int SFA_tmem = BLOCK_N;
constexpr int SFB_tmem = SFA_tmem + 4 * (TMA_BLOCK_K / MMA_K);
constexpr uint32_t idesc = (1U << 7U) | (1U << 10U)
| ((uint32_t)BLOCK_N >> 3U << 17U)
| ((uint32_t)MMA_M >> 7U << 27U);
if (warp_id == 0) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem_base), "r"(TMEM_COLS));
} else if (warp_id == 1 && elect_sync()) {
// Best-effort tensormap prefetch for the first few work items.
for (int i = 0; i < num_items && i < 8; ++i) {
const ProblemInfo* prob = &global_probs[work_items[i].problem_idx];
asm volatile("prefetch.tensormap [%0];" :: "l"(&prob->A_tmap) : "memory");
if constexpr (BLOCK_N == 256) {
asm volatile("prefetch.tensormap [%0];" :: "l"(&prob->B_tmap_256) : "memory");
} else {
asm volatile("prefetch.tensormap [%0];" :: "l"(&prob->B_tmap) : "memory");
}
}
}
__syncthreads();
__shared__ int shared_work_idx;
int work_idx = blockIdx.x;
if (work_idx < num_items) {
if (tid == 0) {
for (int i = 0; i < NS; ++i) {
// mbarrier[stage]: TMA completion barrier.
// Two producer warps arrive (A/SFA and B/SFB).
mbarrier_init(mbar_base + i * 8, 2);
// mbarrier[NS+stage]: stage reuse barrier.
// The MMA warp commits once per stage.
mbarrier_init(mbar_base + (NS + i) * 8, 1);
if constexpr (CLUSTER_SIZE > 1) {
if (cta_rank == 0) {
// Only CTA rank 0 initializes the cluster-wide barrier used to
// guard multicast stage reuse.
mbarrier_init(mbar_base + (2*NS + i) * 8, CLUSTER_SIZE);
}
}
}
asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
}
__syncthreads();
}
// Ensure all CTAs have initialized mbarriers before any multicast TMA.
if constexpr (CLUSTER_SIZE > 1) {
cluster_sync();
}
while (work_idx < num_items) {
const WorkItem& work = work_items[work_idx];
const ProblemInfo& prob = global_probs[work.problem_idx];
const int m_offset = work.tile_m * BLOCK_M;
const int n_offset = work.tile_n * BLOCK_N;
const int K = prob.K;
const int num_k_iters = K / TMA_BLOCK_K;
if ((warp_id == TMA_WARP || warp_id == TMA_WARP_B) && elect_sync()) {
constexpr uint64_t cache_A = EVICT_LAST;
constexpr uint64_t cache_B = EVICT_FIRST;
const bool do_A = (warp_id == TMA_WARP);
const bool do_B = (warp_id == TMA_WARP_B);
auto issue_tma = [&](int k_iter, int stage) {
const int mbar_addr = mbar_base + stage * 8;
const int stage_base = smem_base + stage * STAGE_SIZE;
const int off_k = k_iter * TMA_BLOCK_K;
// Program expect_tx before issuing any TMA that completes to this
// barrier. A completion arriving before expect_tx is set can leave the
// consumer stuck in mbarrier_wait.
const int expect_bytes = do_A
? (TMA_A_SMEM_BYTES + TMA_SFA_SMEM_BYTES)
: (TMA_B_SMEM_BYTES + TMA_SFB_SMEM_BYTES);
mbarrier_arrive_expect_tx(mbar_addr, expect_bytes);
if (do_A) {
if constexpr (CLUSTER_SIZE > 1) {
// Cluster path: CTA rank 0 multicasts A to the whole cluster.
// dst and mbarrier are passed as shared::cluster addresses.
if (cta_rank == 0) {
uint16_t mc = (1 << CLUSTER_SIZE) - 1;
int cluster_dst = mapa_cta_to_cluster(stage_base, 0);
int cluster_mbar = mapa_cta_to_cluster(mbar_addr, 0);
tma_3d_gmem2smem_multicast(cluster_dst, &prob.A_tmap, 0, m_offset, off_k / 256, cluster_mbar, mc);
}
} else {
tma_3d_gmem2smem<1>(stage_base, &prob.A_tmap, 0, m_offset, off_k / 256, mbar_addr, cache_A);
}
// SFA scale blocks are indexed by (m_tile, k_blk) and stored as
// 512B blocks (matching tcgen05_cp_nvfp4 granularity).
const int rest_k = K / 16 / 4;
const int k_blk = off_k / (16 * 4);
const char* SFA_src = prob.SFA_ptr + ((m_offset / 128) * rest_k + k_blk) * 512;
tma_gmem2smem(stage_base + SFA_off, SFA_src, TMA_SFA_SMEM_BYTES, mbar_addr, cache_A);
} else if (do_B) {
if constexpr (BLOCK_N == 256) {
tma_3d_gmem2smem<1>(stage_base + B_off, &prob.B_tmap_256, 0, n_offset, off_k / 256, mbar_addr, cache_B);
} else {
tma_3d_gmem2smem<1>(stage_base + B_off, &prob.B_tmap, 0, n_offset, off_k / 256, mbar_addr, cache_B);
}
// SFB scale blocks are indexed by (n_tile, k_blk).
const int rest_k = K / 16 / 4;
const int k_blk = off_k / (16 * 4);
if constexpr (BLOCK_N == 256) {
constexpr int SFB_HALF_BYTES = 128 * (TMA_BLOCK_K / 16);
const char* SFB_src0 = prob.SFB_ptr + ((n_offset / 128) * rest_k + k_blk) * 512;
const char* SFB_src1 = prob.SFB_ptr + (((n_offset / 128) + 1) * rest_k + k_blk) * 512;
tma_gmem2smem(stage_base + SFB_off, SFB_src0, SFB_HALF_BYTES, mbar_addr, cache_B);
tma_gmem2smem(stage_base + SFB_off + SFB_HALF_BYTES, SFB_src1, SFB_HALF_BYTES, mbar_addr, cache_B);
} else {
const char* SFB_src = prob.SFB_ptr + ((n_offset / 128) * rest_k + k_blk) * 512;
tma_gmem2smem(stage_base + SFB_off, SFB_src, TMA_SFB_SMEM_BYTES, mbar_addr, cache_B);
}
}
};
for (int k_iter = 0; k_iter < NS && k_iter < num_k_iters; k_iter++) {
issue_tma(k_iter, k_iter);
}
int stage = 0;
int mma_phase = 0;
for (int k_iter = NS; k_iter < num_k_iters; k_iter++) {
if constexpr (CLUSTER_SIZE > 1) {
if (do_A && cta_rank == 0) {
// A is shared across the cluster via multicast. Before reusing a
// ring-buffer stage for the next multicast, rank 0 must wait for
// all CTAs to finish consuming the current stage.
mbarrier_wait(mbar_base + (2*NS + stage) * 8, mma_phase);
} else {
mbarrier_wait(mbar_base + (NS + stage) * 8, mma_phase);
}
} else {
mbarrier_wait(mbar_base + (NS + stage) * 8, mma_phase);
}
issue_tma(k_iter, stage);
stage++;
if (stage == NS) {
stage = 0;
mma_phase ^= 1;
}
}
}
else if (warp_id == MMA_WARP && elect_sync()) {
auto make_desc_AB = [](int addr) -> uint64_t {
const int SBO = 8 * 128;
// Descriptor encoding is coupled to the shared-memory swizzle and the
// tcgen05 operand layout. SBO matches 128B swizzle.
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
auto make_desc_SF = [](int addr) -> uint64_t {
// Scale-factor loads use a different stride (16B) but the same address
// encoding (16B units).
const int SBO = 8 * 16;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
};
int stage = 0;
int tma_phase = 0;
for (int k_iter = 0; k_iter < num_k_iters; k_iter++) {
mbarrier_wait(mbar_base + stage * 8, tma_phase);
const int stage_base = smem_base + stage * STAGE_SIZE;
const uint64_t SF_desc = make_desc_SF(0);
const uint64_t SFA_desc = SF_desc + ((uint64_t)(stage_base + SFA_off) >> 4ULL);
const uint64_t SFB_desc = SF_desc + ((uint64_t)(stage_base + SFB_off) >> 4ULL);
// Copy scale factors from shared memory into TMEM.
#pragma unroll
for (int k = 0; k < TMA_BLOCK_K / MMA_K; k++) {
tcgen05_cp_nvfp4(SFA_tmem + k * 4, SFA_desc + (uint64_t)k * (512ULL >> 4ULL));
if constexpr (BLOCK_N == 256) {
constexpr uint64_t SFB_HALF_DESC = (uint64_t)(128 * (TMA_BLOCK_K / 16)) >> 4ULL;
tcgen05_cp_nvfp4(SFB_tmem + k * 8, SFB_desc + (uint64_t)k * (512ULL >> 4ULL));
tcgen05_cp_nvfp4(SFB_tmem + k * 8 + 4, SFB_desc + SFB_HALF_DESC + (uint64_t)k * (512ULL >> 4ULL));
} else {
tcgen05_cp_nvfp4(SFB_tmem + k * 4, SFB_desc + (uint64_t)k * (512ULL >> 4ULL));
}
}
// MMA loop over the 256-wide K tile in 64-wide chunks.
#pragma unroll
for (int k = 0; k < TMA_BLOCK_K / MMA_K; k++) {
uint64_t a_desc = make_desc_AB(stage_base + k * 32);
uint64_t b_desc = make_desc_AB(stage_base + B_off + k * 32);
const int scale_A_tmem = SFA_tmem + k * 4;
int scale_B_tmem;
if constexpr (BLOCK_N == 256) {
scale_B_tmem = SFB_tmem + k * 8;
} else {
scale_B_tmem = SFB_tmem + k * 4;
}
// First MMA uses D=0, subsequent MMAs accumulate.
const int enable_input_d = (k_iter == 0 && k == 0) ? 0 : 1;
tcgen05_mma_nvfp4(a_desc, b_desc, idesc, scale_A_tmem, scale_B_tmem, enable_input_d);
}
tcgen05_commit(mbar_base + (NS + stage) * 8);
if constexpr (CLUSTER_SIZE > 1) {
int cm = mapa_cta_to_cluster(mbar_base + (2*NS + stage) * 8, 0);
mbarrier_arrive_cluster(cm);
}
stage++;
if (stage == NS) {
stage = 0;
tma_phase ^= 1;
}
}
const int last_stage = (num_k_iters - 1) % NS;
const int last_phase = ((num_k_iters - 1) / NS) % 2;
mbarrier_wait(mbar_base + (NS + last_stage) * 8, last_phase);
}
__syncthreads();
asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");
const bool short_k = (K <= 2048);
const bool full_tile = (m_offset + BLOCK_M <= prob.M) && (n_offset + BLOCK_N <= prob.N);
if (short_k && full_tile && prob.Cs1 == 1) {
epilogue_store_fulltile_contiguous<BLOCK_M, BLOCK_N>(prob, m_offset, n_offset, tid, warp_id, lane_id);
} else {
const bool adaptive_m_epilogue = (prob.M <= LOW_M_THRESHOLD) || short_k;
if (adaptive_m_epilogue) {
epilogue_store<BLOCK_M, BLOCK_N, true>(prob, m_offset, n_offset, tid, warp_id, lane_id);
} else {
epilogue_store<BLOCK_M, BLOCK_N, false>(prob, m_offset, n_offset, tid, warp_id, lane_id);
}
}
// In cluster mode we must synchronize across CTAs before re-initializing
// mbarriers; __syncthreads is CTA-local and does not order cluster-wide
// mbarrier arrivals.
if constexpr (PERSISTENT && CLUSTER_SIZE > 1) {
cluster_sync();
}
if constexpr (PERSISTENT) {
if (warp_id == TMA_WARP && elect_sync()) {
// Static grid-stride work distribution avoids global atomics and
// smooths the tail when num_items slightly exceeds one wave.
shared_work_idx = work_idx + gridDim.x;
if (shared_work_idx < num_items) {
for (int i = 0; i < NS; ++i) {
mbarrier_init(mbar_base + i * 8, 2);
mbarrier_init(mbar_base + (NS + i) * 8, 1);
if constexpr (CLUSTER_SIZE > 1) {
if (cta_rank == 0) {
mbarrier_init(mbar_base + (2*NS + i) * 8, CLUSTER_SIZE);
}
}
}
asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
}
}
}
if constexpr (PERSISTENT) {
__syncthreads();
}
// Cluster barrier after re-init: ensure all CTAs see re-initialized mbarriers.
if constexpr (PERSISTENT && CLUSTER_SIZE > 1) {
cluster_sync();
}
if constexpr (PERSISTENT) {
work_idx = shared_work_idx;
} else {
break;
}
}
if (warp_id == 0) {
tcgen05_dealloc_cols_cta1(0, TMEM_COLS);
}
}
void init_AB_tmap_u4(
CUtensorMap *tmap,
const void *ptr,
uint64_t global_height, uint64_t global_width,
uint32_t shared_height, uint32_t shared_width
) {
TORCH_CHECK(ptr != nullptr, "ptr is null");
TORCH_CHECK(((uintptr_t)ptr % 16) == 0, "ptr must be 16-byte aligned");
TORCH_CHECK(global_width >= 256 && (global_width % 256) == 0, "K must be multiple of 256");
TORCH_CHECK(shared_width == 256, "shared_width must be 256");
constexpr uint32_t rank = 3;
uint64_t globalDim[rank] = {256, global_height, global_width / 256};
uint64_t globalStrides[rank-1] = {global_width / 2, 128};
uint32_t boxDim[rank] = {256, shared_height, 1};
uint32_t elementStrides[rank] = {1, 1, 1};
// Swizzle must match the shared-memory layout expected by tcgen05.
constexpr CUtensorMapSwizzle swizzle = CU_TENSOR_MAP_SWIZZLE_128B;
// Cache cuTensorMap templates by shape.
struct ShapeKey { uint64_t gh, gw; uint32_t sh, sw; };
struct ShapeHash {
size_t operator()(const ShapeKey& k) const noexcept {
uint64_t h = k.gh;
h ^= (k.gw + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2));
h ^= ((uint64_t)k.sh << 32) ^ (uint64_t)k.sw;
return (size_t)h;
}
};
struct ShapeEq {
bool operator()(const ShapeKey& a, const ShapeKey& b) const noexcept {
return a.gh == b.gh && a.gw == b.gw && a.sh == b.sh && a.sw == b.sw;
}
};
struct PtrKey { uint64_t gh, gw; uint32_t sh, sw; const void* ptr; };
struct PtrHash {
size_t operator()(const PtrKey& k) const noexcept {
uint64_t h = k.gh;
h ^= (k.gw + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2));
h ^= ((uint64_t)k.sh << 32) ^ (uint64_t)k.sw;
h ^= ((uint64_t)k.ptr >> 4);
return (size_t)h;
}
};
struct PtrEq {
bool operator()(const PtrKey& a, const PtrKey& b) const noexcept {
return a.gh == b.gh && a.gw == b.gw && a.sh == b.sh && a.sw == b.sw && a.ptr == b.ptr;
}
};
static thread_local std::unordered_map<ShapeKey, CUtensorMap, ShapeHash, ShapeEq> tmpl_cache;
static thread_local std::unordered_map<PtrKey, CUtensorMap, PtrHash, PtrEq> ptr_cache;
PtrKey pkey{global_height, global_width, shared_height, shared_width, ptr};
auto pit = ptr_cache.find(pkey);
if (pit != ptr_cache.end()) { *tmap = pit->second; return; }
ShapeKey skey{global_height, global_width, shared_height, shared_width};
auto sit = tmpl_cache.find(skey);
if (sit == tmpl_cache.end()) {
CUtensorMap tmp;
auto err = cuTensorMapEncodeTiled(
&tmp, CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
rank, (void*)ptr, globalDim, globalStrides, boxDim, elementStrides,
CU_TENSOR_MAP_INTERLEAVE_NONE, swizzle,
CU_TENSOR_MAP_L2_PROMOTION_NONE, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
);
TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapEncodeTiled failed");
sit = tmpl_cache.emplace(skey, tmp).first;
}
CUtensorMap tmp = sit->second;
auto err = cuTensorMapReplaceAddress(&tmp, (void*)ptr);
TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapReplaceAddress failed");
ptr_cache.emplace(pkey, tmp);
*tmap = tmp;
}
std::vector<at::Tensor> group_gemm(
std::vector<at::Tensor> A_list,
std::vector<at::Tensor> B_list,
std::vector<at::Tensor> C_list,
std::vector<at::Tensor> sfa_list,
std::vector<at::Tensor> sfb_list,
at::Tensor sizes_cpu
) {
int64_t G = A_list.size();
auto dev = A_list[0].device();
c10::cuda::CUDAGuard device_guard(dev);
auto sizes_accessor = sizes_cpu.accessor<int64_t, 2>();
// SM count is used for occupancy-based launch shaping.
// Hardcoded for the target environment (B200 / SM100) to keep the logic
// simple and deterministic.
constexpr int sm_count = 148;
static bool attrs_set = false;
if (!attrs_set) {
// Two pipeline depths:
// - HI: deeper pipeline, more overlap, higher shared-memory footprint.
// - LO: shallower pipeline, lower shared-memory footprint (can improve
// occupancy when shared memory is the limiter).
constexpr int NS_DEEP_HI = 6;
constexpr int NS_DEEP_LO = 3;
constexpr int NS_WIDE_HI = 4;
constexpr int NS_WIDE_LO = 2;
constexpr int MBAR_BYTES_DEEP_HI = ((2 * NS_DEEP_HI * 8 + 63) & ~63);
constexpr int MBAR_BYTES_DEEP_LO = ((2 * NS_DEEP_LO * 8 + 63) & ~63);
constexpr int MBAR_BYTES_WIDE_HI = ((2 * NS_WIDE_HI * 8 + 63) & ~63);
constexpr int MBAR_BYTES_WIDE_LO = ((2 * NS_WIDE_LO * 8 + 63) & ~63);
constexpr int MBAR_BYTES_WIDE_CLUSTER_HI = ((3 * NS_WIDE_HI * 8 + 63) & ~63);
constexpr int MBAR_BYTES_WIDE_CLUSTER_LO = ((3 * NS_WIDE_LO * 8 + 63) & ~63);
// BLOCK_N=128
constexpr int STAGE_128_K256 = 128 * (256 / 2) + 128 * (256 / 2) + 128 * (256 / 16) + 128 * (256 / 16);
constexpr int SMEM_128_HI_K256 = STAGE_128_K256 * NS_DEEP_HI + MBAR_BYTES_DEEP_HI;
constexpr int SMEM_128_LO_K256 = STAGE_128_K256 * NS_DEEP_LO + MBAR_BYTES_DEEP_LO;
// BLOCK_N=256
constexpr int STAGE_256_K256 = 128 * (256 / 2) + 256 * (256 / 2) + 128 * (256 / 16) + 256 * (256 / 16);
constexpr int SMEM_256_HI_K256 = STAGE_256_K256 * NS_WIDE_HI + MBAR_BYTES_WIDE_HI;
constexpr int SMEM_256_LO_K256 = STAGE_256_K256 * NS_WIDE_LO + MBAR_BYTES_WIDE_LO;
constexpr int SMEM_256_CLUSTER_HI_K256 = STAGE_256_K256 * NS_WIDE_HI + MBAR_BYTES_WIDE_CLUSTER_HI;
constexpr int SMEM_256_CLUSTER_LO_K256 = STAGE_256_K256 * NS_WIDE_LO + MBAR_BYTES_WIDE_CLUSTER_LO;
// Variants: {persistent} x {BLOCK_N} x {NS}
// BLOCK_N=128
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 128, NS_DEEP_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_128_HI_K256));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_128_HI_K256));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 128, NS_DEEP_LO>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_128_LO_K256));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_LO>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_128_LO_K256));
// BLOCK_N=256 without cluster
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_HI_K256));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_HI_K256));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_LO>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_LO_K256));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_LO_K256));
// BLOCK_N=256 with cluster multicast (CLUSTER_SIZE=4)
constexpr int CL = 4;
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_CLUSTER_HI_K256));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_CLUSTER_HI_K256));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_LO, CL>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_CLUSTER_LO_K256));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO, CL>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_CLUSTER_LO_K256));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeNonPortableClusterSizeAllowed, 1));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeNonPortableClusterSizeAllowed, 1));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_LO, CL>), cudaFuncAttributeNonPortableClusterSizeAllowed, 1));
CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO, CL>), cudaFuncAttributeNonPortableClusterSizeAllowed, 1));
attrs_set = true;
}
// Cache occupancy for launch shaping (per process).
static thread_local bool occ_set = false;
static thread_local int occ_128_hi = 0, occ_128_lo = 0;
static thread_local int occ_256_hi = 0, occ_256_lo = 0;
static thread_local int occ_256_cluster_hi = 0, occ_256_cluster_lo = 0;
if (!occ_set) {
constexpr int NS_DEEP_HI = 6;
constexpr int NS_DEEP_LO = 3;
constexpr int NS_WIDE_HI = 4;
constexpr int NS_WIDE_LO = 2;
constexpr int MBAR_BYTES_DEEP_HI = ((2 * NS_DEEP_HI * 8 + 63) & ~63);
constexpr int MBAR_BYTES_DEEP_LO = ((2 * NS_DEEP_LO * 8 + 63) & ~63);
constexpr int MBAR_BYTES_WIDE_HI = ((2 * NS_WIDE_HI * 8 + 63) & ~63);
constexpr int MBAR_BYTES_WIDE_LO = ((2 * NS_WIDE_LO * 8 + 63) & ~63);
constexpr int MBAR_BYTES_WIDE_CLUSTER_HI = ((3 * NS_WIDE_HI * 8 + 63) & ~63);
constexpr int MBAR_BYTES_WIDE_CLUSTER_LO = ((3 * NS_WIDE_LO * 8 + 63) & ~63);
constexpr int STAGE_128_K256 = 128 * (256 / 2) + 128 * (256 / 2) + 128 * (256 / 16) + 128 * (256 / 16);
constexpr int STAGE_256_K256 = 128 * (256 / 2) + 256 * (256 / 2) + 128 * (256 / 16) + 256 * (256 / 16);
constexpr int SMEM_128_HI_K256 = STAGE_128_K256 * NS_DEEP_HI + MBAR_BYTES_DEEP_HI;
constexpr int SMEM_128_LO_K256 = STAGE_128_K256 * NS_DEEP_LO + MBAR_BYTES_DEEP_LO;
constexpr int SMEM_256_HI_K256 = STAGE_256_K256 * NS_WIDE_HI + MBAR_BYTES_WIDE_HI;
constexpr int SMEM_256_LO_K256 = STAGE_256_K256 * NS_WIDE_LO + MBAR_BYTES_WIDE_LO;
constexpr int SMEM_256_CLUSTER_HI_K256 = STAGE_256_K256 * NS_WIDE_HI + MBAR_BYTES_WIDE_CLUSTER_HI;
constexpr int SMEM_256_CLUSTER_LO_K256 = STAGE_256_K256 * NS_WIDE_LO + MBAR_BYTES_WIDE_CLUSTER_LO;
constexpr int THREADS = TMA_NUM_WARPS * WARP_SIZE;
CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&occ_128_hi, grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_HI>, THREADS, SMEM_128_HI_K256));
CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&occ_128_lo, grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_LO>, THREADS, SMEM_128_LO_K256));
CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&occ_256_hi, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI>, THREADS, SMEM_256_HI_K256));
CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&occ_256_lo, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO>, THREADS, SMEM_256_LO_K256));
CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&occ_256_cluster_hi, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI, 4>, THREADS, SMEM_256_CLUSTER_HI_K256));
CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&occ_256_cluster_lo, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO, 4>, THREADS, SMEM_256_CLUSTER_LO_K256));
TORCH_CHECK(occ_128_hi > 0 && occ_128_lo > 0 && occ_256_hi > 0 && occ_256_lo > 0
&& occ_256_cluster_hi > 0 && occ_256_cluster_lo > 0,
"occupancy query returned zero blocks/SM");
occ_set = true;
}
std::vector<ProblemInfo> problem_infos(G);
static thread_local std::vector<WorkItem> cached_work_items_128;
static thread_local std::vector<WorkItem> cached_work_items_256;
static thread_local std::vector<WorkItem> cached_work_items_256_cluster;
static thread_local uint64_t cached_work_hash = 0;
static thread_local bool cached_work_valid = false;
std::vector<uint8_t> active(G, 0);
std::vector<uint8_t> use_256(G, 0);
std::vector<uint8_t> use_256_cluster(G, 0);
std::vector<int> num_tiles_m(G, 0);
std::vector<int> num_tiles_n(G, 0);
std::vector<int64_t> Ms(G, 0), Ns(G, 0), Ks(G, 0);
int64_t total_tiles = 0;
for (int64_t i = 0; i < G; i++) {
const int64_t M = sizes_accessor[i][0], N = sizes_accessor[i][1], K = sizes_accessor[i][2];
Ms[(size_t)i] = M; Ns[(size_t)i] = N; Ks[(size_t)i] = K;
if (A_list[i].stride(1) != 1 || B_list[i].stride(1) != 1) {
continue;
}
active[(size_t)i] = 1;
bool is_256 = (N >= 4096) && (K >= 2048) && ((N & 255) == 0);
bool is_256_cluster = is_256 && (K > 2048);
use_256[(size_t)i] = is_256;
use_256_cluster[(size_t)i] = is_256_cluster;
int block_m = 128;
int block_n = is_256 ? 256 : 128;
num_tiles_m[(size_t)i] = ceil_div((int)M, block_m);
num_tiles_n[(size_t)i] = ceil_div((int)N, block_n);
total_tiles += (int64_t)num_tiles_m[(size_t)i] * (int64_t)num_tiles_n[(size_t)i];
}
constexpr int tma_block_k = 256;
uint64_t work_hash = 1469598103934665603ULL;
work_hash = hash_combine_u64(work_hash, (uint64_t)tma_block_k);
for (int64_t i = 0; i < G; i++) {
const bool is_active = (active[(size_t)i] != 0);
if (!is_active) {
work_hash = hash_combine_u64(work_hash, 0);
continue;
}
const int64_t M = Ms[(size_t)i];
const int64_t N = Ns[(size_t)i];
const bool is_256 = (use_256[(size_t)i] != 0);
const bool is_256_cluster = (use_256_cluster[(size_t)i] != 0);
work_hash = hash_combine_u64(work_hash, (uint64_t)M);
work_hash = hash_combine_u64(work_hash, (uint64_t)N);
work_hash = hash_combine_u64(work_hash, (uint64_t)is_256);
work_hash = hash_combine_u64(work_hash, (uint64_t)is_256_cluster);
}
if (!cached_work_valid || cached_work_hash != work_hash) {
cached_work_items_128.clear();
cached_work_items_256.clear();
cached_work_items_256_cluster.clear();
cached_work_items_128.reserve(G * 32);
cached_work_items_256.reserve(G * 32);
cached_work_items_256_cluster.reserve(G * 32);
// Work scheduling.
// 128-wide path: iterate by tile_n across groups so a wave tends to touch
// the same B strip across different problems (better L2 locality).
int max_tn_128 = 0;
for (int64_t i = 0; i < G; i++) {
if (!active[(size_t)i]) continue;
if (!use_256[(size_t)i]) {
max_tn_128 = std::max(max_tn_128, num_tiles_n[(size_t)i]);
}
}
// 128-wide tiles: interleave by tile_n across groups
for (int tn = 0; tn < max_tn_128; tn++) {
for (int64_t i = 0; i < G; i++) {
if (!active[(size_t)i] || use_256[(size_t)i]) continue;
if (tn >= num_tiles_n[(size_t)i]) continue;
for (int tm = 0; tm < num_tiles_m[(size_t)i]; tm++) {
cached_work_items_128.push_back({(int)i, tm, tn});
}
}
}
// 256-wide path: split by K.
// - K > 2048: clustered multicast path.
// - K <= 2048: non-cluster path to avoid cluster overhead.
constexpr int CLUSTER_SIZE_256 = 4;
for (int64_t i = 0; i < G; i++) {
if (!active[(size_t)i] || !use_256[(size_t)i]) continue;
for (int tm = 0; tm < num_tiles_m[(size_t)i]; tm++) {
for (int tn = 0; tn < num_tiles_n[(size_t)i]; tn++) {
if (use_256_cluster[(size_t)i]) {
cached_work_items_256_cluster.push_back({(int)i, tm, tn});
} else {
cached_work_items_256.push_back({(int)i, tm, tn});
}
}
if (use_256_cluster[(size_t)i]) {
// Pad clustered work to a multiple of CLUSTER_SIZE so every
// launched cluster is full. Duplicating an existing tile is safe:
// it deterministically writes the same result.
int remainder = num_tiles_n[(size_t)i] % CLUSTER_SIZE_256;
if (remainder != 0) {
for (int p = 0; p < CLUSTER_SIZE_256 - remainder; p++) {
cached_work_items_256_cluster.push_back({(int)i, tm, 0});
}
}
}
}
}
cached_work_hash = work_hash;
cached_work_valid = true;
}
uint64_t probs_hash = 1469598103934665603ULL;
probs_hash = hash_combine_u64(probs_hash, (uint64_t)tma_block_k);
for (int64_t i = 0; i < G; i++) {
if (!active[(size_t)i]) continue;
const int64_t M = Ms[(size_t)i], N = Ns[(size_t)i], K = Ks[(size_t)i];
ProblemInfo& p = problem_infos[i];
p.M = M; p.N = N; p.K = K;
p.Cs0 = C_list[i].stride(0); p.Cs1 = C_list[i].stride(1);
p.C_ptr = (half*)C_list[i].data_ptr();
p.SFA_ptr = (const char*)sfa_list[i].data_ptr();
p.SFB_ptr = (const char*)sfb_list[i].data_ptr();
init_AB_tmap_u4(&p.A_tmap, A_list[i].data_ptr(), A_list[i].size(0), K, 128, 256);
init_AB_tmap_u4(&p.B_tmap, B_list[i].data_ptr(), B_list[i].size(0), K, 128, 256);
if (use_256[(size_t)i]) {
init_AB_tmap_u4(&p.B_tmap_256, B_list[i].data_ptr(), B_list[i].size(0), K, 256, 256);
} else {
p.B_tmap_256 = p.B_tmap;
}
probs_hash = hash_combine_u64(probs_hash, (uint64_t)i);
probs_hash = hash_combine_u64(probs_hash, (uint64_t)M);
probs_hash = hash_combine_u64(probs_hash, (uint64_t)N);
probs_hash = hash_combine_u64(probs_hash, (uint64_t)K);
probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)A_list[i].data_ptr());
probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)B_list[i].data_ptr());
probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)p.C_ptr);
probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)p.SFA_ptr);
probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)p.SFB_ptr);
probs_hash = hash_combine_u64(probs_hash, (uint64_t)p.Cs0);
probs_hash = hash_combine_u64(probs_hash, (uint64_t)p.Cs1);
}
if (cached_work_items_128.empty() && cached_work_items_256.empty() && cached_work_items_256_cluster.empty()) return C_list;
auto options = at::TensorOptions().dtype(at::kByte).device(dev);
static thread_local at::Tensor d_probs_cache;
static thread_local at::Tensor d_work_cache_128;
static thread_local at::Tensor d_work_cache_256;
static thread_local at::Tensor d_work_cache_256_cluster;
static thread_local uint64_t last_probs_hash = 0;
static thread_local uint64_t last_work_hash = 0;
static thread_local bool last_hash_valid = false;
const int64_t probs_bytes = (int64_t)(G * sizeof(ProblemInfo));
const int64_t work_bytes_128 = (int64_t)(cached_work_items_128.size() * sizeof(WorkItem));
const int64_t work_bytes_256 = (int64_t)(cached_work_items_256.size() * sizeof(WorkItem));
const int64_t work_bytes_256_cluster = (int64_t)(cached_work_items_256_cluster.size() * sizeof(WorkItem));
bool probs_realloc = false;
if (!d_probs_cache.defined() || d_probs_cache.device() != dev || d_probs_cache.scalar_type() != at::kByte || d_probs_cache.numel() < probs_bytes) {
d_probs_cache = at::empty({probs_bytes}, options);
probs_realloc = true;
}
if (work_bytes_128 > 0 && (!d_work_cache_128.defined() || d_work_cache_128.device() != dev || d_work_cache_128.scalar_type() != at::kByte || d_work_cache_128.numel() < work_bytes_128)) {
d_work_cache_128 = at::empty({work_bytes_128}, options);
}
if (work_bytes_256 > 0 && (!d_work_cache_256.defined() || d_work_cache_256.device() != dev || d_work_cache_256.scalar_type() != at::kByte || d_work_cache_256.numel() < work_bytes_256)) {
d_work_cache_256 = at::empty({work_bytes_256}, options);
}
if (work_bytes_256_cluster > 0 && (!d_work_cache_256_cluster.defined() || d_work_cache_256_cluster.device() != dev || d_work_cache_256_cluster.scalar_type() != at::kByte || d_work_cache_256_cluster.numel() < work_bytes_256_cluster)) {
d_work_cache_256_cluster = at::empty({work_bytes_256_cluster}, options);
}
if (probs_realloc || !last_hash_valid || last_probs_hash != probs_hash) {
CUDA_CHECK(cudaMemcpyAsync(d_probs_cache.data_ptr(), problem_infos.data(), G * sizeof(ProblemInfo), cudaMemcpyHostToDevice));
last_probs_hash = probs_hash;
}
if (!last_hash_valid || last_work_hash != work_hash) {
if (work_bytes_128 > 0) {
CUDA_CHECK(cudaMemcpyAsync(d_work_cache_128.data_ptr(), cached_work_items_128.data(), work_bytes_128, cudaMemcpyHostToDevice));
}
if (work_bytes_256 > 0) {
CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256.data_ptr(), cached_work_items_256.data(), work_bytes_256, cudaMemcpyHostToDevice));
}
if (work_bytes_256_cluster > 0) {
CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256_cluster.data_ptr(), cached_work_items_256_cluster.data(), work_bytes_256_cluster, cudaMemcpyHostToDevice));
}
last_work_hash = work_hash;
}
last_hash_valid = true;
constexpr int NS_DEEP_HI = 6;
constexpr int NS_DEEP_LO = 3;
constexpr int NS_WIDE_HI = 4;
constexpr int NS_WIDE_LO = 2;
if (!cached_work_items_128.empty()) {
int num_items_128 = (int)cached_work_items_128.size();
// Choose pipeline depth by expected waves.
const int wave_hi = sm_count * occ_128_hi;
const int wave_lo = sm_count * occ_128_lo;
// Only switch to low-smem variant when the grid is large enough that reducing
// waves is likely to outweigh reduced pipeline overlap.
const bool use_lo = (wave_lo > wave_hi) && (num_items_128 > 2 * wave_hi);
const int ns = use_lo ? NS_DEEP_LO : NS_DEEP_HI;
const int occ = use_lo ? occ_128_lo : occ_128_hi;
const int wave_cap = sm_count * occ;
constexpr int MBAR_HI = ((2 * NS_DEEP_HI * 8 + 63) & ~63);
constexpr int MBAR_LO = ((2 * NS_DEEP_LO * 8 + 63) & ~63);
const int stage_size_128 = 128 * (tma_block_k / 2) + 128 * (tma_block_k / 2) + 128 * (tma_block_k / 16) + 128 * (tma_block_k / 16);
const int SMEM_SIZE_128 = stage_size_128 * ns + (use_lo ? MBAR_LO : MBAR_HI);
// Launch shaping: if more CTAs than one full wave, use persistent grid-stride.
const bool persistent_128 = (num_items_128 > wave_cap);
const int launch_ctas_128 = persistent_128 ? wave_cap : num_items_128;
if (persistent_128) {
if (use_lo) {
grouped_gemm_kernel_v4<true, 128, 128, NS_DEEP_LO><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128);
} else {
grouped_gemm_kernel_v4<true, 128, 128, NS_DEEP_HI><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128);
}
} else {
if (use_lo) {
grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_LO><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128);
} else {
grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_HI><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128);
}
}
}
if (!cached_work_items_256.empty()) {
int num_items_256 = (int)cached_work_items_256.size();
const int wave_hi = sm_count * occ_256_hi;
const int wave_lo = sm_count * occ_256_lo;
const bool use_lo = (wave_lo > wave_hi) && (num_items_256 > 2 * wave_hi);
const int ns = use_lo ? NS_WIDE_LO : NS_WIDE_HI;
const int occ = use_lo ? occ_256_lo : occ_256_hi;
const int wave_cap = sm_count * occ;
constexpr int MBAR_HI = ((2 * NS_WIDE_HI * 8 + 63) & ~63);
constexpr int MBAR_LO = ((2 * NS_WIDE_LO * 8 + 63) & ~63);
const int stage_size_256 = 128 * (tma_block_k / 2) + 256 * (tma_block_k / 2) + 128 * (tma_block_k / 16) + 256 * (tma_block_k / 16);
const int SMEM_SIZE_256 = stage_size_256 * ns + (use_lo ? MBAR_LO : MBAR_HI);
const bool persistent_256 = (num_items_256 > wave_cap);
const int launch_ctas_256 = persistent_256 ? wave_cap : num_items_256;
if (persistent_256) {
if (use_lo) {
grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_LO><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256);
} else {
grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256);
}
} else {
if (use_lo) {
grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256);
} else {
grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256);
}
}
}
if (!cached_work_items_256_cluster.empty()) {
int num_items_256_cluster = (int)cached_work_items_256_cluster.size();
constexpr int CL256 = 4;
const int wave_hi = sm_count * occ_256_cluster_hi;
const int wave_lo = sm_count * occ_256_cluster_lo;
const bool use_lo = (wave_lo > wave_hi) && (num_items_256_cluster > 2 * wave_hi);
const int ns = use_lo ? NS_WIDE_LO : NS_WIDE_HI;
const int occ = use_lo ? occ_256_cluster_lo : occ_256_cluster_hi;
// Align wave_cap to cluster size so grid-stride preserves cluster alignment.
const int wave_cap = (sm_count * occ / CL256) * CL256;
constexpr int MBAR_HI = ((3 * NS_WIDE_HI * 8 + 63) & ~63);
constexpr int MBAR_LO = ((3 * NS_WIDE_LO * 8 + 63) & ~63);
const int stage_size_256 = 128 * (tma_block_k / 2) + 256 * (tma_block_k / 2) + 128 * (tma_block_k / 16) + 256 * (tma_block_k / 16);
const int SMEM_SIZE_256 = stage_size_256 * ns + (use_lo ? MBAR_LO : MBAR_HI);
const bool persistent_256 = (num_items_256_cluster > wave_cap);
// Grid must be multiple of cluster size.
int launch_ctas_256 = persistent_256 ? wave_cap : num_items_256_cluster;
// num_items_256_cluster is already padded to multiple of CL256 by work scheduling.
const ProblemInfo* d_probs_ptr = (const ProblemInfo*)d_probs_cache.data_ptr();
const WorkItem* d_work_ptr = (const WorkItem*)d_work_cache_256_cluster.data_ptr();
cudaLaunchConfig_t config = {};
config.gridDim = dim3(launch_ctas_256);
config.blockDim = dim3(TMA_NUM_WARPS * WARP_SIZE);
config.dynamicSmemBytes = SMEM_SIZE_256;
cudaLaunchAttribute launch_attrs[1];
launch_attrs[0].id = cudaLaunchAttributeClusterDimension;
launch_attrs[0].val.clusterDim = {CL256, 1, 1};
config.attrs = launch_attrs;
config.numAttrs = 1;
if (persistent_256) {
if (use_lo) {
CUDA_CHECK(cudaLaunchKernelEx(&config, grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_LO, CL256>,
d_probs_ptr, d_work_ptr, num_items_256_cluster));
} else {
CUDA_CHECK(cudaLaunchKernelEx(&config, grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI, CL256>,
d_probs_ptr, d_work_ptr, num_items_256_cluster));
}
} else {
if (use_lo) {
CUDA_CHECK(cudaLaunchKernelEx(&config, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO, CL256>,
d_probs_ptr, d_work_ptr, num_items_256_cluster));
} else {
CUDA_CHECK(cudaLaunchKernelEx(&config, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI, CL256>,
d_probs_ptr, d_work_ptr, num_items_256_cluster));
}
}
}
CUDA_CHECK(cudaGetLastError());
return C_list;
}
TORCH_LIBRARY(my_module, m) {
m.def("group_gemm(Tensor[] a, Tensor[] b, Tensor[] c, Tensor[] sfa, Tensor[] sfb, Tensor sizes) -> Tensor[]");
m.impl("group_gemm", &group_gemm);
}
"""
load_inline(
"group_gemm",
cpp_sources="",
cuda_sources=CUDA_SRC,
verbose=True,
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",
"-Xptxas=-v",
],
extra_ldflags=["-lcuda"],
)
group_gemm = torch.ops.my_module.group_gemm
@lru_cache(maxsize=128)
def _sizes_cpu_cached(key: tuple[tuple[int, int, int], ...]) -> torch.Tensor:
return torch.tensor(key, dtype=torch.int64, device="cpu")
def custom_kernel(data: input_t) -> output_t:
abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = data
A_list = [t[0] for t in abc_tensors]
B_list = [t[1] for t in abc_tensors]
C_list = [t[2] for t in abc_tensors]
sfa_list = [t[0] for t in sfasfb_reordered_tensors]
sfb_list = [t[1] for t in sfasfb_reordered_tensors]
key = tuple(tuple(int(v) for v in x) for x in problem_sizes)
sizes_cpu = _sizes_cpu_cached(key)
out = group_gemm(A_list, B_list, C_list, sfa_list, sfb_list, sizes_cpu)
return cast(output_t, out)
scrolls · 1385 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 490526.
#!POPCORN leaderboard nvfp4_group_gemm- #!POPCORN gpu NVIDIAfrom __future__ import annotations+ from functools import lru_cachefrom typing import castimport torchfrom torch.utils.cpp_extension import load_inline+ from task import input_t, output_t+"""g: 8; k: [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168]; m: [80, 176, 128, 72, 64, 248, 96, 160]; n: [4096, 4096, 4096, 4096, 4096, 4096, 4096, 4096]; seed: 1111- ⏱ 60.3 ± 0.02 µs- ⚡ 60.2 µs 🐌 60.3 µs+ ⏱ 47.5 ± 0.00 µs+ ⚡ 47.4 µs 🐌 47.5 µsg: 8; k: [2048, 2048, 2048, 2048, 2048, 2048, 2048, 2048]; m: [40, 76, 168, 72, 164, 148, 196, 160]; n: [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168]; seed: 1111- ⏱ 48.3 ± 0.05 µs- ⚡ 47.9 µs 🐌 48.6 µs+ ⏱ 45.0 ± 0.04 µs+ ⚡ 44.6 µs 🐌 45.3 µsg: 2; k: [4096, 4096]; m: [192, 320]; n: [3072, 3072]; seed: 1111- ⏱ 19.6 ± 0.02 µs- ⚡ 19.3 µs 🐌 20.0 µs+ ⏱ 14.3 ± 0.01 µs+ ⚡ 13.9 µs 🐌 14.6 µsg: 2; k: [1536, 1536]; m: [128, 384]; n: [4096, 4096]; seed: 1111- ⏱ 15.3 ± 0.02 µs- ⚡ 14.9 µs 🐌 15.6 µs+ ⏱ 10.5 ± 0.01 µs+ ⚡ 10.4 µs 🐌 10.7 µs"""- from task import input_t, output_t-CUDA_SRC = """#include <vector>#include <unordered_map>⋯ 16 unchanged linesTORCH_CHECK(_err == cudaSuccess, "CUDA error: ", cudaGetErrorString(_err)); \\} while (0)+ static inline uint64_t hash_combine_u64(uint64_t h, uint64_t x) {+ h ^= x;+ h *= 1099511628211ULL;+ return h;+ }+constexpr int WARP_SIZE = 32;constexpr int MMA_K = 64;⋯ 6 unchanged linesint tile_n;};+ // Per-problem metadata.+ // Align on 128 byte boundary, useful since this is read by many CTAs.struct __align__(128) ProblemInfo {CUtensorMap A_tmap;CUtensorMap B_tmap;⋯ 5 unchanged linesint64_t Cs0, Cs1;};- __device__ inline- constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; }+ // tcgen05 descriptors encode shared-memory addresses in 16-byte units.+ // Mask to the HW-supported address width and drop the 16B alignment bits.+ __device__ inline constexpr uint64_t desc_encode(uint64_t x) {+ return (x & 0x3'FFFFULL) >> 4ULL;+ }- __device__- uint32_t elect_sync() {+ // elect.sync: use this to have a single lane issue TMA/tcgen05 instructions while the+ // whole warp stays converged.+ __device__ uint32_t elect_sync() {uint32_t pred = 0;asm volatile("{\\n\\t"⋯ 7 unchanged linesreturn pred;}+ __device__ inline uint32_t get_cluster_ctarank() {+ uint32_t rank;+ asm volatile("mov.u32 %0, %%cluster_ctarank;" : "=r"(rank));+ return rank;+ }++ __device__ inline void cluster_sync() {+ asm volatile("barrier.cluster.arrive;" ::: "memory");+ asm volatile("barrier.cluster.wait;" ::: "memory");+ }++ // Shared-memory mbarrier helpers.+ // Used for:+ // - TMA completion barrier: consumer waits for bytes to arrive in shared memory.+ // - Stage reuse barrier: producer waits until MMA is done with a stage before+ // overwriting it.__device__ inline void mbarrier_init(int mbar_addr, int count) {asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));}+ // Program the expected byte count for a TMA stage and arrive.+ // This must happen before issuing any cp.async.bulk.* that completes to the+ // barrier.__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::cta.b64 _, [%0], %1;":: "r"(mbar_addr), "r"(size) : "memory");}⋯ 10 unchanged lines);}+ // TMA: 3D tensor-map load from global -> shared memory.+ // The (x,y,z) coordinates correspond to the CUtensorMap encoding in+ // init_AB_tmap_u4.template <int CTA_GROUP = 1>__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) {⋯ 5 unchanged lines);}+ __device__ inline int mapa_cta_to_cluster(int cta_addr, int dest_cta) {+ int cluster_addr;+ asm volatile("mapa.shared::cluster.u32 %0, %1, %2;" : "=r"(cluster_addr) : "r"(cta_addr), "r"(dest_cta));+ return cluster_addr;+ }++ __device__ inline void mbarrier_arrive_cluster(int mbar_cluster_addr) {+ asm volatile("mbarrier.arrive.shared::cluster.b64 _, [%0];"+ :: "r"(mbar_cluster_addr) : "memory");+ }++ // Cluster multicast variant of 3D TMA.+ // dst and mbar_addr are in shared::cluster address space.+ // The multicast mask selects which CTAs in the cluster receive the data.+ __device__ inline void tma_3d_gmem2smem_multicast(int dst, const void *tmap_ptr, int x, int y, int z,+ int mbar_addr, uint16_t multicast_mask) {+ asm volatile(+ "cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.cta_group::1 "+ "[%0], [%1, {%2, %3, %4}], [%5], %6;"+ :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "h"(multicast_mask)+ : "memory"+ );+ }++ // Linear bulk copy global -> shared.+ // Used for scale-factor tensors (SFA/SFB).__device__ inline void tma_gmem2smem(int dst, const void *src, int size, int mbar_addr, uint64_t cache_policy) {asm volatile("cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "⋯ 85 unchanged linesconstexpr int TMA_WARP_B = 6;constexpr int MMA_WARP = 5;+ // Epilogue: read fp32 accumulators from TMEM (tcgen05.ld) and store fp16 C.+ // Only threads with tid < BLOCK_M participate; this maps 4 warps (0..3) to the+ // 128 output rows, with each warp handling a 32-row stripe.template <int BLOCK_M, int BLOCK_N, bool LOW_M_EPILOGUE>__device__ __forceinline__ void epilogue_store(const ProblemInfo& prob,⋯ 25 unchanged linesm_iters = (remaining <= 16) ? 1 : 2;}+ // Lane mapping: each lane owns two columns (half2) and one of 8 rows.const int lane_row = lane_id >> 2;const int lane_col = (lane_id & 3) * 2;++ // We load/store in 128-column halves so tcgen05.ld has a fixed shape.constexpr int HALF_N = 128;const int halves = BLOCK_N / HALF_N;⋯ 1 unchanged linesfor (int half_idx = 0; half_idx < halves; ++half_idx) {float tmp[HALF_N / 2];const int col_base = half_idx * HALF_N;+ // TMEM coordinates are relative to the CTA's output tile.tcgen05_ld_16x256b_x16(tmp, warp_id * 32 + m * 16, col_base);asm volatile("tcgen05.wait::ld.sync.aligned;\\n");⋯ 9 unchanged linesconst int idx = i * 4;const int col = i * 8 + lane_col;const int h2_idx = col >> 1;- row0_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 0]), __float2half_rn(tmp[idx + 1]));- row1_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 2]), __float2half_rn(tmp[idx + 3]));+ row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));+ row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));}continue;}⋯ 9 unchanged linesconst int col = i * 8 + lane_col;const int h2_idx = col >> 1;if (row0_in) {- row0_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 0]), __float2half_rn(tmp[idx + 1]));+ row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));}if (row1_in) {- row1_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 2]), __float2half_rn(tmp[idx + 3]));+ row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));}}} else {⋯ 2 unchanged linesconst int idx = i * 4;const int col = n_offset + col_base + i * 8 + lane_col;if (col < N) {- const half h00 = __float2half_rn(tmp[idx + 0]);- const half h01 = __float2half_rn(tmp[idx + 1]);- const half h10 = __float2half_rn(tmp[idx + 2]);- const half h11 = __float2half_rn(tmp[idx + 3]);+ const half2 h2_row0 = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));+ const half2 h2_row1 = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));if (row0_in) {if (col + 1 < N) {- reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + col)[0] = __halves2half2(h00, h01);+ reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + col)[0] = h2_row0;} else {- C_ptr[row0 * Cs0 + col] = h00;+ C_ptr[row0 * Cs0 + col] = __low2half(h2_row0);}}if (row1_in) {if (col + 1 < N) {- reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + col)[0] = __halves2half2(h10, h11);+ reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + col)[0] = h2_row1;} else {- C_ptr[row1 * Cs0 + col] = h10;+ C_ptr[row1 * Cs0 + col] = __low2half(h2_row1);}}}⋯ 7 unchanged linesconst int idx = i * 4;const int col = n_offset + col_base + i * 8 + lane_col;if (col < N) {- const half h00 = __float2half_rn(tmp[idx + 0]);- const half h01 = __float2half_rn(tmp[idx + 1]);- const half h10 = __float2half_rn(tmp[idx + 2]);- const half h11 = __float2half_rn(tmp[idx + 3]);+ const half2 h2_row0 = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));+ const half2 h2_row1 = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));+ const half h00 = __low2half(h2_row0);+ const half h01 = __high2half(h2_row0);+ const half h10 = __low2half(h2_row1);+ const half h11 = __high2half(h2_row1);if (row0_in) {C_ptr[row0 * Cs0 + col * Cs1] = h00;if (col + 1 < N) C_ptr[row0 * Cs0 + (col + 1) * Cs1] = h01;⋯ 9 unchanged lines}}- template <int BLOCK_N>- __device__ __forceinline__ void issue_tma_for_tile(+ template <int BLOCK_M, int BLOCK_N>+ __device__ __forceinline__ void epilogue_store_fulltile_contiguous(const ProblemInfo& prob,- int m_offset, int n_offset,- int k_iter, int stage,- int smem_base, int mbar_base,- bool do_A, bool do_B,- uint64_t cache_A, uint64_t cache_B+ int m_offset,+ int n_offset,+ int tid,+ int warp_id,+ int lane_id) {- constexpr int TMA_A_SMEM_BYTES_L = 128 * (TMA_BLOCK_K / 2);- constexpr int TMA_B_SMEM_BYTES_L = BLOCK_N * (TMA_BLOCK_K / 2);- constexpr int TMA_SFA_SMEM_BYTES_L = MMA_M * (TMA_BLOCK_K / 16);- constexpr int TMA_SFB_SMEM_BYTES_L = BLOCK_N * (TMA_BLOCK_K / 16);- constexpr int B_off_L = TMA_A_SMEM_BYTES_L;- constexpr int SFA_off_L = B_off_L + TMA_B_SMEM_BYTES_L;- constexpr int SFB_off_L = SFA_off_L + TMA_SFA_SMEM_BYTES_L;- constexpr int STAGE_SIZE_L = TMA_A_SMEM_BYTES_L + TMA_B_SMEM_BYTES_L + TMA_SFA_SMEM_BYTES_L + TMA_SFB_SMEM_BYTES_L;+ if (tid >= BLOCK_M) return;- const int K = prob.K;- const int mbar_addr = mbar_base + stage * 8;- const int stage_base = smem_base + stage * STAGE_SIZE_L;- const int off_k = k_iter * TMA_BLOCK_K;+ half* C_ptr = prob.C_ptr;+ const int64_t Cs0 = prob.Cs0;- int expect_bytes = 0;- if (do_A) {- tma_3d_gmem2smem<1>(stage_base, &prob.A_tmap, 0, m_offset, off_k / 256, mbar_addr, cache_A);+ const int warp_row_base = m_offset + warp_id * 32;- const int rest_k = K / 16 / 4;- const int k_blk = off_k / (16 * 4);- const char* SFA_src = prob.SFA_ptr + ((m_offset / 128) * rest_k + k_blk) * 512;- tma_gmem2smem(stage_base + SFA_off_L, SFA_src, TMA_SFA_SMEM_BYTES_L, mbar_addr, cache_A);+ // Lane mapping: each lane owns two columns (half2) and one of 8 rows.+ const int lane_row = lane_id >> 2;+ const int lane_col = (lane_id & 3) * 2;- expect_bytes = TMA_A_SMEM_BYTES_L + TMA_SFA_SMEM_BYTES_L;- } else if (do_B) {- if constexpr (BLOCK_N == 256) {- tma_3d_gmem2smem<1>(stage_base + B_off_L, &prob.B_tmap_256, 0, n_offset, off_k / 256, mbar_addr, cache_B);- } else {- tma_3d_gmem2smem<1>(stage_base + B_off_L, &prob.B_tmap, 0, n_offset, off_k / 256, mbar_addr, cache_B);- }+ constexpr int HALF_N = 128;+ const int halves = BLOCK_N / HALF_N;- const int rest_k = K / 16 / 4;- const int k_blk = off_k / (16 * 4);- if constexpr (BLOCK_N == 256) {- constexpr int SFB_HALF_BYTES_L = 128 * (TMA_BLOCK_K / 16);- const char* SFB_src0 = prob.SFB_ptr + ((n_offset / 128) * rest_k + k_blk) * 512;- const char* SFB_src1 = prob.SFB_ptr + (((n_offset / 128) + 1) * rest_k + k_blk) * 512;- tma_gmem2smem(stage_base + SFB_off_L, SFB_src0, SFB_HALF_BYTES_L, mbar_addr, cache_B);- tma_gmem2smem(stage_base + SFB_off_L + SFB_HALF_BYTES_L, SFB_src1, SFB_HALF_BYTES_L, mbar_addr, cache_B);- } else {- const char* SFB_src = prob.SFB_ptr + ((n_offset / 128) * rest_k + k_blk) * 512;- tma_gmem2smem(stage_base + SFB_off_L, SFB_src, TMA_SFB_SMEM_BYTES_L, mbar_addr, cache_B);- }+ #pragma unroll+ for (int m = 0; m < 2; ++m) {+ const int row0 = warp_row_base + m * 16 + lane_row;+ const int row1 = row0 + 8;- expect_bytes = TMA_B_SMEM_BYTES_L + TMA_SFB_SMEM_BYTES_L;- }+ #pragma unroll+ for (int half_idx = 0; half_idx < halves; ++half_idx) {+ float tmp[HALF_N / 2];+ const int col_base = half_idx * HALF_N;+ tcgen05_ld_16x256b_x16(tmp, warp_id * 32 + m * 16, col_base);+ asm volatile("tcgen05.wait::ld.sync.aligned;\\n");- mbarrier_arrive_expect_tx(mbar_addr, expect_bytes);+ half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset + col_base);+ half2* row1_ptr = reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset + col_base);+ #pragma unroll+ for (int i = 0; i < HALF_N / 8; i++) {+ const int idx = i * 4;+ const int col = i * 8 + lane_col;+ const int h2_idx = col >> 1;+ row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));+ row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));+ }+ }+ }}- template <bool PERSISTENT, int BLOCK_M, int BLOCK_N, int NS>+ template <bool PERSISTENT, int BLOCK_M, int BLOCK_N, int NS, int CLUSTER_SIZE = 1>__global__ __launch_bounds__(TMA_NUM_WARPS * WARP_SIZE)void grouped_gemm_kernel_v4(const ProblemInfo* __restrict__ global_probs,const WorkItem* __restrict__ work_items,- int num_items,- int* __restrict__ work_counter+ int num_items) {constexpr int TMA_A_SMEM_BYTES = BLOCK_M * (TMA_BLOCK_K / 2);constexpr int TMA_B_SMEM_BYTES = BLOCK_N * (TMA_BLOCK_K / 2);⋯ 4 unchanged linesconst int lane_id = tid % WARP_SIZE;const int warp_id = tid / WARP_SIZE;+ uint32_t cta_rank = 0;+ if constexpr (CLUSTER_SIZE > 1) {+ cta_rank = get_cluster_ctarank();+ }++ // Shared memory is used as a multi-stage ring buffer.+ // Per stage: [A tile][B tile][SFA][SFB]. After all stages we place mbarriers.extern __shared__ __align__(1024) char smem_ptr[];const int smem_base = static_cast<int>(__cvta_generic_to_shared(smem_ptr));⋯ 2 unchanged linesconstexpr int SFB_off = SFA_off + TMA_SFA_SMEM_BYTES;const int mbar_base = smem_base + STAGE_SIZE * NS;-++ // TMEM allocation is in columns. We need 2 columns per output column because+ // accumulators are fp32.constexpr int TMEM_COLS = BLOCK_N * 2;constexpr int SFA_tmem = BLOCK_N;constexpr int SFB_tmem = SFA_tmem + 4 * (TMA_BLOCK_K / MMA_K);⋯ 4 unchanged linesif (warp_id == 0) {asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem_base), "r"(TMEM_COLS));- }- else if (warp_id == 1 && elect_sync()) {+ } else if (warp_id == 1 && elect_sync()) {+ // Best-effort tensormap prefetch for the first few work items.+for (int i = 0; i < num_items && i < 8; ++i) {const ProblemInfo* prob = &global_probs[work_items[i].problem_idx];asm volatile("prefetch.tensormap [%0];" :: "l"(&prob->A_tmap) : "memory");⋯ 7 unchanged lines__syncthreads();__shared__ int shared_work_idx;- __shared__ int tma_start_k;int work_idx = blockIdx.x;if (work_idx < num_items) {if (tid == 0) {- tma_start_k = 0;for (int i = 0; i < NS; ++i) {- // Two producer warps (A/SFA and B/SFB) arrive on the TMA barrier.+ // mbarrier[stage]: TMA completion barrier.+ // Two producer warps arrive (A/SFA and B/SFB).mbarrier_init(mbar_base + i * 8, 2);++ // mbarrier[NS+stage]: stage reuse barrier.+ // The MMA warp commits once per stage.mbarrier_init(mbar_base + (NS + i) * 8, 1);++ if constexpr (CLUSTER_SIZE > 1) {+ if (cta_rank == 0) {+ // Only CTA rank 0 initializes the cluster-wide barrier used to+ // guard multicast stage reuse.+ mbarrier_init(mbar_base + (2*NS + i) * 8, CLUSTER_SIZE);+ }+ }}asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");}__syncthreads();}+ // Ensure all CTAs have initialized mbarriers before any multicast TMA.+ if constexpr (CLUSTER_SIZE > 1) {+ cluster_sync();+ }+while (work_idx < num_items) {const WorkItem& work = work_items[work_idx];const ProblemInfo& prob = global_probs[work.problem_idx];⋯ 7 unchanged linesconstexpr uint64_t cache_A = EVICT_LAST;constexpr uint64_t cache_B = EVICT_FIRST;- const bool do_A = (warp_id == TMA_WARP);- const bool do_B = (warp_id == TMA_WARP_B);+ const bool do_A = (warp_id == TMA_WARP);+ const bool do_B = (warp_id == TMA_WARP_B);- // If the previous tile's overlap already prefetched first NS stages,- // tma_start_k == NS and this loop is a no-op.- for (int k_iter = tma_start_k; k_iter < NS && k_iter < num_k_iters; k_iter++) {- issue_tma_for_tile<BLOCK_N>(prob, m_offset, n_offset, k_iter, k_iter,- smem_base, mbar_base, do_A, do_B, cache_A, cache_B);+ auto issue_tma = [&](int k_iter, int stage) {+ const int mbar_addr = mbar_base + stage * 8;+ const int stage_base = smem_base + stage * STAGE_SIZE;+ const int off_k = k_iter * TMA_BLOCK_K;++ // Program expect_tx before issuing any TMA that completes to this+ // barrier. A completion arriving before expect_tx is set can leave the+ // consumer stuck in mbarrier_wait.+ const int expect_bytes = do_A+ ? (TMA_A_SMEM_BYTES + TMA_SFA_SMEM_BYTES)+ : (TMA_B_SMEM_BYTES + TMA_SFB_SMEM_BYTES);+ mbarrier_arrive_expect_tx(mbar_addr, expect_bytes);++ if (do_A) {+ if constexpr (CLUSTER_SIZE > 1) {+ // Cluster path: CTA rank 0 multicasts A to the whole cluster.+ // dst and mbarrier are passed as shared::cluster addresses.+ if (cta_rank == 0) {+ uint16_t mc = (1 << CLUSTER_SIZE) - 1;+ int cluster_dst = mapa_cta_to_cluster(stage_base, 0);+ int cluster_mbar = mapa_cta_to_cluster(mbar_addr, 0);+ tma_3d_gmem2smem_multicast(cluster_dst, &prob.A_tmap, 0, m_offset, off_k / 256, cluster_mbar, mc);+ }+ } else {+ tma_3d_gmem2smem<1>(stage_base, &prob.A_tmap, 0, m_offset, off_k / 256, mbar_addr, cache_A);+ }++ // SFA scale blocks are indexed by (m_tile, k_blk) and stored as+ // 512B blocks (matching tcgen05_cp_nvfp4 granularity).+ const int rest_k = K / 16 / 4;+ const int k_blk = off_k / (16 * 4);+ const char* SFA_src = prob.SFA_ptr + ((m_offset / 128) * rest_k + k_blk) * 512;+ tma_gmem2smem(stage_base + SFA_off, SFA_src, TMA_SFA_SMEM_BYTES, mbar_addr, cache_A);+ } else if (do_B) {+ if constexpr (BLOCK_N == 256) {+ tma_3d_gmem2smem<1>(stage_base + B_off, &prob.B_tmap_256, 0, n_offset, off_k / 256, mbar_addr, cache_B);+ } else {+ tma_3d_gmem2smem<1>(stage_base + B_off, &prob.B_tmap, 0, n_offset, off_k / 256, mbar_addr, cache_B);+ }++ // SFB scale blocks are indexed by (n_tile, k_blk).+ const int rest_k = K / 16 / 4;+ const int k_blk = off_k / (16 * 4);+ if constexpr (BLOCK_N == 256) {+ constexpr int SFB_HALF_BYTES = 128 * (TMA_BLOCK_K / 16);+ const char* SFB_src0 = prob.SFB_ptr + ((n_offset / 128) * rest_k + k_blk) * 512;+ const char* SFB_src1 = prob.SFB_ptr + (((n_offset / 128) + 1) * rest_k + k_blk) * 512;+ tma_gmem2smem(stage_base + SFB_off, SFB_src0, SFB_HALF_BYTES, mbar_addr, cache_B);+ tma_gmem2smem(stage_base + SFB_off + SFB_HALF_BYTES, SFB_src1, SFB_HALF_BYTES, mbar_addr, cache_B);+ } else {+ const char* SFB_src = prob.SFB_ptr + ((n_offset / 128) * rest_k + k_blk) * 512;+ tma_gmem2smem(stage_base + SFB_off, SFB_src, TMA_SFB_SMEM_BYTES, mbar_addr, cache_B);+ }+ }+ };++ for (int k_iter = 0; k_iter < NS && k_iter < num_k_iters; k_iter++) {+ issue_tma(k_iter, k_iter);}+ int stage = 0;+ int mma_phase = 0;for (int k_iter = NS; k_iter < num_k_iters; k_iter++) {- const int stage = k_iter % NS;- const int mma_phase = (k_iter / NS - 1) % 2;- mbarrier_wait(mbar_base + (NS + stage) * 8, mma_phase);- issue_tma_for_tile<BLOCK_N>(prob, m_offset, n_offset, k_iter, stage,- smem_base, mbar_base, do_A, do_B, cache_A, cache_B);+ if constexpr (CLUSTER_SIZE > 1) {+ if (do_A && cta_rank == 0) {+ // A is shared across the cluster via multicast. Before reusing a+ // ring-buffer stage for the next multicast, rank 0 must wait for+ // all CTAs to finish consuming the current stage.+ mbarrier_wait(mbar_base + (2*NS + stage) * 8, mma_phase);+ } else {+ mbarrier_wait(mbar_base + (NS + stage) * 8, mma_phase);+ }+ } else {+ mbarrier_wait(mbar_base + (NS + stage) * 8, mma_phase);+ }+ issue_tma(k_iter, stage);+ stage++;+ if (stage == NS) {+ stage = 0;+ mma_phase ^= 1;+ }}}else if (warp_id == MMA_WARP && elect_sync()) {auto make_desc_AB = [](int addr) -> uint64_t {const int SBO = 8 * 128;- // Keep descriptor swizzle fixed (matches known-good layout).+ // Descriptor encoding is coupled to the shared-memory swizzle and the+ // tcgen05 operand layout. SBO matches 128B swizzle.return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);};auto make_desc_SF = [](int addr) -> uint64_t {- // Keep SF descriptor encoding fixed; tcgen05_cp_nvfp4 expects this layout.+ // Scale-factor loads use a different stride (16B) but the same address+ // encoding (16B units).const int SBO = 8 * 16;return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);};+ int stage = 0;+ int tma_phase = 0;for (int k_iter = 0; k_iter < num_k_iters; k_iter++) {- const int stage = k_iter % NS;- const int tma_phase = (k_iter / NS) % 2;mbarrier_wait(mbar_base + stage * 8, tma_phase);const int stage_base = smem_base + stage * STAGE_SIZE;⋯ 2 unchanged linesconst uint64_t SFA_desc = SF_desc + ((uint64_t)(stage_base + SFA_off) >> 4ULL);const uint64_t SFB_desc = SF_desc + ((uint64_t)(stage_base + SFB_off) >> 4ULL);+ // Copy scale factors from shared memory into TMEM.#pragma unrollfor (int k = 0; k < TMA_BLOCK_K / MMA_K; k++) {tcgen05_cp_nvfp4(SFA_tmem + k * 4, SFA_desc + (uint64_t)k * (512ULL >> 4ULL));⋯ 6 unchanged lines}}+ // MMA loop over the 256-wide K tile in 64-wide chunks.#pragma unrollfor (int k = 0; k < TMA_BLOCK_K / MMA_K; k++) {uint64_t a_desc = make_desc_AB(stage_base + k * 32);uint64_t b_desc = make_desc_AB(stage_base + B_off + k * 32);- const int scale_A_tmem = SFA_tmem + k * 4 + (work.tile_m % (MMA_M / BLOCK_M)) * (BLOCK_M / 32);+ const int scale_A_tmem = SFA_tmem + k * 4;int scale_B_tmem;if constexpr (BLOCK_N == 256) {scale_B_tmem = SFB_tmem + k * 8;⋯ 1 unchanged linesscale_B_tmem = SFB_tmem + k * 4;}+ // First MMA uses D=0, subsequent MMAs accumulate.const int enable_input_d = (k_iter == 0 && k == 0) ? 0 : 1;tcgen05_mma_nvfp4(a_desc, b_desc, idesc, scale_A_tmem, scale_B_tmem, enable_input_d);}tcgen05_commit(mbar_base + (NS + stage) * 8);+ if constexpr (CLUSTER_SIZE > 1) {+ int cm = mapa_cta_to_cluster(mbar_base + (2*NS + stage) * 8, 0);+ mbarrier_arrive_cluster(cm);+ }++ stage++;+ if (stage == NS) {+ stage = 0;+ tma_phase ^= 1;+ }}const int last_stage = (num_k_iters - 1) % NS;⋯ 2 unchanged lines}__syncthreads();+ asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");- if constexpr (PERSISTENT) {- // === OVERLAPPED: epilogue (warps 0-3) + next-tile TMA prefetch (warps 4,6) ===-- // Warps 0-3: fence + epilogue (reads TMEM, writes global C)- if (warp_id < 4) {- asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");- const bool low_m = (prob.M <= LOW_M_THRESHOLD);- if (low_m) {- epilogue_store<BLOCK_M, BLOCK_N, true>(prob, m_offset, n_offset, tid, warp_id, lane_id);- } else {- epilogue_store<BLOCK_M, BLOCK_N, false>(prob, m_offset, n_offset, tid, warp_id, lane_id);- }+ const bool short_k = (K <= 2048);+ const bool full_tile = (m_offset + BLOCK_M <= prob.M) && (n_offset + BLOCK_N <= prob.N);+ if (short_k && full_tile && prob.Cs1 == 1) {+ epilogue_store_fulltile_contiguous<BLOCK_M, BLOCK_N>(prob, m_offset, n_offset, tid, warp_id, lane_id);+ } else {+ const bool adaptive_m_epilogue = (prob.M <= LOW_M_THRESHOLD) || short_k;+ if (adaptive_m_epilogue) {+ epilogue_store<BLOCK_M, BLOCK_N, true>(prob, m_offset, n_offset, tid, warp_id, lane_id);+ } else {+ epilogue_store<BLOCK_M, BLOCK_N, false>(prob, m_offset, n_offset, tid, warp_id, lane_id);}+ }- // Warp 4 (TMA_WARP): get next work index + reinitialize barriers+ // In cluster mode we must synchronize across CTAs before re-initializing+ // mbarriers; __syncthreads is CTA-local and does not order cluster-wide+ // mbarrier arrivals.+ if constexpr (PERSISTENT && CLUSTER_SIZE > 1) {+ cluster_sync();+ }++ if constexpr (PERSISTENT) {if (warp_id == TMA_WARP && elect_sync()) {- if (work_counter != nullptr) {- shared_work_idx = atomicAdd(work_counter, 1);- } else {- shared_work_idx = work_idx + gridDim.x;- }- tma_start_k = 0;+ // Static grid-stride work distribution avoids global atomics and+ // smooths the tail when num_items slightly exceeds one wave.+ shared_work_idx = work_idx + gridDim.x;if (shared_work_idx < num_items) {for (int i = 0; i < NS; ++i) {mbarrier_init(mbar_base + i * 8, 2);mbarrier_init(mbar_base + (NS + i) * 8, 1);+ if constexpr (CLUSTER_SIZE > 1) {+ if (cta_rank == 0) {+ mbarrier_init(mbar_base + (2*NS + i) * 8, CLUSTER_SIZE);+ }+ }}asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");}}+ }- // Named barrier: sync warps 4 and 6 so warp 6 knows barriers are initialized- if (warp_id == TMA_WARP || warp_id == TMA_WARP_B) {- asm volatile("bar.sync 1, 64;" ::: "memory");- }+ if constexpr (PERSISTENT) {+ __syncthreads();+ }- // TMA warps: prefetch first NS stages for the next tile (overlapped with epilogue)- if ((warp_id == TMA_WARP || warp_id == TMA_WARP_B) && elect_sync()) {- if (shared_work_idx < num_items) {- const WorkItem& nw = work_items[shared_work_idx];- const ProblemInfo& np = global_probs[nw.problem_idx];- const int nm = nw.tile_m * BLOCK_M;- const int nn = nw.tile_n * BLOCK_N;- const int n_k_iters = np.K / TMA_BLOCK_K;+ // Cluster barrier after re-init: ensure all CTAs see re-initialized mbarriers.+ if constexpr (PERSISTENT && CLUSTER_SIZE > 1) {+ cluster_sync();+ }- const bool next_do_A = (warp_id == TMA_WARP);- const bool next_do_B = (warp_id == TMA_WARP_B);-- for (int ki = 0; ki < NS && ki < n_k_iters; ki++) {- issue_tma_for_tile<BLOCK_N>(np, nm, nn, ki, ki,- smem_base, mbar_base, next_do_A, next_do_B, EVICT_LAST, EVICT_FIRST);- }-- if (next_do_A) {- tma_start_k = (n_k_iters < NS) ? n_k_iters : NS;- }- }- }-- __syncthreads();+ if constexpr (PERSISTENT) {work_idx = shared_work_idx;} else {- // Non-persistent: standard epilogue, no overlap needed- asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");- const bool low_m = (prob.M <= LOW_M_THRESHOLD);- if (low_m) {- epilogue_store<BLOCK_M, BLOCK_N, true>(prob, m_offset, n_offset, tid, warp_id, lane_id);- } else {- epilogue_store<BLOCK_M, BLOCK_N, false>(prob, m_offset, n_offset, tid, warp_id, lane_id);- }- __syncthreads();break;}}⋯ 20 unchanged linesuint32_t boxDim[rank] = {256, shared_height, 1};uint32_t elementStrides[rank] = {1, 1, 1};- // Keep swizzle fixed (matches known-good layout).+ // Swizzle must match the shared-memory layout expected by tcgen05.constexpr CUtensorMapSwizzle swizzle = CU_TENSOR_MAP_SWIZZLE_128B;// Cache cuTensorMap templates by shape.⋯ 68 unchanged linesc10::cuda::CUDAGuard device_guard(dev);auto sizes_accessor = sizes_cpu.accessor<int64_t, 2>();- // SM count is used for tail-smoothing and occupancy-based launch shaping.- // Hardcode to the target device (B200 in Modal) to avoid stale per-thread caches.+ // SM count is used for occupancy-based launch shaping.+ // Hardcoded for the target environment (B200 / SM100) to keep the logic+ // simple and deterministic.constexpr int sm_count = 148;static bool attrs_set = false;if (!attrs_set) {- // Two pipeline depths: high (more overlap) and low (lower dynamic smem).- // Low variants are intended to allow 2 CTAs/SM when shared memory is the limiter.+ // Two pipeline depths:+ // - HI: deeper pipeline, more overlap, higher shared-memory footprint.+ // - LO: shallower pipeline, lower shared-memory footprint (can improve+ // occupancy when shared memory is the limiter).constexpr int NS_DEEP_HI = 6;constexpr int NS_DEEP_LO = 3;constexpr int NS_WIDE_HI = 4;⋯ 3 unchanged linesconstexpr int MBAR_BYTES_DEEP_LO = ((2 * NS_DEEP_LO * 8 + 63) & ~63);constexpr int MBAR_BYTES_WIDE_HI = ((2 * NS_WIDE_HI * 8 + 63) & ~63);constexpr int MBAR_BYTES_WIDE_LO = ((2 * NS_WIDE_LO * 8 + 63) & ~63);+ constexpr int MBAR_BYTES_WIDE_CLUSTER_HI = ((3 * NS_WIDE_HI * 8 + 63) & ~63);+ constexpr int MBAR_BYTES_WIDE_CLUSTER_LO = ((3 * NS_WIDE_LO * 8 + 63) & ~63);// BLOCK_N=128constexpr int STAGE_128_K256 = 128 * (256 / 2) + 128 * (256 / 2) + 128 * (256 / 16) + 128 * (256 / 16);⋯ 4 unchanged linesconstexpr int STAGE_256_K256 = 128 * (256 / 2) + 256 * (256 / 2) + 128 * (256 / 16) + 256 * (256 / 16);constexpr int SMEM_256_HI_K256 = STAGE_256_K256 * NS_WIDE_HI + MBAR_BYTES_WIDE_HI;constexpr int SMEM_256_LO_K256 = STAGE_256_K256 * NS_WIDE_LO + MBAR_BYTES_WIDE_LO;+ constexpr int SMEM_256_CLUSTER_HI_K256 = STAGE_256_K256 * NS_WIDE_HI + MBAR_BYTES_WIDE_CLUSTER_HI;+ constexpr int SMEM_256_CLUSTER_LO_K256 = STAGE_256_K256 * NS_WIDE_LO + MBAR_BYTES_WIDE_CLUSTER_LO;// Variants: {persistent} x {BLOCK_N} x {NS}// BLOCK_N=128⋯ 2 unchanged linesCUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 128, NS_DEEP_LO>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_128_LO_K256));CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_LO>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_128_LO_K256));- // BLOCK_N=256+ // BLOCK_N=256 without clusterCUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_HI_K256));CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_HI_K256));CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_LO>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_LO_K256));CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_LO_K256));+ // BLOCK_N=256 with cluster multicast (CLUSTER_SIZE=4)+ constexpr int CL = 4;+ CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_CLUSTER_HI_K256));+ CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_CLUSTER_HI_K256));+ CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_LO, CL>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_CLUSTER_LO_K256));+ CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO, CL>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_CLUSTER_LO_K256));+ CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeNonPortableClusterSizeAllowed, 1));+ CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeNonPortableClusterSizeAllowed, 1));+ CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_LO, CL>), cudaFuncAttributeNonPortableClusterSizeAllowed, 1));+ CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO, CL>), cudaFuncAttributeNonPortableClusterSizeAllowed, 1));+attrs_set = true;}⋯ 1 unchanged linesstatic thread_local bool occ_set = false;static thread_local int occ_128_hi = 0, occ_128_lo = 0;static thread_local int occ_256_hi = 0, occ_256_lo = 0;+ static thread_local int occ_256_cluster_hi = 0, occ_256_cluster_lo = 0;if (!occ_set) {constexpr int NS_DEEP_HI = 6;constexpr int NS_DEEP_LO = 3;⋯ 3 unchanged linesconstexpr int MBAR_BYTES_DEEP_LO = ((2 * NS_DEEP_LO * 8 + 63) & ~63);constexpr int MBAR_BYTES_WIDE_HI = ((2 * NS_WIDE_HI * 8 + 63) & ~63);constexpr int MBAR_BYTES_WIDE_LO = ((2 * NS_WIDE_LO * 8 + 63) & ~63);+ constexpr int MBAR_BYTES_WIDE_CLUSTER_HI = ((3 * NS_WIDE_HI * 8 + 63) & ~63);+ constexpr int MBAR_BYTES_WIDE_CLUSTER_LO = ((3 * NS_WIDE_LO * 8 + 63) & ~63);constexpr int STAGE_128_K256 = 128 * (256 / 2) + 128 * (256 / 2) + 128 * (256 / 16) + 128 * (256 / 16);constexpr int STAGE_256_K256 = 128 * (256 / 2) + 256 * (256 / 2) + 128 * (256 / 16) + 256 * (256 / 16);constexpr int SMEM_128_HI_K256 = STAGE_128_K256 * NS_DEEP_HI + MBAR_BYTES_DEEP_HI;constexpr int SMEM_128_LO_K256 = STAGE_128_K256 * NS_DEEP_LO + MBAR_BYTES_DEEP_LO;constexpr int SMEM_256_HI_K256 = STAGE_256_K256 * NS_WIDE_HI + MBAR_BYTES_WIDE_HI;constexpr int SMEM_256_LO_K256 = STAGE_256_K256 * NS_WIDE_LO + MBAR_BYTES_WIDE_LO;+ constexpr int SMEM_256_CLUSTER_HI_K256 = STAGE_256_K256 * NS_WIDE_HI + MBAR_BYTES_WIDE_CLUSTER_HI;+ constexpr int SMEM_256_CLUSTER_LO_K256 = STAGE_256_K256 * NS_WIDE_LO + MBAR_BYTES_WIDE_CLUSTER_LO;constexpr int THREADS = TMA_NUM_WARPS * WARP_SIZE;CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(⋯ 4 unchanged lines&occ_256_hi, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI>, THREADS, SMEM_256_HI_K256));CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(&occ_256_lo, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO>, THREADS, SMEM_256_LO_K256));+ CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(+ &occ_256_cluster_hi, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI, 4>, THREADS, SMEM_256_CLUSTER_HI_K256));+ CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(+ &occ_256_cluster_lo, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO, 4>, THREADS, SMEM_256_CLUSTER_LO_K256));- TORCH_CHECK(occ_128_hi > 0 && occ_128_lo > 0 && occ_256_hi > 0 && occ_256_lo > 0,+ TORCH_CHECK(occ_128_hi > 0 && occ_128_lo > 0 && occ_256_hi > 0 && occ_256_lo > 0+ && occ_256_cluster_hi > 0 && occ_256_cluster_lo > 0,"occupancy query returned zero blocks/SM");occ_set = true;}std::vector<ProblemInfo> problem_infos(G);- std::vector<WorkItem> work_items_128;- std::vector<WorkItem> work_items_256;+ static thread_local std::vector<WorkItem> cached_work_items_128;+ static thread_local std::vector<WorkItem> cached_work_items_256;+ static thread_local std::vector<WorkItem> cached_work_items_256_cluster;+ static thread_local uint64_t cached_work_hash = 0;+ static thread_local bool cached_work_valid = false;std::vector<uint8_t> active(G, 0);std::vector<uint8_t> use_256(G, 0);+ std::vector<uint8_t> use_256_cluster(G, 0);std::vector<int> num_tiles_m(G, 0);std::vector<int> num_tiles_n(G, 0);std::vector<int64_t> Ms(G, 0), Ns(G, 0), Ks(G, 0);+ int64_t total_tiles = 0;for (int64_t i = 0; i < G; i++) {const int64_t M = sizes_accessor[i][0], N = sizes_accessor[i][1], K = sizes_accessor[i][2];Ms[(size_t)i] = M; Ns[(size_t)i] = N; Ks[(size_t)i] = K;⋯ 2 unchanged lines}active[(size_t)i] = 1;bool is_256 = (N >= 4096) && (K >= 2048) && ((N & 255) == 0);+ bool is_256_cluster = is_256 && (K > 2048);use_256[(size_t)i] = is_256;+ use_256_cluster[(size_t)i] = is_256_cluster;int block_m = 128;int block_n = is_256 ? 256 : 128;num_tiles_m[(size_t)i] = ceil_div((int)M, block_m);num_tiles_n[(size_t)i] = ceil_div((int)N, block_n);++ total_tiles += (int64_t)num_tiles_m[(size_t)i] * (int64_t)num_tiles_n[(size_t)i];}- const int tma_block_k = 256;+ constexpr int tma_block_k = 256;- work_items_128.reserve(G * 32);- work_items_256.reserve(G * 32);-- // Cross-group N-strip interleaving: iterate by tile_n first across all- // groups so that CTAs in the same wave load the same B N-strip from- // different groups, maximizing L2 hit rate for B data.- int max_tn_128 = 0, max_tn_256 = 0;+ uint64_t work_hash = 1469598103934665603ULL;+ work_hash = hash_combine_u64(work_hash, (uint64_t)tma_block_k);for (int64_t i = 0; i < G; i++) {- if (!active[(size_t)i]) continue;- if (use_256[(size_t)i]) {- max_tn_256 = std::max(max_tn_256, num_tiles_n[(size_t)i]);- } else {- max_tn_128 = std::max(max_tn_128, num_tiles_n[(size_t)i]);- }+ const bool is_active = (active[(size_t)i] != 0);+ if (!is_active) {+ work_hash = hash_combine_u64(work_hash, 0);+ continue;+ }+ const int64_t M = Ms[(size_t)i];+ const int64_t N = Ns[(size_t)i];+ const bool is_256 = (use_256[(size_t)i] != 0);+ const bool is_256_cluster = (use_256_cluster[(size_t)i] != 0);++ work_hash = hash_combine_u64(work_hash, (uint64_t)M);+ work_hash = hash_combine_u64(work_hash, (uint64_t)N);+ work_hash = hash_combine_u64(work_hash, (uint64_t)is_256);+ work_hash = hash_combine_u64(work_hash, (uint64_t)is_256_cluster);}- // 128-wide tiles: interleave by tile_n across groups- for (int tn = 0; tn < max_tn_128; tn++) {+ if (!cached_work_valid || cached_work_hash != work_hash) {+ cached_work_items_128.clear();+ cached_work_items_256.clear();+ cached_work_items_256_cluster.clear();+ cached_work_items_128.reserve(G * 32);+ cached_work_items_256.reserve(G * 32);+ cached_work_items_256_cluster.reserve(G * 32);++ // Work scheduling.+ // 128-wide path: iterate by tile_n across groups so a wave tends to touch+ // the same B strip across different problems (better L2 locality).+ int max_tn_128 = 0;for (int64_t i = 0; i < G; i++) {- if (!active[(size_t)i] || use_256[(size_t)i]) continue;- if (tn >= num_tiles_n[(size_t)i]) continue;- for (int tm = 0; tm < num_tiles_m[(size_t)i]; tm++) {- work_items_128.push_back({(int)i, tm, tn});+ if (!active[(size_t)i]) continue;+ if (!use_256[(size_t)i]) {+ max_tn_128 = std::max(max_tn_128, num_tiles_n[(size_t)i]);}}- }- // 256-wide tiles: interleave by tile_n across groups- for (int tn = 0; tn < max_tn_256; tn++) {+ // 128-wide tiles: interleave by tile_n across groups+ for (int tn = 0; tn < max_tn_128; tn++) {+ for (int64_t i = 0; i < G; i++) {+ if (!active[(size_t)i] || use_256[(size_t)i]) continue;+ if (tn >= num_tiles_n[(size_t)i]) continue;+ for (int tm = 0; tm < num_tiles_m[(size_t)i]; tm++) {+ cached_work_items_128.push_back({(int)i, tm, tn});+ }+ }+ }++ // 256-wide path: split by K.+ // - K > 2048: clustered multicast path.+ // - K <= 2048: non-cluster path to avoid cluster overhead.+ constexpr int CLUSTER_SIZE_256 = 4;for (int64_t i = 0; i < G; i++) {if (!active[(size_t)i] || !use_256[(size_t)i]) continue;- if (tn >= num_tiles_n[(size_t)i]) continue;for (int tm = 0; tm < num_tiles_m[(size_t)i]; tm++) {- work_items_256.push_back({(int)i, tm, tn});+ for (int tn = 0; tn < num_tiles_n[(size_t)i]; tn++) {+ if (use_256_cluster[(size_t)i]) {+ cached_work_items_256_cluster.push_back({(int)i, tm, tn});+ } else {+ cached_work_items_256.push_back({(int)i, tm, tn});+ }+ }+ if (use_256_cluster[(size_t)i]) {+ // Pad clustered work to a multiple of CLUSTER_SIZE so every+ // launched cluster is full. Duplicating an existing tile is safe:+ // it deterministically writes the same result.+ int remainder = num_tiles_n[(size_t)i] % CLUSTER_SIZE_256;+ if (remainder != 0) {+ for (int p = 0; p < CLUSTER_SIZE_256 - remainder; p++) {+ cached_work_items_256_cluster.push_back({(int)i, tm, 0});+ }+ }+ }}}++ cached_work_hash = work_hash;+ cached_work_valid = true;}+ uint64_t probs_hash = 1469598103934665603ULL;+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)tma_block_k);for (int64_t i = 0; i < G; i++) {if (!active[(size_t)i]) continue;const int64_t M = Ms[(size_t)i], N = Ns[(size_t)i], K = Ks[(size_t)i];⋯ 12 unchanged lines} else {p.B_tmap_256 = p.B_tmap;}++ probs_hash = hash_combine_u64(probs_hash, (uint64_t)i);+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)M);+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)N);+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)K);+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)A_list[i].data_ptr());+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)B_list[i].data_ptr());+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)p.C_ptr);+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)p.SFA_ptr);+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)p.SFB_ptr);+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)p.Cs0);+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)p.Cs1);}- if (work_items_128.empty() && work_items_256.empty()) return C_list;+ if (cached_work_items_128.empty() && cached_work_items_256.empty() && cached_work_items_256_cluster.empty()) return C_list;auto options = at::TensorOptions().dtype(at::kByte).device(dev);static thread_local at::Tensor d_probs_cache;static thread_local at::Tensor d_work_cache_128;static thread_local at::Tensor d_work_cache_256;+ static thread_local at::Tensor d_work_cache_256_cluster;+ static thread_local uint64_t last_probs_hash = 0;+ static thread_local uint64_t last_work_hash = 0;+ static thread_local bool last_hash_valid = false;const int64_t probs_bytes = (int64_t)(G * sizeof(ProblemInfo));- const int64_t work_bytes_128 = (int64_t)(work_items_128.size() * sizeof(WorkItem));- const int64_t work_bytes_256 = (int64_t)(work_items_256.size() * sizeof(WorkItem));+ const int64_t work_bytes_128 = (int64_t)(cached_work_items_128.size() * sizeof(WorkItem));+ const int64_t work_bytes_256 = (int64_t)(cached_work_items_256.size() * sizeof(WorkItem));+ const int64_t work_bytes_256_cluster = (int64_t)(cached_work_items_256_cluster.size() * sizeof(WorkItem));+ bool probs_realloc = false;if (!d_probs_cache.defined() || d_probs_cache.device() != dev || d_probs_cache.scalar_type() != at::kByte || d_probs_cache.numel() < probs_bytes) {d_probs_cache = at::empty({probs_bytes}, options);+ probs_realloc = true;}if (work_bytes_128 > 0 && (!d_work_cache_128.defined() || d_work_cache_128.device() != dev || d_work_cache_128.scalar_type() != at::kByte || d_work_cache_128.numel() < work_bytes_128)) {d_work_cache_128 = at::empty({work_bytes_128}, options);⋯ 1 unchanged linesif (work_bytes_256 > 0 && (!d_work_cache_256.defined() || d_work_cache_256.device() != dev || d_work_cache_256.scalar_type() != at::kByte || d_work_cache_256.numel() < work_bytes_256)) {d_work_cache_256 = at::empty({work_bytes_256}, options);}+ if (work_bytes_256_cluster > 0 && (!d_work_cache_256_cluster.defined() || d_work_cache_256_cluster.device() != dev || d_work_cache_256_cluster.scalar_type() != at::kByte || d_work_cache_256_cluster.numel() < work_bytes_256_cluster)) {+ d_work_cache_256_cluster = at::empty({work_bytes_256_cluster}, options);+ }- CUDA_CHECK(cudaMemcpyAsync(d_probs_cache.data_ptr(), problem_infos.data(), G * sizeof(ProblemInfo), cudaMemcpyHostToDevice));- if (work_bytes_128 > 0) {- CUDA_CHECK(cudaMemcpyAsync(d_work_cache_128.data_ptr(), work_items_128.data(), work_bytes_128, cudaMemcpyHostToDevice));+ if (probs_realloc || !last_hash_valid || last_probs_hash != probs_hash) {+ CUDA_CHECK(cudaMemcpyAsync(d_probs_cache.data_ptr(), problem_infos.data(), G * sizeof(ProblemInfo), cudaMemcpyHostToDevice));+ last_probs_hash = probs_hash;}- if (work_bytes_256 > 0) {- CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256.data_ptr(), work_items_256.data(), work_bytes_256, cudaMemcpyHostToDevice));+ if (!last_hash_valid || last_work_hash != work_hash) {+ if (work_bytes_128 > 0) {+ CUDA_CHECK(cudaMemcpyAsync(d_work_cache_128.data_ptr(), cached_work_items_128.data(), work_bytes_128, cudaMemcpyHostToDevice));+ }+ if (work_bytes_256 > 0) {+ CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256.data_ptr(), cached_work_items_256.data(), work_bytes_256, cudaMemcpyHostToDevice));+ }+ if (work_bytes_256_cluster > 0) {+ CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256_cluster.data_ptr(), cached_work_items_256_cluster.data(), work_bytes_256_cluster, cudaMemcpyHostToDevice));+ }+ last_work_hash = work_hash;}+ last_hash_valid = true;constexpr int NS_DEEP_HI = 6;constexpr int NS_DEEP_LO = 3;constexpr int NS_WIDE_HI = 4;constexpr int NS_WIDE_LO = 2;- if (!work_items_128.empty()) {- int num_items_128 = (int)work_items_128.size();+ if (!cached_work_items_128.empty()) {+ int num_items_128 = (int)cached_work_items_128.size();// Choose pipeline depth by expected waves.const int wave_hi = sm_count * occ_128_hi;const int wave_lo = sm_count * occ_128_lo;// Only switch to low-smem variant when the grid is large enough that reducing// waves is likely to outweigh reduced pipeline overlap.- const bool use_lo = (tma_block_k == 256) && (wave_lo > wave_hi) && (num_items_128 > 2 * wave_hi);+ const bool use_lo = (wave_lo > wave_hi) && (num_items_128 > 2 * wave_hi);const int ns = use_lo ? NS_DEEP_LO : NS_DEEP_HI;const int occ = use_lo ? occ_128_lo : occ_128_hi;⋯ 8 unchanged linesconst bool persistent_128 = (num_items_128 > wave_cap);const int launch_ctas_128 = persistent_128 ? wave_cap : num_items_128;- if (tma_block_k == 256) {- if (persistent_128) {- if (use_lo) {- grouped_gemm_kernel_v4<true, 128, 128, NS_DEEP_LO><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(- (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128, nullptr);- } else {- grouped_gemm_kernel_v4<true, 128, 128, NS_DEEP_HI><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(- (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128, nullptr);- }+ if (persistent_128) {+ if (use_lo) {+ grouped_gemm_kernel_v4<true, 128, 128, NS_DEEP_LO><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(+ (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128);} else {- if (use_lo) {- grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_LO><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(- (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128, nullptr);- } else {- grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_HI><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(- (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128, nullptr);- }+ grouped_gemm_kernel_v4<true, 128, 128, NS_DEEP_HI><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(+ (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128);}+ } else {+ if (use_lo) {+ grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_LO><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(+ (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128);+ } else {+ grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_HI><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(+ (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128);+ }}}- if (!work_items_256.empty()) {- int num_items_256 = (int)work_items_256.size();+ if (!cached_work_items_256.empty()) {+ int num_items_256 = (int)cached_work_items_256.size();const int wave_hi = sm_count * occ_256_hi;const int wave_lo = sm_count * occ_256_lo;- const bool use_lo = (tma_block_k == 256) && (wave_lo > wave_hi) && (num_items_256 > 2 * wave_hi);+ const bool use_lo = (wave_lo > wave_hi) && (num_items_256 > 2 * wave_hi);const int ns = use_lo ? NS_WIDE_LO : NS_WIDE_HI;const int occ = use_lo ? occ_256_lo : occ_256_hi;⋯ 7 unchanged linesconst bool persistent_256 = (num_items_256 > wave_cap);const int launch_ctas_256 = persistent_256 ? wave_cap : num_items_256;- if (tma_block_k == 256) {- if (persistent_256) {- if (use_lo) {- grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_LO><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(- (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256, nullptr);- } else {- grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(- (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256, nullptr);- }+ if (persistent_256) {+ if (use_lo) {+ grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_LO><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(+ (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256);} else {- if (use_lo) {- grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(- (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256, nullptr);- } else {- grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(- (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256, nullptr);- }+ grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(+ (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256);}+ } else {+ if (use_lo) {+ grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_LO><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(+ (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256);+ } else {+ grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(+ (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256);+ }}}++ if (!cached_work_items_256_cluster.empty()) {+ int num_items_256_cluster = (int)cached_work_items_256_cluster.size();+ constexpr int CL256 = 4;++ const int wave_hi = sm_count * occ_256_cluster_hi;+ const int wave_lo = sm_count * occ_256_cluster_lo;+ const bool use_lo = (wave_lo > wave_hi) && (num_items_256_cluster > 2 * wave_hi);+⋯ diff truncated
scrolls · 1201 diff lines total
Best evidence level for this revision: reported
JSON