submission 493749
XoTic · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1041 lines, June 9 Researcher Reciprocity License v1.0.
v4b.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-493749?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:0bc24e80860c3f5ead9b6776601a961277edb00cf3808f88e0fb6b21872a03e4
license declaredunknown
license concludedunknown
authorsXoTic
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>shared-memory
extern __shared__ __align__(1024) char smem_ptr[];stages = 4
constexpr int NUM_STAGES = 4;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
v4b.py1041 lines
#!POPCORN leaderboard nvfp4_group_gemm
#!POPCORN gpu NVIDIA
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
"""
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
⏱ 89.3 ± 0.09 µs
⚡ 88.7 µs 🐌 89.6 µ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
⏱ 82.7 ± 0.05 µs
⚡ 82.7 µs 🐌 82.8 µs
g: 2; k: [4096, 4096]; m: [192, 320]; n: [3072, 3072]; seed: 1111
⏱ 29.5 ± 0.03 µs
⚡ 29.1 µs 🐌 29.8 µs
g: 2; k: [1536, 1536]; m: [128, 384]; n: [4096, 4096]; seed: 1111
⏱ 16.4 ± 0.02 µs
⚡ 16.0 µs 🐌 16.6 µs
"""
CUDA_SRC = """
#include <vector>
#include <unordered_map>
#include <cstdint>
#include <cstdio>
#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) {
// 64-bit FNV-1a variant
h ^= x;
h *= 1099511628211ULL;
return h;
}
constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;
// Cache hints
constexpr uint64_t EVICT_NORMAL = 0x1000000000000000;
constexpr uint64_t EVICT_FIRST = 0x12F0000000000000;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000;
// Work item for persistent kernel
struct WorkItem {
int problem_idx;
int tile_m;
int tile_n;
};
// Global Problem Info stored in Global Memory
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, Cs2;
};
__device__ inline
constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; };
__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 void mbarrier_init(int mbar_addr, int count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
}
__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)
);
}
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 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)
);
}
// TMEM load helpers
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"
);
}
// ============================================================================
// KERNEL CONFIGURATION
// ============================================================================
constexpr int TMA_BLOCK_N = 128;
constexpr int TMA_BLOCK_K = 256;
constexpr int NUM_STAGES = 4;
constexpr int TMA_NUM_WARPS = 8;
constexpr int MMA_M = 128;
constexpr int MBAR_BYTES = ((2 * NUM_STAGES * 8 + 63) & ~63);
constexpr int LOW_M_THRESHOLD = 96;
// Warp assignments (8 warps total):
// Warp 0-3: epilogue helpers
// Warp 4: TMA producer
// Warp 5: MMA consumer (single warp issues tcgen05.mma.cta_group::1)
// Warp 6-7: additional helpers
constexpr int TMA_WARP = 4;
constexpr int MMA_WARP = 5;
// ============================================================================
// EPILOGUE
// ============================================================================
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;
}
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;
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;
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] = __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]));
}
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] = __halves2half2(__float2half_rn(tmp[idx + 0]), __float2half_rn(tmp[idx + 1]));
}
if (row1_in) {
row1_ptr[h2_idx] = __halves2half2(__float2half_rn(tmp[idx + 2]), __float2half_rn(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 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]);
if (row0_in) {
if (col + 1 < N) {
reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + col)[0] = __halves2half2(h00, h01);
} else {
C_ptr[row0 * Cs0 + col] = h00;
}
}
if (row1_in) {
if (col + 1 < N) {
reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + col)[0] = __halves2half2(h10, h11);
} else {
C_ptr[row1 * Cs0 + col] = h10;
}
}
}
}
}
} 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 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]);
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;
}
}
}
}
}
}
}
// ============================================================================
// MAIN KERNEL
// ============================================================================
template <bool PERSISTENT, int BLOCK_M, int BLOCK_N>
__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
) {
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;
// Shared Memory Setup
extern __shared__ __align__(1024) char smem_ptr[];
const int smem_base = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
// Offsets within each stage
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;
// Mbarriers
const int mbar_base = smem_base + STAGE_SIZE * NUM_STAGES;
// TMEM addresses
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);
// Allocate TMEM
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()) {
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();
// Persistent work counter
__shared__ int shared_work_idx;
// Get first work item
int work_idx;
if constexpr (PERSISTENT) {
if (tid == 0) {
shared_work_idx = atomicAdd(work_counter, 1);
}
__syncthreads();
work_idx = shared_work_idx;
} else {
work_idx = blockIdx.x;
}
// Main processing loop
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;
// Initialize mbarriers per tile
if (tid == 0) {
for (int i = 0; i < NUM_STAGES; ++i) {
mbarrier_init(mbar_base + i * 8, 1);
mbarrier_init(mbar_base + (NUM_STAGES + i) * 8, 1);
}
asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
}
__syncthreads();
// ================================================================
// PRODUCER WARP: TMA (warp 4)
// ================================================================
if (warp_id == TMA_WARP && elect_sync()) {
// With M-tile-major work ordering, A stays in L2 across N-tiles
constexpr uint64_t cache_A = EVICT_LAST;
constexpr uint64_t cache_B = EVICT_FIRST;
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;
// TMA loads
tma_3d_gmem2smem<1>(stage_base, &prob.A_tmap, 0, m_offset, off_k / 256, mbar_addr, cache_A);
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);
}
// Scale factor loads
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);
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);
}
mbarrier_arrive_expect_tx(mbar_addr, STAGE_SIZE);
};
// Pipeline priming: issue first NUM_STAGES TMAs without waiting
for (int k_iter = 0; k_iter < NUM_STAGES && k_iter < num_k_iters; k_iter++) {
issue_tma(k_iter, k_iter);
}
// Steady state: wait for MMA, then issue TMA
for (int k_iter = NUM_STAGES; k_iter < num_k_iters; k_iter++) {
const int stage = k_iter % NUM_STAGES;
const int mma_phase = (k_iter / NUM_STAGES - 1) % 2;
mbarrier_wait(mbar_base + (NUM_STAGES + stage) * 8, mma_phase);
issue_tma(k_iter, stage);
}
}
// ================================================================
// CONSUMER WARP: MMA (warp 5)
// Single elected thread issues tcgen05.mma.cta_group::1
// ================================================================
else if (warp_id == MMA_WARP && elect_sync()) {
auto make_desc_AB = [](int addr) -> uint64_t {
const int SBO = 8 * 128;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
auto make_desc_SF = [](int addr) -> uint64_t {
const int SBO = 8 * 16;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
};
for (int k_iter = 0; k_iter < num_k_iters; k_iter++) {
const int stage = k_iter % NUM_STAGES;
const int tma_phase = (k_iter / NUM_STAGES) % 2;
mbarrier_wait(mbar_base + stage * 8, tma_phase);
const int stage_base = smem_base + stage * STAGE_SIZE;
// Copy scale factors to TMEM
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);
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));
}
}
// Issue MMA
for (int k1 = 0; k1 < TMA_BLOCK_K / 256; k1++) {
for (int k2 = 0; k2 < 256 / MMA_K; k2++) {
uint64_t a_desc = make_desc_AB(stage_base + k1 * BLOCK_M * 128 + k2 * 32);
uint64_t b_desc = make_desc_AB(stage_base + B_off + k1 * BLOCK_N * 128 + k2 * 32);
int k_sf = k1 * 4 + k2;
const int scale_A_tmem = SFA_tmem + k_sf * 4 + (work.tile_m % (MMA_M / BLOCK_M)) * (BLOCK_M / 32);
int scale_B_tmem;
if constexpr (BLOCK_N == 256) {
scale_B_tmem = SFB_tmem + k_sf * 8;
} else {
scale_B_tmem = SFB_tmem + k_sf * 4;
}
const int enable_input_d = (k_iter == 0 && k1 == 0 && k2 == 0) ? 0 : 1;
tcgen05_mma_nvfp4(a_desc, b_desc, idesc, scale_A_tmem, scale_B_tmem, enable_input_d);
}
}
tcgen05_commit(mbar_base + (NUM_STAGES + stage) * 8);
}
// Wait for final commit
if (num_k_iters > 0) {
const int last_stage = (num_k_iters - 1) % NUM_STAGES;
const int last_phase = ((num_k_iters - 1) / NUM_STAGES) % 2;
mbarrier_wait(mbar_base + (NUM_STAGES + last_stage) * 8, last_phase);
}
}
// ================================================================
// SYNCHRONIZATION & EPILOGUE
// ================================================================
__syncthreads();
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();
// Get next work item
if constexpr (PERSISTENT) {
if (tid == 0) {
shared_work_idx = atomicAdd(work_counter, 1);
}
__syncthreads();
work_idx = shared_work_idx;
} else {
break;
}
}
// Deallocate TMEM
if (warp_id == 0) {
tcgen05_dealloc_cols_cta1(0, TMEM_COLS);
}
}
// ============================================================================
// TENSOR MAP INITIALIZATION
// ============================================================================
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");
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, shared_width / 256};
uint32_t elementStrides[rank] = {1, 1, 1};
// cuTensorMapEncodeTiled is a relatively expensive driver call.
// Cache a per-shape template and then patch only the base address.
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, CU_TENSOR_MAP_SWIZZLE_128B,
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;
}
struct PadCacheEntry {
at::Tensor buf;
size_t zeroed_from = 0; // byte offset; bytes in [zeroed_from, end) are guaranteed zero
};
static at::Tensor pad_u4_tensor_cached(const at::Tensor& src, int64_t padded_m, PadCacheEntry* entry) {
TORCH_CHECK(entry != nullptr, "pad_u4_tensor_cached: entry is null");
if (src.size(0) == padded_m && ((uintptr_t)src.data_ptr() & 0xF) == 0) return src;
auto new_sizes = src.sizes().vec();
new_sizes[0] = padded_m;
bool reuse_ok = entry->buf.defined() && entry->buf.dim() == (int)new_sizes.size();
if (reuse_ok) {
// Allow reusing a larger leading-dimension buffer (avoid reallocs when M/N shrink),
// but require trailing dimensions to match exactly.
if (entry->buf.size(0) < padded_m) reuse_ok = false;
for (int d = 1; d < entry->buf.dim(); d++) {
if (entry->buf.size(d) != new_sizes[(size_t)d]) { reuse_ok = false; break; }
}
if (reuse_ok && (((uintptr_t)entry->buf.data_ptr() & 0xF) != 0)) reuse_ok = false;
}
const bool need_new =
!entry->buf.defined() ||
entry->buf.device() != src.device() ||
entry->buf.scalar_type() != src.scalar_type() ||
!reuse_ok;
if (need_new) {
entry->buf = at::empty(new_sizes, src.options());
entry->zeroed_from = (size_t)entry->buf.nbytes(); // nothing guaranteed yet
}
const size_t copy_bytes = (size_t)src.nbytes();
const size_t total_bytes = (size_t)entry->buf.nbytes();
TORCH_CHECK(copy_bytes <= total_bytes, "pad_u4_tensor_cached: size mismatch");
CUDA_CHECK(cudaMemcpyAsync(entry->buf.data_ptr(), src.data_ptr(), copy_bytes, cudaMemcpyDeviceToDevice));
if (copy_bytes < entry->zeroed_from) {
CUDA_CHECK(cudaMemsetAsync((char*)entry->buf.data_ptr() + copy_bytes, 0, entry->zeroed_from - copy_bytes));
entry->zeroed_from = copy_bytes;
} else {
entry->zeroed_from = copy_bytes;
}
return entry->buf;
}
// ============================================================================
// HOST ENTRY POINT
// ============================================================================
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>();
static bool attrs_set = false;
if (!attrs_set) {
constexpr int STAGE_SIZE_128 = 128 * 128 + 128 * 128 + 128 * 16 + 128 * 16;
constexpr int SMEM_SIZE_128 = STAGE_SIZE_128 * NUM_STAGES + MBAR_BYTES;
constexpr int STAGE_SIZE_64 = 64 * 128 + 128 * 128 + 128 * 16 + 128 * 16;
constexpr int SMEM_SIZE_64 = STAGE_SIZE_64 * NUM_STAGES + MBAR_BYTES;
constexpr int STAGE_SIZE_256 = 128 * 128 + 256 * 128 + 128 * 16 + 256 * 16;
constexpr int SMEM_SIZE_256 = STAGE_SIZE_256 * NUM_STAGES + MBAR_BYTES;
CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<true, 128, 128>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_128));
CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<false, 128, 128>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_128));
CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<true, 64, 128>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_64));
CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<false, 64, 128>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_64));
CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<true, 128, 256>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_256));
CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<false, 128, 256>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_256));
attrs_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_64;
static thread_local std::vector<WorkItem> cached_work_items_256;
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_64(G, 0);
std::vector<uint8_t> use_256(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);
uint64_t work_hash = 1469598103934665603ULL;
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) {
work_hash = hash_combine_u64(work_hash, 0);
continue;
}
active[(size_t)i] = 1;
bool is_64 = (M <= 64) && (N <= 2048);
bool is_256 = (!is_64) && (N >= 4096) && (K >= 2048) && ((N & 255) == 0);
use_64[(size_t)i] = is_64;
use_256[(size_t)i] = is_256;
int block_m = is_64 ? 64 : 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);
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_64);
work_hash = hash_combine_u64(work_hash, (uint64_t)is_256);
}
if (!cached_work_valid || cached_work_hash != work_hash) {
cached_work_items_128.clear();
cached_work_items_64.clear();
cached_work_items_256.clear();
cached_work_items_128.reserve(G * 32);
cached_work_items_64.reserve(G * 32);
cached_work_items_256.reserve(G * 32);
for (int64_t i = 0; i < G; i++) {
if (!active[(size_t)i]) continue;
// M-tile-major ordering: consecutive work items share the same M-tile
// so A data stays hot in L2 while iterating N-tiles
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_64[(size_t)i]) {
cached_work_items_64.push_back({(int)i, tm, tn});
} else if (use_256[(size_t)i]) {
cached_work_items_256.push_back({(int)i, tm, tn});
} else {
cached_work_items_128.push_back({(int)i, tm, tn});
}
}
}
}
cached_work_hash = work_hash;
cached_work_valid = true;
}
static thread_local std::vector<PadCacheEntry> A_pad_cache;
static thread_local std::vector<PadCacheEntry> B_pad_cache;
if ((int64_t)A_pad_cache.size() < G) A_pad_cache.resize((size_t)G);
if ((int64_t)B_pad_cache.size() < G) B_pad_cache.resize((size_t)G);
uint64_t probs_hash = 1469598103934665603ULL;
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];
int block_m = use_64[(size_t)i] ? 64 : 128;
int block_n = use_256[(size_t)i] ? 256 : 128;
int64_t padded_M = ((M + block_m - 1) / block_m) * block_m;
int64_t padded_N = ((N + block_n - 1) / block_n) * block_n;
at::Tensor A = pad_u4_tensor_cached(A_list[i], padded_M, &A_pad_cache[(size_t)i]);
at::Tensor B = pad_u4_tensor_cached(B_list[i], padded_N, &B_pad_cache[(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.Cs2 = C_list[i].stride(2);
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.data_ptr(), A.size(0), K, block_m, TMA_BLOCK_K);
init_AB_tmap_u4(&p.B_tmap, B.data_ptr(), B.size(0), K, 128, TMA_BLOCK_K);
if (use_256[(size_t)i]) {
init_AB_tmap_u4(&p.B_tmap_256, B.data_ptr(), B.size(0), K, 256, TMA_BLOCK_K);
} 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.data_ptr());
probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)B.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);
probs_hash = hash_combine_u64(probs_hash, (uint64_t)p.Cs2);
}
if (cached_work_items_128.empty() && cached_work_items_64.empty() && cached_work_items_256.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_64;
static thread_local at::Tensor d_work_cache_256;
static thread_local at::Tensor d_counter_cache;
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_64 = (int64_t)(cached_work_items_64.size() * sizeof(WorkItem));
const int64_t work_bytes_256 = (int64_t)(cached_work_items_256.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_64 > 0 && (!d_work_cache_64.defined() || d_work_cache_64.device() != dev || d_work_cache_64.scalar_type() != at::kByte || d_work_cache_64.numel() < work_bytes_64)) {
d_work_cache_64 = at::empty({work_bytes_64}, 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 (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_64 > 0) {
CUDA_CHECK(cudaMemcpyAsync(d_work_cache_64.data_ptr(), cached_work_items_64.data(), work_bytes_64, cudaMemcpyHostToDevice));
}
if (work_bytes_256 > 0) {
CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256.data_ptr(), cached_work_items_256.data(), work_bytes_256, cudaMemcpyHostToDevice));
}
last_work_hash = work_hash;
}
last_hash_valid = true;
constexpr int MAX_CTAS = 264;
if (!d_counter_cache.defined() || d_counter_cache.device() != dev || d_counter_cache.scalar_type() != at::kInt || d_counter_cache.numel() != 1) {
d_counter_cache = at::empty({1}, options.dtype(at::kInt));
}
if (!cached_work_items_128.empty()) {
int num_items_128 = (int)cached_work_items_128.size();
constexpr int STAGE_SIZE_128 = 128 * 128 + 128 * 128 + 128 * 16 + 128 * 16;
constexpr int SMEM_SIZE_128 = STAGE_SIZE_128 * NUM_STAGES + MBAR_BYTES;
if (num_items_128 > MAX_CTAS) {
CUDA_CHECK(cudaMemsetAsync(d_counter_cache.data_ptr(), 0, sizeof(int)));
grouped_gemm_kernel_v4<true, 128, 128><<<MAX_CTAS, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128, (int*)d_counter_cache.data_ptr());
} else {
grouped_gemm_kernel_v4<false, 128, 128><<<num_items_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 (!cached_work_items_64.empty()) {
int num_items_64 = (int)cached_work_items_64.size();
constexpr int STAGE_SIZE_64 = 64 * 128 + 128 * 128 + 128 * 16 + 128 * 16;
constexpr int SMEM_SIZE_64 = STAGE_SIZE_64 * NUM_STAGES + MBAR_BYTES;
if (num_items_64 > MAX_CTAS) {
CUDA_CHECK(cudaMemsetAsync(d_counter_cache.data_ptr(), 0, sizeof(int)));
grouped_gemm_kernel_v4<true, 64, 128><<<MAX_CTAS, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_64>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_64.data_ptr(), num_items_64, (int*)d_counter_cache.data_ptr());
} else {
grouped_gemm_kernel_v4<false, 64, 128><<<num_items_64, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_64>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_64.data_ptr(), num_items_64, nullptr);
}
}
if (!cached_work_items_256.empty()) {
int num_items_256 = (int)cached_work_items_256.size();
constexpr int STAGE_SIZE_256 = 128 * 128 + 256 * 128 + 128 * 16 + 256 * 16;
constexpr int SMEM_SIZE_256 = STAGE_SIZE_256 * NUM_STAGES + MBAR_BYTES;
if (num_items_256 > MAX_CTAS) {
CUDA_CHECK(cudaMemsetAsync(d_counter_cache.data_ptr(), 0, sizeof(int)));
grouped_gemm_kernel_v4<true, 128, 256><<<MAX_CTAS, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(
(ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256, (int*)d_counter_cache.data_ptr());
} else {
grouped_gemm_kernel_v4<false, 128, 256><<<num_items_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);
}
}
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
def custom_kernel(data: input_t) -> output_t:
abc_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]
global _SIZES_CPU_CACHE
try:
_SIZES_CPU_CACHE
except NameError:
_SIZES_CPU_CACHE = {}
key = tuple(tuple(x) for x in problem_sizes)
sizes_cpu = _SIZES_CPU_CACHE.get(key)
if sizes_cpu is None:
sizes_cpu = torch.tensor(key, dtype=torch.int64, device="cpu")
if len(_SIZES_CPU_CACHE) > 128:
_SIZES_CPU_CACHE.clear()
_SIZES_CPU_CACHE[key] = sizes_cpu
group_gemm(A_list, B_list, C_list, sfa_list, sfb_list, sizes_cpu)
return C_list
scrolls · 1041 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 492630.
⋯ 6 unchanged lines"""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- ⏱ 169 ± 0.1 µs- ⚡ 169 µs 🐌 170 µs+ ⏱ 89.3 ± 0.09 µs+ ⚡ 88.7 µs 🐌 89.6 µ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- ⏱ 159 ± 0.1 µs- ⚡ 159 µs 🐌 159 µs+ ⏱ 82.7 ± 0.05 µs+ ⚡ 82.7 µs 🐌 82.8 µsg: 2; k: [4096, 4096]; m: [192, 320]; n: [3072, 3072]; seed: 1111- ⏱ 45.6 ± 0.04 µs- ⚡ 45.5 µs 🐌 45.7 µs+ ⏱ 29.5 ± 0.03 µs+ ⚡ 29.1 µs 🐌 29.8 µsg: 2; k: [1536, 1536]; m: [128, 384]; n: [4096, 4096]; seed: 1111- ⏱ 20.0 ± 0.02 µs- ⚡ 19.9 µs 🐌 20.0 µs+ ⏱ 16.4 ± 0.02 µs+ ⚡ 16.0 µs 🐌 16.6 µs"""CUDA_SRC = """⋯ 19 unchanged linesTORCH_CHECK(_err == cudaSuccess, "CUDA error: ", cudaGetErrorString(_err)); \\} while (0)+ static inline uint64_t hash_combine_u64(uint64_t h, uint64_t x) {+ // 64-bit FNV-1a variant+ h ^= x;+ h *= 1099511628211ULL;+ return h;+ }+constexpr int WARP_SIZE = 32;+ constexpr int MMA_K = 64;// Cache hintsconstexpr uint64_t EVICT_NORMAL = 0x1000000000000000;constexpr uint64_t EVICT_FIRST = 0x12F0000000000000;+ constexpr uint64_t EVICT_LAST = 0x14F0000000000000;// Work item for persistent kernelstruct WorkItem {⋯ 6 unchanged linesstruct __align__(128) ProblemInfo {CUtensorMap A_tmap;CUtensorMap B_tmap;- const char* SFA_ptr; // points to underlying contiguous storage in [l, mn/128, (k/16)/4, 32, 4, 4]+ CUtensorMap B_tmap_256;+ const char* SFA_ptr;const char* SFB_ptr;half* C_ptr;int M, N, K;int64_t Cs0, Cs1, Cs2;};- // Helper functions__device__ inlineconstexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; };+ __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 void mbarrier_init(int mbar_addr, int count) {asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));}⋯ 3 unchanged lines:: "r"(mbar_addr), "r"(size) : "memory");}- __device__ inline void mbarrier_inval(int mbar_addr) {- asm volatile("mbarrier.inval.shared::cta.b64 [%0];" :: "r"(mbar_addr) : "memory");- }-__device__ void mbarrier_wait(int mbar_addr, int phase) {uint32_t ticks = 0x989680;asm volatile(⋯ 7 unchanged lines);}- // 3D TMA loadtemplate <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);}- // 1D TMA load- template <int CTA_GROUP = 1>- __device__ inline void tma_1d_gmem2smem(int dst, const void *tmap_ptr, int x,- int mbar_addr, uint64_t cache_policy) {- asm volatile(- "cp.async.bulk.tensor.1d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%5.L2::cache_hint "- "[%0], [%1, {%2}], [%3], %4;"- :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(mbar_addr), "l"(cache_policy), "n"(CTA_GROUP)- : "memory"- );- }-__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 "⋯ 3 unchanged lines);}- template <int CTA_GROUP = 1>- __device__ __forceinline__ void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {- asm volatile(- "tcgen05.cp.cta_group::%2.32x128b.warpx4 [%0], %1;\\n"- :: "r"(taddr), "l"(s_desc), "n"(CTA_GROUP)- );+ __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));}- template <int CTA_GROUP = 1>- __device__ __forceinline__ void tcgen05_commit(int mbar_addr) {+ __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), "n"(CTA_GROUP) : "memory"+ "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];\\n"+ :: "r"(mbar_addr) : "memory");}- template <int CTA_GROUP = 1>- __device__ __forceinline__ void tcgen05_mma_nvfp4(- int d_tmem, uint64_t a_desc, uint64_t b_desc, uint32_t i_desc,+ __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::%7.kind::mxf4nvf4.block_scale.block16 "- " [%0], %1, %2, %3, [%4], [%5], p;\\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),- "n"(CTA_GROUP)+ "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d));}- // TMEM load helper+ // TMEM load helpersstruct SHAPE {static constexpr char _16x256b[] = ".16x256b";};⋯ 36 unchanged lines);}- // SMEM Layout:- // [ProblemInfo] (aligned to 128)- // [Stage 0]- // [Stage 1]- // [Stage 2]- // [Mbarriers] (6 mbarriers: 3 for TMA complete, 3 for MMA commit)+ // ============================================================================+ // KERNEL CONFIGURATION+ // ============================================================================- constexpr int TMA_BLOCK_M = 128;constexpr int TMA_BLOCK_N = 128;constexpr int TMA_BLOCK_K = 256;- constexpr int NUM_STAGES_MAIN = 3; // Triple buffering- constexpr int NUM_STAGES_LOW_M = 2; // Lower SMEM footprint for small-M tiles- constexpr int TMA_NUM_WARPS = 4;- constexpr int TMEM_COLS = TMA_BLOCK_N * 2;+ constexpr int NUM_STAGES = 4;+ constexpr int TMA_NUM_WARPS = 8;+ constexpr int MMA_M = 128;+ constexpr int MBAR_BYTES = ((2 * NUM_STAGES * 8 + 63) & ~63);- constexpr int TMA_A_SMEM_BYTES = TMA_BLOCK_M * (TMA_BLOCK_K / 2); // 16KB- constexpr int TMA_B_SMEM_BYTES = TMA_BLOCK_N * (TMA_BLOCK_K / 2); // 16KB- constexpr int TMA_SFA_SMEM_BYTES = TMA_BLOCK_M * (TMA_BLOCK_K / 16); // 2KB- constexpr int TMA_SFB_SMEM_BYTES = TMA_BLOCK_N * (TMA_BLOCK_K / 16); // 2KB- constexpr int STAGE_SIZE = TMA_A_SMEM_BYTES + TMA_B_SMEM_BYTES + TMA_SFA_SMEM_BYTES + TMA_SFB_SMEM_BYTES; // ~36KB-- constexpr int MBAR_BYTES_MAIN = ((2 * NUM_STAGES_MAIN * 8 + 63) & ~63);- constexpr int MBAR_BYTES_LOW_M = ((2 * NUM_STAGES_LOW_M * 8 + 63) & ~63);- constexpr int SMEM_SIZE_MAIN = STAGE_SIZE * NUM_STAGES_MAIN + MBAR_BYTES_MAIN;- constexpr int SMEM_SIZE_LOW_M = STAGE_SIZE * NUM_STAGES_LOW_M + MBAR_BYTES_LOW_M;constexpr int LOW_M_THRESHOLD = 96;- template <bool LOW_M_EPILOGUE>+ // Warp assignments (8 warps total):+ // Warp 0-3: epilogue helpers+ // Warp 4: TMA producer+ // Warp 5: MMA consumer (single warp issues tcgen05.mma.cta_group::1)+ // Warp 6-7: additional helpers+ constexpr int TMA_WARP = 4;+ constexpr int MMA_WARP = 5;++ // ============================================================================+ // EPILOGUE+ // ============================================================================++ template <int BLOCK_M, int BLOCK_N, bool LOW_M_EPILOGUE>__device__ __forceinline__ void epilogue_store(const ProblemInfo& prob,int m_offset,⋯ 2 unchanged linesint warp_id,int lane_id) {- if (tid >= TMA_BLOCK_M) return;+ if (tid >= BLOCK_M) return;const int M = prob.M;const int N = prob.N;⋯ 1 unchanged linesconst int64_t Cs0 = prob.Cs0;const int64_t Cs1 = prob.Cs1;- const bool full_n = (n_offset + TMA_BLOCK_N <= N);- const bool full_m = (m_offset + TMA_BLOCK_M <= M);+ 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);⋯ 8 unchanged linesconst 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;for (int m = 0; m < m_iters; ++m) {- float tmp[TMA_BLOCK_N / 2];- tcgen05_ld_16x256b_x16(tmp, warp_id * 32 + m * 16, 0);- asm volatile("tcgen05.wait::ld.sync.aligned;\\n");+ 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");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);- half2* row1_ptr = reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset);+ 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 < TMA_BLOCK_N / 8; i++) {+ 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;⋯ 6 unchanged linesconst 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) : nullptr;- half2* row1_ptr = row1_in ? reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset) : nullptr;+ 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 < TMA_BLOCK_N / 8; i++) {+ 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;⋯ 6 unchanged lines}} else {#pragma unroll- for (int i = 0; i < TMA_BLOCK_N / 8; i++) {+ for (int i = 0; i < HALF_N / 8; i++) {const int idx = i * 4;- const int col = n_offset + i * 8 + lane_col;+ 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]);⋯ 19 unchanged lines} else {const bool row0_in = row0 < M;const bool row1_in = row1 < M;- if (full_n) {- #pragma unroll- for (int i = 0; i < TMA_BLOCK_N / 8; i++) {- const int idx = i * 4;- const int col = n_offset + i * 8 + lane_col;+ #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 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]);if (row0_in) {C_ptr[row0 * Cs0 + col * Cs1] = h00;- C_ptr[row0 * Cs0 + (col + 1) * Cs1] = h01;+ if (col + 1 < N) C_ptr[row0 * Cs0 + (col + 1) * Cs1] = h01;}if (row1_in) {C_ptr[row1 * Cs0 + col * Cs1] = h10;- C_ptr[row1 * Cs0 + (col + 1) * Cs1] = h11;+ if (col + 1 < N) C_ptr[row1 * Cs0 + (col + 1) * Cs1] = h11;}}- } else {- #pragma unroll- for (int i = 0; i < TMA_BLOCK_N / 8; i++) {- const int idx = i * 4;- const int col = n_offset + 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]);- 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;- }- }- }}}+ }}}- // Helper to issue TMA loads for a given stage- __device__ __forceinline__ void issue_tma_loads(- int A_smem, int B_smem, int SFA_smem, int SFB_smem, int mbar_addr,- const ProblemInfo* prob,- int m_offset, int n_offset, int k_iter,- int m_tile_idx, int n_tile_idx, int sf_bytes_per_m_tile, int sf_k_per_iter- ) {- const int off_k = k_iter * TMA_BLOCK_K;- tma_3d_gmem2smem<1>(A_smem, &prob->A_tmap, 0, m_offset, off_k / 256, mbar_addr, EVICT_NORMAL);- tma_3d_gmem2smem<1>(B_smem, &prob->B_tmap, 0, n_offset, off_k / 256, mbar_addr, EVICT_FIRST);-- // Scale factors live in the underlying contiguous storage order:- // [l=1, mn/128, (k/16)/4, 32, 4, 4], with each (32,4,4) tile = 512 bytes.- // For BLOCK_K=256 we need 4 consecutive 512B tiles per k_iter (total 2048B).- const int rest_k = prob->K / 64; // (K/16)/4- const int k_blk = off_k / 64; // (off_k/16)/4- const char* SFA_src = prob->SFA_ptr + (int64_t)(m_tile_idx * rest_k + k_blk) * 512;- const char* SFB_src = prob->SFB_ptr + (int64_t)(n_tile_idx * rest_k + k_blk) * 512;- tma_gmem2smem(SFA_smem, SFA_src, TMA_SFA_SMEM_BYTES, mbar_addr, EVICT_NORMAL);- tma_gmem2smem(SFB_smem, SFB_src, TMA_SFB_SMEM_BYTES, mbar_addr, EVICT_FIRST);- mbarrier_arrive_expect_tx(mbar_addr, STAGE_SIZE);- }+ // ============================================================================+ // MAIN KERNEL+ // ============================================================================- // Single Kernel (template for main/low-M variants, with optional persistence)- template <int NUM_STAGES, bool LOW_M_EPILOGUE, bool PERSISTENT>+ template <bool PERSISTENT, int BLOCK_M, int BLOCK_N>__global__ __launch_bounds__(TMA_NUM_WARPS * WARP_SIZE)- void grouped_gemm_tcgen_tma_v3_persistent(+ void grouped_gemm_kernel_v4(const ProblemInfo* __restrict__ global_probs,const WorkItem* __restrict__ work_items,int num_items,- int* __restrict__ work_counter // Only used when PERSISTENT=true+ int* __restrict__ work_counter) {+ 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;- // Shared Memory Setup (constant addresses)+ // Shared Memory Setupextern __shared__ __align__(1024) char smem_ptr[];const int smem_base = static_cast<int>(__cvta_generic_to_shared(smem_ptr));- // Pipeline Buffers- const int stages_base = smem_base;- int A_smem[NUM_STAGES];- #pragma unroll- for (int i = 0; i < NUM_STAGES; ++i) {- A_smem[i] = stages_base + STAGE_SIZE * i;- }- const int B_off = TMA_A_SMEM_BYTES;- const int SFA_off = B_off + TMA_B_SMEM_BYTES;- const int SFB_off = SFA_off + TMA_SFA_SMEM_BYTES;+ // Offsets within each stage+ 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;++ // Mbarriers+ const int mbar_base = smem_base + STAGE_SIZE * NUM_STAGES;- // Mbarriers (constant addresses)- const int mbar_base = stages_base + STAGE_SIZE * NUM_STAGES;- int tma_mbar[NUM_STAGES];- int mma_mbar[NUM_STAGES];- #pragma unroll- for (int i = 0; i < NUM_STAGES; ++i) {- tma_mbar[i] = mbar_base + i * 8;- mma_mbar[i] = mbar_base + (NUM_STAGES + i) * 8;- }+ // TMEM addresses+ 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);- // TMEM addresses (constant)- const int tmem_base = 0;- const int d_tmem = tmem_base;- const int sfa_tmem = tmem_base + TMA_BLOCK_N;- const int sfb_tmem = sfa_tmem + 4 * (TMA_BLOCK_K / 64);- constexpr uint32_t idesc = (1U << 7U) | (1U << 10U) | ((uint32_t)TMA_BLOCK_N >> 3U << 17U) | ((uint32_t)TMA_BLOCK_M >> 7U << 27U);-- // Allocate TMEM ONCE+ // Allocate TMEMif (warp_id == 0) {- asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(stages_base), "r"(TMEM_COLS));+ 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()) {+ 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 work index (only used in persistent mode)+ // Persistent work counter__shared__ int shared_work_idx;- // ================================================================- // WORK DISPATCH: Persistent loop vs single-tile- // ================================================================+ // Get first work itemint work_idx;if constexpr (PERSISTENT) {- // Fetch work atomicallyif (tid == 0) {shared_work_idx = atomicAdd(work_counter, 1);}__syncthreads();work_idx = shared_work_idx;} else {- // Non-persistent: each CTA handles one tile via blockIdxwork_idx = blockIdx.x;}- // Main processing loop (runs once for non-persistent, loops for persistent)+ // Main processing loopwhile (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 * TMA_BLOCK_M;- const int n_offset = work.tile_n * TMA_BLOCK_N;+ 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 - 1) / TMA_BLOCK_K;+ const int num_k_iters = K / TMA_BLOCK_K;- // Scale factor offsets- const int m_tile_idx = m_offset / TMA_BLOCK_M;- const int n_tile_idx = n_offset / TMA_BLOCK_N;- const int sf_bytes_per_m_tile = TMA_BLOCK_M * (K / 16);- const int sf_k_per_iter = TMA_SFA_SMEM_BYTES;-- // Initialize or reinit mbarriers- // Note: mbarrier_inval is not strictly needed if we ensure all threads synced and state is clean via init-+ // Initialize mbarriers per tileif (tid == 0) {- for(int i=0; i<NUM_STAGES; ++i) mbarrier_init(tma_mbar[i], 1);- for(int i=0; i<NUM_STAGES; ++i) mbarrier_init(mma_mbar[i], 1);+ for (int i = 0; i < NUM_STAGES; ++i) {+ mbarrier_init(mbar_base + i * 8, 1);+ mbarrier_init(mbar_base + (NUM_STAGES + i) * 8, 1);+ }asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");}__syncthreads();- // ----------------------------------------------------------------- // PRODUCER WARP (Warp 0): Issues TMA- // ----------------------------------------------------------------- if (warp_id == 0) {- if (lane_id == 0) {- for (int k_iter = 0; k_iter < num_k_iters; ++k_iter) {- const int stage = k_iter % NUM_STAGES;- // Wait for buffer to be free (consumed by MMA)- // Ideally, mma_mbar tracks "buffer consumed".- // Init count 1.- // Logic:- // k=0: buffer is free (init state).- // But we need to sync with consumer?- // Let's assume mma_mbar is signaled when MMA is done using the buffer.- // Initial state: buffers are free. mma_mbar should NOT block for first use.- // But mbarrier logic: wait() blocks until phase flips.- // We need to manage phases carefully.-- // Correct logic:- // TMA thread waits for buffer to be available.- // For k < NUM_STAGES, buffers are initially available.- // For k >= NUM_STAGES, wait for previous usage to complete.-- if (k_iter >= NUM_STAGES) {- // Wait for stage to be released by MMA- // Corresponding k was k_iter - NUM_STAGES- mbarrier_wait(mma_mbar[stage], (k_iter - NUM_STAGES) / NUM_STAGES); // Wait for phase flip?- // Or just use the same phase logic as coupled.- // Coupled used: mbarrier_wait(mma_mbar[next_stage], ...)- }+ // ================================================================+ // PRODUCER WARP: TMA (warp 4)+ // ================================================================+ if (warp_id == TMA_WARP && elect_sync()) {+ // With M-tile-major work ordering, A stays in L2 across N-tiles+ constexpr uint64_t cache_A = EVICT_LAST;+ constexpr uint64_t cache_B = EVICT_FIRST;- // Issue TMA- const int stage_base = A_smem[stage];- issue_tma_loads(stage_base, stage_base + B_off, stage_base + SFA_off, stage_base + SFB_off,- tma_mbar[stage], &prob, m_offset, n_offset, k_iter,- m_tile_idx, n_tile_idx, sf_bytes_per_m_tile, sf_k_per_iter);- }+ 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;++ // TMA loads+ tma_3d_gmem2smem<1>(stage_base, &prob.A_tmap, 0, m_offset, off_k / 256, mbar_addr, cache_A);+ 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);+ }++ // Scale factor loads+ 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);+ 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);+ }++ mbarrier_arrive_expect_tx(mbar_addr, STAGE_SIZE);+ };++ // Pipeline priming: issue first NUM_STAGES TMAs without waiting+ for (int k_iter = 0; k_iter < NUM_STAGES && k_iter < num_k_iters; k_iter++) {+ issue_tma(k_iter, k_iter);}- }- // ----------------------------------------------------------------- // CONSUMER WARP (Warp 1): Issues MMA- // ----------------------------------------------------------------- else if (warp_id == 1) {- if (lane_id == 0) {- for (int k_iter = 0; k_iter < num_k_iters; ++k_iter) {- const int stage = k_iter % NUM_STAGES;-- // Wait for data ready (TMA complete)- mbarrier_wait(tma_mbar[stage], k_iter / NUM_STAGES);-- const int stage_base = A_smem[stage];-- // Issue MMA- constexpr uint64_t SF_desc = (desc_encode(8 * 16) << 32ULL) | (1ULL << 46ULL);- 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);+ // Steady state: wait for MMA, then issue TMA+ for (int k_iter = NUM_STAGES; k_iter < num_k_iters; k_iter++) {+ const int stage = k_iter % NUM_STAGES;+ const int mma_phase = (k_iter / NUM_STAGES - 1) % 2;+ mbarrier_wait(mbar_base + (NUM_STAGES + stage) * 8, mma_phase);+ issue_tma(k_iter, stage);+ }+ }- constexpr uint64_t AB_desc = (desc_encode(8 * 128) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);- const uint64_t A_desc = AB_desc | ((uint64_t)stage_base >> 4ULL);- const uint64_t B_desc = AB_desc | ((uint64_t)(stage_base + B_off) >> 4ULL);+ // ================================================================+ // CONSUMER WARP: MMA (warp 5)+ // Single elected thread issues tcgen05.mma.cta_group::1+ // ================================================================+ else if (warp_id == MMA_WARP && elect_sync()) {+ auto make_desc_AB = [](int addr) -> uint64_t {+ const int SBO = 8 * 128;+ return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);+ };+ auto make_desc_SF = [](int addr) -> uint64_t {+ const int SBO = 8 * 16;+ return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);+ };- #pragma unroll- for (int k = 0; k < (TMA_BLOCK_K / 64); k++) {- tcgen05_cp_nvfp4<1>(sfa_tmem + k * 4, SFA_desc + (uint64_t)k * (512ULL >> 4ULL));- tcgen05_cp_nvfp4<1>(sfb_tmem + k * 4, SFB_desc + (uint64_t)k * (512ULL >> 4ULL));- }+ for (int k_iter = 0; k_iter < num_k_iters; k_iter++) {+ const int stage = k_iter % NUM_STAGES;+ const int tma_phase = (k_iter / NUM_STAGES) % 2;+ mbarrier_wait(mbar_base + stage * 8, tma_phase);- #pragma unroll- for (int k2 = 0; k2 < (TMA_BLOCK_K / 64); k2++) {- const uint64_t a_desc = A_desc + (uint64_t)k2 * (32ULL >> 4ULL);- const uint64_t b_desc = B_desc + (uint64_t)k2 * (32ULL >> 4ULL);- const int enable_input_d = (k_iter == 0 && k2 == 0) ? 0 : 1;- tcgen05_mma_nvfp4<1>(d_tmem, a_desc, b_desc, idesc, sfa_tmem + k2 * 4, sfb_tmem + k2 * 4, enable_input_d);- }-- // Commit MMA -> Signals mma_mbar[stage]- // When commit reaches mma_mbar, it allows TMA producer to reuse buffer for next phase- tcgen05_commit<1>(mma_mbar[stage]);+ const int stage_base = smem_base + stage * STAGE_SIZE;++ // Copy scale factors to TMEM+ 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);++ 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));}-- // Final Wait- // Wait for last commit to ensure all instructions retired?- // Actually, we must ensure all work is done before exiting.- // Wait for the last commit to finish- if (num_k_iters > 0) {- const int last_stage = (num_k_iters - 1) % NUM_STAGES;- mbarrier_wait(mma_mbar[last_stage], (num_k_iters - 1) / NUM_STAGES);- }+ }++ // Issue MMA+ for (int k1 = 0; k1 < TMA_BLOCK_K / 256; k1++) {+ for (int k2 = 0; k2 < 256 / MMA_K; k2++) {+ uint64_t a_desc = make_desc_AB(stage_base + k1 * BLOCK_M * 128 + k2 * 32);+ uint64_t b_desc = make_desc_AB(stage_base + B_off + k1 * BLOCK_N * 128 + k2 * 32);++ int k_sf = k1 * 4 + k2;+ const int scale_A_tmem = SFA_tmem + k_sf * 4 + (work.tile_m % (MMA_M / BLOCK_M)) * (BLOCK_M / 32);+ int scale_B_tmem;+ if constexpr (BLOCK_N == 256) {+ scale_B_tmem = SFB_tmem + k_sf * 8;+ } else {+ scale_B_tmem = SFB_tmem + k_sf * 4;+ }++ const int enable_input_d = (k_iter == 0 && k1 == 0 && k2 == 0) ? 0 : 1;+ tcgen05_mma_nvfp4(a_desc, b_desc, idesc, scale_A_tmem, scale_B_tmem, enable_input_d);+ }+ }++ tcgen05_commit(mbar_base + (NUM_STAGES + stage) * 8);}- }-- // ----------------------------------------------------------------- // SYNCHRONIZATION- // ----------------------------------------------------------------- __syncthreads(); // Wait for all warps to finish-- // Ensure MMA results visible- asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");- // Epilogue- epilogue_store<LOW_M_EPILOGUE>(prob, m_offset, n_offset, tid, warp_id, lane_id);+ // Wait for final commit+ if (num_k_iters > 0) {+ const int last_stage = (num_k_iters - 1) % NUM_STAGES;+ const int last_phase = ((num_k_iters - 1) / NUM_STAGES) % 2;+ mbarrier_wait(mbar_base + (NUM_STAGES + last_stage) * 8, last_phase);+ }+ }- __syncthreads();-- // Loop control: continue (persistent) or break (non-persistent)+ // ================================================================+ // SYNCHRONIZATION & EPILOGUE+ // ================================================================+ __syncthreads();+ 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();++ // Get next work itemif constexpr (PERSISTENT) {- // Fetch next workif (tid == 0) {shared_work_idx = atomicAdd(work_counter, 1);}__syncthreads();work_idx = shared_work_idx;} else {- // Non-persistent: exit after single tilebreak;}- } // End main loop-- // Deallocate TMEM once at end+ }++ // Deallocate TMEMif (warp_id == 0) {- tcgen05_dealloc_cols_cta1(tmem_base, TMEM_COLS);+ tcgen05_dealloc_cols_cta1(0, TMEM_COLS);}}+ // ============================================================================+ // TENSOR MAP INITIALIZATION+ // ============================================================================- // Tensor Map Initializationvoid 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, "init_AB_tmap_u4: ptr is null");+ 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");-+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, shared_width / 256};uint32_t elementStrides[rank] = {1, 1, 1};- // cuTensorMapEncodeTiled is relatively expensive on the host; amortize by caching- // a per-shape template and then patching only the base address each call.- struct Key {- uint64_t gh, gw;- uint32_t sh, sw;- };- struct KeyHash {- size_t operator()(const Key& k) const noexcept {+ // cuTensorMapEncodeTiled is a relatively expensive driver call.+ // Cache a per-shape template and then patch only the base address.+ 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 KeyEq {- bool operator()(const Key& a, const Key& b) const noexcept {+ 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;}};- static std::unordered_map<Key, CUtensorMap, KeyHash, KeyEq> tmpl_cache;+ 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;+ }+ };- Key key{global_height, global_width, shared_height, shared_width};- auto it = tmpl_cache.find(key);- if (it == tmpl_cache.end()) {+ 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,- CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,+ &tmp, CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,rank, (void*)ptr, globalDim, globalStrides, boxDim, elementStrides,- CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,- CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_128B,- CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,- CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE+ CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,+ CU_TENSOR_MAP_L2_PROMOTION_NONE, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);- TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapEncodeTiled failed for AB template");- it = tmpl_cache.emplace(key, tmp).first;+ TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapEncodeTiled failed");+ sit = tmpl_cache.emplace(skey, tmp).first;}- // Copy template then patch base address.- *tmap = it->second;- auto err = cuTensorMapReplaceAddress(tmap, (void*)ptr);- TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapReplaceAddress failed for AB");+ CUtensorMap tmp = sit->second;+ auto err = cuTensorMapReplaceAddress(&tmp, (void*)ptr);+ TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapReplaceAddress failed");+ ptr_cache.emplace(pkey, tmp);+ *tmap = tmp;}- // Device-side padding/alignment for AB tensors.- //- // TMA with CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE requires all accessed elements to be- // in-bounds, so we pad the leading dimension to 128 (tile height) and zero-fill- // the tail. We cache by source data_ptr() + padded_M to make repeated calls cheap.- static at::Tensor pad_u4_tensor_m128(- const at::Tensor& src,- int64_t padded_m- ) {- TORCH_CHECK(src.is_cuda(), "pad_u4_tensor_m128: src must be CUDA");- TORCH_CHECK(src.numel() > 0, "pad_u4_tensor_m128: empty tensor");- TORCH_CHECK(padded_m >= src.size(0), "padded_m must be >= src.size(0)");+ struct PadCacheEntry {+ at::Tensor buf;+ size_t zeroed_from = 0; // byte offset; bytes in [zeroed_from, end) are guaranteed zero+ };- const bool needs_pad = (src.size(0) != padded_m);- const bool needs_align = (((uintptr_t)src.data_ptr() & 0xF) != 0);- if (!needs_pad && !needs_align) return src;+ static at::Tensor pad_u4_tensor_cached(const at::Tensor& src, int64_t padded_m, PadCacheEntry* entry) {+ TORCH_CHECK(entry != nullptr, "pad_u4_tensor_cached: entry is null");+ if (src.size(0) == padded_m && ((uintptr_t)src.data_ptr() & 0xF) == 0) return src;- auto new_sizes = src.sizes().vec();- new_sizes[0] = padded_m;- at::Tensor dst = at::empty(new_sizes, src.options());+ auto new_sizes = src.sizes().vec();+ new_sizes[0] = padded_m;- // src is created by the benchmark generator and is contiguous; use a single- // D2D memcpy and then zero the padded tail.- const size_t copy_bytes = (size_t)src.nbytes();- const size_t total_bytes = (size_t)dst.nbytes();- TORCH_CHECK(copy_bytes <= total_bytes, "pad_u4_tensor_m128: size mismatch");-- CUDA_CHECK(cudaMemcpyAsync(dst.data_ptr(), src.data_ptr(), copy_bytes, cudaMemcpyDeviceToDevice));- if (total_bytes > copy_bytes) {- CUDA_CHECK(cudaMemsetAsync((char*)dst.data_ptr() + copy_bytes, 0, total_bytes - copy_bytes));+ bool reuse_ok = entry->buf.defined() && entry->buf.dim() == (int)new_sizes.size();+ if (reuse_ok) {+ // Allow reusing a larger leading-dimension buffer (avoid reallocs when M/N shrink),+ // but require trailing dimensions to match exactly.+ if (entry->buf.size(0) < padded_m) reuse_ok = false;+ for (int d = 1; d < entry->buf.dim(); d++) {+ if (entry->buf.size(d) != new_sizes[(size_t)d]) { reuse_ok = false; break; }}+ if (reuse_ok && (((uintptr_t)entry->buf.data_ptr() & 0xF) != 0)) reuse_ok = false;+ }- return dst;+ const bool need_new =+ !entry->buf.defined() ||+ entry->buf.device() != src.device() ||+ entry->buf.scalar_type() != src.scalar_type() ||+ !reuse_ok;++ if (need_new) {+ entry->buf = at::empty(new_sizes, src.options());+ entry->zeroed_from = (size_t)entry->buf.nbytes(); // nothing guaranteed yet+ }++ const size_t copy_bytes = (size_t)src.nbytes();+ const size_t total_bytes = (size_t)entry->buf.nbytes();+ TORCH_CHECK(copy_bytes <= total_bytes, "pad_u4_tensor_cached: size mismatch");++ CUDA_CHECK(cudaMemcpyAsync(entry->buf.data_ptr(), src.data_ptr(), copy_bytes, cudaMemcpyDeviceToDevice));+ if (copy_bytes < entry->zeroed_from) {+ CUDA_CHECK(cudaMemsetAsync((char*)entry->buf.data_ptr() + copy_bytes, 0, entry->zeroed_from - copy_bytes));+ entry->zeroed_from = copy_bytes;+ } else {+ entry->zeroed_from = copy_bytes;+ }+ return entry->buf;}- // Host Entry Point+ // ============================================================================+ // HOST ENTRY POINT+ // ============================================================================+std::vector<at::Tensor> group_gemm(std::vector<at::Tensor> A_list,std::vector<at::Tensor> B_list,⋯ 3 unchanged linesat::Tensor sizes_cpu) {int64_t G = A_list.size();- TORCH_CHECK(B_list.size() == G && C_list.size() == G, "A/B/C list sizes must match");- TORCH_CHECK(sfa_list.size() == G && sfb_list.size() == G, "sfa/sfb list sizes must match");- TORCH_CHECK(sizes_cpu.device().is_cpu() && sizes_cpu.scalar_type() == at::kLong, "sizes must be CPU int64");-- TORCH_CHECK(A_list[0].is_cuda(), "A must be CUDA");auto dev = A_list[0].device();c10::cuda::CUDAGuard device_guard(dev);-auto sizes_accessor = sizes_cpu.accessor<int64_t, 2>();-- // Set shared memory attributes for all kernel variants (once)+static bool attrs_set = false;if (!attrs_set) {- auto set_attr = [&](auto kernel, int smem_size) {- CUDA_CHECK(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));- };- set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_MAIN, false, false>, SMEM_SIZE_MAIN);- set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_MAIN, false, true>, SMEM_SIZE_MAIN);- set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_LOW_M, true, false>, SMEM_SIZE_LOW_M);- set_attr(grouped_gemm_tcgen_tma_v3_persistent<NUM_STAGES_LOW_M, true, true>, SMEM_SIZE_LOW_M);+ constexpr int STAGE_SIZE_128 = 128 * 128 + 128 * 128 + 128 * 16 + 128 * 16;+ constexpr int SMEM_SIZE_128 = STAGE_SIZE_128 * NUM_STAGES + MBAR_BYTES;+ constexpr int STAGE_SIZE_64 = 64 * 128 + 128 * 128 + 128 * 16 + 128 * 16;+ constexpr int SMEM_SIZE_64 = STAGE_SIZE_64 * NUM_STAGES + MBAR_BYTES;+ constexpr int STAGE_SIZE_256 = 128 * 128 + 256 * 128 + 128 * 16 + 256 * 16;+ constexpr int SMEM_SIZE_256 = STAGE_SIZE_256 * NUM_STAGES + MBAR_BYTES;+ CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<true, 128, 128>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_128));+ CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<false, 128, 128>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_128));+ CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<true, 64, 128>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_64));+ CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<false, 64, 128>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_64));+ CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<true, 128, 256>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_256));+ CUDA_CHECK(cudaFuncSetAttribute(grouped_gemm_kernel_v4<false, 128, 256>, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_SIZE_256));attrs_set = true;}std::vector<ProblemInfo> problem_infos(G);- std::vector<WorkItem> work_items_main;- std::vector<WorkItem> work_items_low;- work_items_main.reserve(G * 64);- work_items_low.reserve(G * 64);-- // Track min K for persistence heuristic- int min_K = INT_MAX;+ static thread_local std::vector<WorkItem> cached_work_items_128;+ static thread_local std::vector<WorkItem> cached_work_items_64;+ static thread_local std::vector<WorkItem> cached_work_items_256;+ static thread_local uint64_t cached_work_hash = 0;+ static thread_local bool cached_work_valid = false;- std::vector<at::Tensor> keepers;- keepers.reserve(G * 2);+ std::vector<uint8_t> active(G, 0);+ std::vector<uint8_t> use_64(G, 0);+ std::vector<uint8_t> use_256(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);- for (int64_t prob_idx = 0; prob_idx < G; prob_idx++) {- at::Tensor A = A_list[prob_idx];- at::Tensor B = B_list[prob_idx];- at::Tensor C = C_list[prob_idx];- at::Tensor sfa = sfa_list[prob_idx];- at::Tensor sfb = sfb_list[prob_idx];-- int64_t M = sizes_accessor[prob_idx][0];- int64_t N = sizes_accessor[prob_idx][1];- int64_t K = sizes_accessor[prob_idx][2];+ uint64_t work_hash = 1469598103934665603ULL;+ 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) {+ work_hash = hash_combine_u64(work_hash, 0);+ continue;+ }+ active[(size_t)i] = 1;+ bool is_64 = (M <= 64) && (N <= 2048);+ bool is_256 = (!is_64) && (N >= 4096) && (K >= 2048) && ((N & 255) == 0);+ use_64[(size_t)i] = is_64;+ use_256[(size_t)i] = is_256;+ int block_m = is_64 ? 64 : 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);- if (A.stride(1) != 1 || B.stride(1) != 1) continue;- TORCH_CHECK((K % TMA_BLOCK_K) == 0, "K must be multiple of ", TMA_BLOCK_K);+ 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_64);+ work_hash = hash_combine_u64(work_hash, (uint64_t)is_256);+ }- // TMA requires K to be compatible and AB pointers to be aligned.- // For partial tiles along M/N, we pad the leading dimension to 128 rows- // and zero-fill, but keep prob_info.M/N as the *true* sizes for epilogue.- const int64_t padded_M = ((M + TMA_BLOCK_M - 1) / TMA_BLOCK_M) * TMA_BLOCK_M;- const int64_t padded_N = ((N + TMA_BLOCK_N - 1) / TMA_BLOCK_N) * TMA_BLOCK_N;-- A = pad_u4_tensor_m128(A, padded_M);- B = pad_u4_tensor_m128(B, padded_N);- keepers.push_back(A);- keepers.push_back(B);+ if (!cached_work_valid || cached_work_hash != work_hash) {+ cached_work_items_128.clear();+ cached_work_items_64.clear();+ cached_work_items_256.clear();+ cached_work_items_128.reserve(G * 32);+ cached_work_items_64.reserve(G * 32);+ cached_work_items_256.reserve(G * 32);+ for (int64_t i = 0; i < G; i++) {+ if (!active[(size_t)i]) continue;+ // M-tile-major ordering: consecutive work items share the same M-tile+ // so A data stays hot in L2 while iterating N-tiles+ 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_64[(size_t)i]) {+ cached_work_items_64.push_back({(int)i, tm, tn});+ } else if (use_256[(size_t)i]) {+ cached_work_items_256.push_back({(int)i, tm, tn});+ } else {+ cached_work_items_128.push_back({(int)i, tm, tn});+ }+ }+ }+ }+ cached_work_hash = work_hash;+ cached_work_valid = true;+ }- ProblemInfo& prob_info = problem_infos[prob_idx];- prob_info.M = (int)M;- prob_info.N = (int)N;- prob_info.K = (int)K;- min_K = std::min(min_K, (int)K);- prob_info.Cs0 = C.stride(0);- prob_info.Cs1 = C.stride(1);- prob_info.Cs2 = C.stride(2);- prob_info.C_ptr = (half*)C.data_ptr();+ static thread_local std::vector<PadCacheEntry> A_pad_cache;+ static thread_local std::vector<PadCacheEntry> B_pad_cache;+ if ((int64_t)A_pad_cache.size() < G) A_pad_cache.resize((size_t)G);+ if ((int64_t)B_pad_cache.size() < G) B_pad_cache.resize((size_t)G);- init_AB_tmap_u4(&prob_info.A_tmap, A.data_ptr(), (uint64_t)A.size(0), (uint64_t)K, TMA_BLOCK_M, TMA_BLOCK_K);- init_AB_tmap_u4(&prob_info.B_tmap, B.data_ptr(), (uint64_t)B.size(0), (uint64_t)K, TMA_BLOCK_N, TMA_BLOCK_K);- // The provided SF tensors are a (non-contiguous) view into a contiguous backing- // storage laid out as [l=1, mn/128, (k/16)/4, 32, 4, 4]. We access the backing- // storage directly via data_ptr().- prob_info.SFA_ptr = (const char*)sfa.data_ptr();- prob_info.SFB_ptr = (const char*)sfb.data_ptr();+ uint64_t probs_hash = 1469598103934665603ULL;+ 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];- int num_tiles_m = ceil_div((int)M, TMA_BLOCK_M);- int num_tiles_n = ceil_div((int)N, TMA_BLOCK_N);+ int block_m = use_64[(size_t)i] ? 64 : 128;+ int block_n = use_256[(size_t)i] ? 256 : 128;+ int64_t padded_M = ((M + block_m - 1) / block_m) * block_m;+ int64_t padded_N = ((N + block_n - 1) / block_n) * block_n;- std::vector<WorkItem>& target = (M <= LOW_M_THRESHOLD) ? work_items_low : work_items_main;- for (int tm = 0; tm < num_tiles_m; tm++) {- for (int tn = 0; tn < num_tiles_n; tn++) {- target.push_back({(int)prob_idx, tm, tn});- }+ at::Tensor A = pad_u4_tensor_cached(A_list[i], padded_M, &A_pad_cache[(size_t)i]);+ at::Tensor B = pad_u4_tensor_cached(B_list[i], padded_N, &B_pad_cache[(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.Cs2 = C_list[i].stride(2);+ 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.data_ptr(), A.size(0), K, block_m, TMA_BLOCK_K);+ init_AB_tmap_u4(&p.B_tmap, B.data_ptr(), B.size(0), K, 128, TMA_BLOCK_K);+ if (use_256[(size_t)i]) {+ init_AB_tmap_u4(&p.B_tmap_256, B.data_ptr(), B.size(0), K, 256, TMA_BLOCK_K);+ } 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.data_ptr());+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)B.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);+ probs_hash = hash_combine_u64(probs_hash, (uint64_t)p.Cs2);}- if (work_items_main.empty() && work_items_low.empty()) return C_list;+ if (cached_work_items_128.empty() && cached_work_items_64.empty() && cached_work_items_256.empty()) return C_list;- // Device allocationsauto options = at::TensorOptions().dtype(at::kByte).device(dev);- at::Tensor d_probs = at::empty({(int64_t)(G * sizeof(ProblemInfo))}, options);-- CUDA_CHECK(cudaMemcpyAsync(d_probs.data_ptr(), problem_infos.data(), G * sizeof(ProblemInfo), cudaMemcpyHostToDevice));+ static thread_local at::Tensor d_probs_cache;+ static thread_local at::Tensor d_work_cache_128;⋯ diff truncated
scrolls · 1201 diff lines total
Best evidence level for this revision: reported
JSON