KernelBench hard · RTX PRO 6000
FP8 GEMM Kimi K3 (256k)
manually audited: clean
Genuine hand-written CUDA C++ tcgen05 (5th-gen tensor core) fp8 GEMM via torch load_inline: persistent warp-specialized pipeline (TMA loader warp + MMA umpire warp + 4-warp epilogue), inline-PTX tcgen05.mma kind::f8f6f4 with fp32 accumulation in tensor memory, per-channel dequant scale applied in an SMEM-staged epilogue, TMA bulk-tensor stores of bf16 output. The session started from a Gluon (triton.experimental.gluon) prototype and converged to raw CUDA+PTX. No torch._scaled_mm / aten._scaled_mm anywhere in solution.py. The off-alignment K=4127 shape is handled by a real zero-padding kernel to K0=4128, re-run from live input data on every forward. Host-side caching is limited to compiled-artifact state (CUtensorMap descriptors keyed on base ptr + dims, padded WEIGHT keyed on weight._version) — never outputs. Per-shape fractions 0.4489/0.3441/ 0.0259/0.6091 (geomean 0.2222); shape 2 (M=32 decode) is the known bandwidth/launch-bound shape where every model lands ~0.03.
Kernel source (redacted)
"""FP8 e4m3 GEMM for B200 (SM100): y = (x @ w.T) * weight_scale, bf16 out.
Hand-written CUDA C++ tcgen05 (5th-gen tensor core) kernel: persistent
warp-specialized pipeline (TMA loader warp + MMA umpire warp + epilogue
warpgroup), fp8 x fp8 MMA with fp32 accumulation in tensor memory, per-channel
dequant scale in an SMEM-staged epilogue. No vendor GEMM calls anywhere.
"""
import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline
OP_TYPE = "gemm"
SUPPORTED_PRECISIONS = ["fp8_e4m3"]
HARDWARE_REQUIRED = ["B200"]
E4M3_MAX = 448.0
_CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <cstdint>
#include <unordered_map>
#define DEV __device__ __forceinline__
DEV static uint32_t elect_one_sync() {
uint32_t pred = 0;
asm volatile(
"{\n"
".reg .pred P1;\n"
"elect.sync _|P1, %1;\n"
"selp.b32 %0, 1, 0, P1;\n"
"}\n"
: "=r"(pred)
: "r"(0xffffffff));
return pred;
}
DEV static void mbar_init(uint64_t* bar, uint32_t count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;\n"
:
: "r"((uint32_t)__cvta_generic_to_shared(bar)), "r"(count));
}
DEV static void mbar_expect(uint64_t* bar, uint32_t bytes) {
asm volatile("mbarrier.arrive.expect_tx.shared::cta.b64 _, [%0], %1;\n"
:
: "r"((uint32_t)__cvta_generic_to_shared(bar)), "r"(bytes));
}
DEV static void mbar_arrive(uint64_t* bar) {
asm volatile("mbarrier.arrive.shared::cta.b64 _, [%0];\n"
:
: "r"((uint32_t)__cvta_generic_to_shared(bar)));
}
DEV static void mbar_wait(uint64_t* bar, uint32_t phase) {
asm volatile(
"{\n"
".reg .pred P1;\n"
"WAITLOOP:\n"
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1;\n"
"@P1 bra DONE;\n"
"bra WAITLOOP;\n"
"DONE:\n"
"}\n"
:
: "r"((uint32_t)__cvta_generic_to_shared(bar)), "r"(phase));
}
DEV static void tma_load_2d(const void* smem_dst, const CUtensorMap* map,
int32_t c0, int32_t c1, uint64_t* bar) {
asm volatile(
"cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4}], [%2];\n"
:
: "r"((uint32_t)__cvta_generic_to_shared(smem_dst)),
"l"((uint64_t)map),
"r"((uint32_t)__cvta_generic_to_shared(bar)),
"r"(c0), "r"(c1)
: "memory");
}
DEV static void tma_store_2d(const CUtensorMap* map, const void* smem_src,
int32_t c0, int32_t c1) {
asm volatile(
"cp.async.bulk.tensor.2d.global.shared::cta.bulk_group [%0, {%2, %3}], [%1];\n"
:
: "l"((uint64_t)map),
"r"((uint32_t)__cvta_generic_to_shared(smem_src)),
"r"(c0), "r"(c1)
: "memory");
}
DEV static void tma_store_commit() { asm volatile("cp.async.bulk.commit_group;"); }
template <int N>
DEV static void tma_store_wait() {
asm volatile("cp.async.bulk.wait_group.read %0;" ::"n"(N));
}
DEV static void fence_async_shared() {
asm volatile("fence.proxy.async.shared::cta;");
}
// smem matrix descriptor (UMMA). layout_type: 2=SW128, 4=SW64, 6=SW32.
DEV static uint64_t smem_desc(uint32_t saddr, uint32_t sbo, uint32_t swz_type) {
uint64_t d = ((uint64_t)(saddr >> 4) & 0x3FFFULL);
d |= ((uint64_t)(sbo >> 4)) << 32;
d |= 1ULL << 46; // version for blackwell
d |= (uint64_t)swz_type << 61;
return d;
}
DEV static void umma_f8(uint64_t da, uint64_t db, uint32_t tmem_d, uint32_t idesc,
uint32_t accum) {
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"setp.ne.b32 p, %4, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::f8f6f4 [%0], %1, %2, %3, {%5, %6, %7, %8}, p; \n\t"
"}\n"
:
: "r"(tmem_d), "l"(da), "l"(db), "r"(idesc), "r"(accum),
"r"(0), "r"(0), "r"(0), "r"(0));
}
DEV static void umma_commit(uint64_t* bar) {
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];\n"
:
: "r"((uint32_t)__cvta_generic_to_shared(bar)));
}
DEV static void tmem_alloc(uint32_t* smem_out, uint32_t ncols) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;\n"
:
: "r"((uint32_t)__cvta_generic_to_shared(smem_out)), "r"(ncols));
}
DEV static void tmem_dealloc(uint32_t taddr, uint32_t ncols) {
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n" ::"r"(taddr),
"r"(ncols));
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;");
}
DEV static void tmem_ld64(float* vals, uint32_t taddr) {
uint32_t* v = (uint32_t*)vals;
asm volatile("tcgen05.ld.sync.aligned.32x32b.x64.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];\n"
: "=r"(v[0]), "=r"(v[1]), "=r"(v[2]), "=r"(v[3]), "=r"(v[4]), "=r"(v[5]),
"=r"(v[6]), "=r"(v[7]), "=r"(v[8]), "=r"(v[9]), "=r"(v[10]), "=r"(v[11]),
"=r"(v[12]), "=r"(v[13]), "=r"(v[14]), "=r"(v[15]), "=r"(v[16]), "=r"(v[17]),
"=r"(v[18]), "=r"(v[19]), "=r"(v[20]), "=r"(v[21]), "=r"(v[22]), "=r"(v[23]),
"=r"(v[24]), "=r"(v[25]), "=r"(v[26]), "=r"(v[27]), "=r"(v[28]), "=r"(v[29]),
"=r"(v[30]), "=r"(v[31]), "=r"(v[32]), "=r"(v[33]), "=r"(v[34]), "=r"(v[35]),
"=r"(v[36]), "=r"(v[37]), "=r"(v[38]), "=r"(v[39]), "=r"(v[40]), "=r"(v[41]),
"=r"(v[42]), "=r"(v[43]), "=r"(v[44]), "=r"(v[45]), "=r"(v[46]), "=r"(v[47]),
"=r"(v[48]), "=r"(v[49]), "=r"(v[50]), "=r"(v[51]), "=r"(v[52]), "=r"(v[53]),
"=r"(v[54]), "=r"(v[55]), "=r"(v[56]), "=r"(v[57]), "=r"(v[58]), "=r"(v[59]),
"=r"(v[60]), "=r"(v[61]), "=r"(v[62]), "=r"(v[63])
: "r"(taddr));
asm volatile("tcgen05.wait::ld.sync.aligned;");
}
DEV static constexpr uint32_t make_idesc(int M, int N) {
return (1u << 4) | ((uint32_t)(N >> 3) << 17) | ((uint32_t)(M >> 4) << 24);
}
// ------------------------------------------------------------------ scheduling
struct TileScheduler {
int start_pid, num_pid_m, num_pid_n, num_pid_in_group, num_pid, num_progs;
int group_m;
__device__ void init(int M, int N, int BM, int BN, int gm) {
start_pid = blockIdx.x;
num_progs = gridDim.x;
num_pid_m = (M + BM - 1) / BM;
num_pid_n = (N + BN - 1) / BN;
num_pid_in_group = gm * num_pid_n;
num_pid = num_pid_m * num_pid_n;
group_m = gm;
}
__device__ int num_tiles() const { return (num_pid - start_pid + num_progs - 1) / num_progs; }
__device__ void get_tile(int idx, int& pm, int& pn) const {
int tile_id = start_pid + idx * num_progs;
int group_id = tile_id / num_pid_in_group;
int first_pid_m = group_id * group_m;
int group_size = min(num_pid_m - first_pid_m, group_m);
pm = first_pid_m + (tile_id % group_size);
pn = (tile_id % num_pid_in_group) / group_size;
}
};
// ------------------------------------------------------------------ kernel
template <int BM, int BN, int BK, int STAGES, int GM>
__global__ void __launch_bounds__(192, 1)
k_gemm(const __grid_constant__ CUtensorMap a_map,
const __grid_constant__ CUtensorMap b_map,
const __grid_constant__ CUtensorMap c_map,
const float* __restrict__ scale, int M, int N, int K) {
constexpr int A_BYTES = BM * BK;
constexpr int B_BYTES = BN * BK;
constexpr int C_COLS = 64; // bf16 columns per c box
constexpr int C_ELEMS = BM * C_COLS; // bf16 elements per c buffer
constexpr int N_CHUNKS = BN / C_COLS;
constexpr int EPI_WARPS = 4;
constexpr int EPI_THREADS = EPI_WARPS * 32;
// m64 accumulator packs rows into 16-lane groups; the two acc buffers are
// interleaved in lane halves: buf b at lane offset 16*b (same columns).
// m128 accumulator: buf b at column offset b * BN.
constexpr int SWZ_TYPE = (BK == 128) ? 2 : 4; // 128B or 64B swizzle
extern __shared__ __align__(1024) uint8_t smem[];
uint8_t* a_smem = smem; // STAGES * A_BYTES
uint8_t* b_smem = a_smem + STAGES * A_BYTES; // STAGES * B_BYTES
__nv_bfloat16* c_smem = (__nv_bfloat16*)(b_smem + STAGES * B_BYTES); // 2 * C_ELEMS
float* s_smem = (float*)(c_smem + 2 * C_ELEMS); // BN floats
__shared__ __align__(8) uint64_t full_bars[STAGES], empty_bars[STAGES];
__shared__ __align__(8) uint64_t acc_full[2], acc_empty[2];
__shared__ uint32_t tmem_addr_sh;
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
if (warp == EPI_WARPS) {
tmem_alloc(&tmem_addr_sh, 512);
} else if (warp == EPI_WARPS + 1 && lane == 0) {
for (int i = 0; i < STAGES; ++i) {
mbar_init(&full_bars[i], 1);
mbar_init(&empty_bars[i], 1);
}
mbar_init(&acc_full[0], 1);
mbar_init(&acc_full[1], 1);
mbar_init(&acc_empty[0], EPI_WARPS);
mbar_init(&acc_empty[1], EPI_WARPS);
}
asm volatile("barrier.sync 1, 192;");
const uint32_t taddr_base = tmem_addr_sh;
const uint32_t idesc = make_idesc(BM, BN);
constexpr uint32_t ACC1_OFF = (BM == 64) ? (16u << 16) : (uint32_t)BN;
TileScheduler sched;
sched.init(M, N, BM, BN, GM);
const int k_iters = (K + BK - 1) / BK;
if (warp == EPI_WARPS) {
// ---------------------------------------------------------- loader
if (elect_one_sync()) {
uint32_t e_phase = 1; // first STAGES waits bypass
int stage = 0;
for (int t = 0; t < sched.num_tiles(); ++t) {
int pm, pn;
sched.get_tile(t, pm, pn);
const int off_m = pm * BM, off_n = pn * BN;
for (int k = 0; k < k_iters; ++k) {
const int s = stage;
mbar_wait(&empty_bars[s], e_phase);
mbar_expect(&full_bars[s], A_BYTES + B_BYTES);
tma_load_2d(a_smem + s * A_BYTES, &a_map, k * BK, off_m, &full_bars[s]);
tma_load_2d(b_smem + s * B_BYTES, &b_map, k * BK, off_n, &full_bars[s]);
if (++stage == STAGES) {
stage = 0;
e_phase ^= 1;
}
}
}
}
} else if (warp == EPI_WARPS + 1) {
// ---------------------------------------------------------- mma umpire
if (elect_one_sync()) {
uint32_t f_phase = 0, ae_phase = 1;
int stage = 0, abuf = 0;
for (int t = 0; t < sched.num_tiles(); ++t) {
mbar_wait(&acc_empty[abuf], ae_phase);
const uint32_t taddr = taddr_base + abuf * ACC1_OFF;
for (int k = 0; k < k_iters; ++k) {
const int s = stage;
mbar_wait(&full_bars[s], f_phase);
uint64_t da = smem_desc((uint32_t)__cvta_generic_to_shared(a_smem + s * A_BYTES),
8 * BK, SWZ_TYPE);
uint64_t db = smem_desc((uint32_t)__cvta_generic_to_shared(b_smem + s * B_BYTES),
8 * BK, SWZ_TYPE);
for (int kk = 0; kk < BK / 32; ++kk) {
umma_f8(da + kk * 2, db + kk * 2, taddr, idesc, (k > 0 || kk > 0) ? 1u : 0u);
}
umma_commit(&empty_bars[s]);
if (++stage == STAGES) {
stage = 0;
f_phase ^= 1;
}
}
umma_commit(&acc_full[abuf]);
if (++abuf == 2) {
abuf = 0;
ae_phase ^= 1;
}
}
}
} else if (warp < EPI_WARPS) {
// ---------------------------------------------------------- epilogue
// m128: each warp owns 32 rows. m64: lanes hold 16 rows per warp in the
// 16-lane half selected by the accumulator buffer index.
int row = 0;
bool row_ok = true;
if (BM == 128) {
row = warp * 32 + lane;
} else {
row_ok = (lane / 16) == 0; // placeholder; abuf-dependent, set in loop
}
uint32_t af_phase = 0;
int abuf = 0;
for (int t = 0; t < sched.num_tiles(); ++t) {
int pm, pn;
sched.get_tile(t, pm, pn);
const int off_m = pm * BM, off_n = pn * BN;
// preload scales for this tile (BN floats) into smem
for (int j = threadIdx.x; j < BN; j += EPI_THREADS) {
const int n = off_n + j;
s_smem[j] = (n < N) ? scale[n] : 0.f;
}
mbar_wait(&acc_full[abuf], af_phase);
const uint32_t taddr = taddr_base + abuf * ACC1_OFF;
for (int c = 0; c < N_CHUNKS; ++c) {
float vals[C_COLS];
tmem_ld64(vals, taddr + c * C_COLS);
if (c == N_CHUNKS - 1 && lane == 0) {
// all tmem reads of this accumulator are done
mbar_arrive(&acc_empty[abuf]);
}
__nv_bfloat16 outv[C_COLS];
const float* sv = s_smem + c * C_COLS;
for (int j = 0; j < C_COLS; ++j) {
outv[j] = __float2bfloat16(vals[j] * sv[j]);
}
// wait for the tma store that used this buffer 2 chunks ago
if (c >= 2 && warp == 0 && elect_one_sync()) {
tma_store_wait<1>();
}
asm volatile("barrier.sync 2, %0;" ::"n"(EPI_THREADS));
__nv_bfloat16* dst = c_smem + (c & 1) * C_ELEMS;
if (BM == 128) {
uint4* d4 = (uint4*)(dst + row * C_COLS);
const uint4* s4 = (const uint4*)outv;
for (int j = 0; j < C_COLS * 2 / 16; ++j) d4[j] = s4[j];
} else {
// m64: lanes 16*abuf..16*abuf+16 of this warp hold rows
// warp*16 + (lane - 16*abuf).
const int l16 = lane - 16 * abuf;
if (l16 >= 0 && l16 < 16) {
uint4* d4 = (uint4*)(dst + (warp * 16 + l16) * C_COLS);
const uint4* s4 = (const uint4*)outv;
for (int j = 0; j < C_COLS * 2 / 16; ++j) d4[j] = s4[j];
}
}
asm volatile("barrier.sync 2, %0;" ::"n"(EPI_THREADS));
if (warp == 0 && elect_one_sync()) {
fence_async_shared();
tma_store_2d(&c_map, c_smem + (c & 1) * C_ELEMS, off_n + c * C_COLS, off_m);
tma_store_commit();
}
}
if (++abuf == 2) {
abuf = 0;
af_phase ^= 1;
}
}
if (warp == 0 && elect_one_sync()) {
tma_store_wait<0>();
}
}
asm volatile("barrier.sync 3, 192;");
if (warp == EPI_WARPS + 1) {
tmem_dealloc(taddr_base, 512);
}
}
// ------------------------------------------------------------------ host
__global__ void k_pad2_funnel(const uint32_t* __restrict__ src_a, uint32_t* __restrict__ dst_a,
const uint32_t* __restrict__ src_b, uint32_t* __restrict__ dst_b,
int K, int dst_words, int src_words_a, int src_words_b, int rows_b) {
const int row = blockIdx.x;
const bool is_a = row < rows_b;
const uint32_t* src = is_a ? src_a : src_b;
uint32_t* dst = is_a ? dst_a : dst_b;
const int src_words = is_a ? src_words_a : src_words_b;
const long long sr = (long long)(is_a ? row : row - rows_b) * K;
const int w0 = (int)(sr >> 2);
const int sh = (int)(sr & 3) * 8;
for (int j = threadIdx.x; j < dst_words; j += blockDim.x) {
const int sw = w0 + j;
const unsigned a = (sw < src_words) ? src[sw] : 0u;
const unsigned b = (sw + 1 < src_words) ? src[sw + 1] : 0u;
unsigned w = (a >> sh);
if (sh != 0) w |= (b << (32 - sh));
const int byte0 = j * 4;
int vb = K - byte0;
vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb);
const unsigned keep = (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1);
dst[(size_t)(is_a ? row : row - rows_b) * dst_words + j] = w & keep;
}
}
// pad kernel: unpack odd-K fp8 rows into u32-aligned rows via funnel shifts
// (128-bit stores + shared adjacent-word loads)
__global__ void k_pad_funnel(const uint32_t* __restrict__ src, uint4* __restrict__ dst,
int K, int dst_words, int src_words) {
const int row = blockIdx.x;
const long long sr = (long long)row * K;
const int w0 = (int)(sr >> 2);
const int sh = (int)(sr & 3) * 8;
const int n_v = dst_words >> 2;
for (int j = threadIdx.x; j < n_v; j += blockDim.x) {
const int sw = w0 + (j << 2);
const unsigned s0 = (sw < src_words) ? src[sw] : 0u;
const unsigned s1 = (sw + 1 < src_words) ? src[sw + 1] : 0u;
const unsigned s2 = (sw + 2 < src_words) ? src[sw + 2] : 0u;
const unsigned s3 = (sw + 3 < src_words) ? src[sw + 3] : 0u;
const unsigned s4 = (sw + 4 < src_words) ? src[sw + 4] : 0u;
uint4 v;
if (sh == 0) {
v.x = s0; v.y = s1; v.z = s2; v.w = s3;
} else {
const int shr = 32 - sh;
v.x = (s0 >> sh) | (s1 << shr);
v.y = (s1 >> sh) | (s2 << shr);
v.z = (s2 >> sh) | (s3 << shr);
v.w = (s3 >> sh) | (s4 << shr);
}
const int base = (j << 2) * 4;
int vb;
vb = K - base; vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb);
v.x &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1);
vb = K - (base + 4); vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb);
v.y &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1);
vb = K - (base + 8); vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb);
v.z &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1);
vb = K - (base + 12); vb = vb < 0 ? 0 : (vb > 4 ? 4 : vb);
v.w &= (vb >= 4) ? 0xFFFFFFFFu : ((1u << (vb * 8)) - 1);
dst[(size_t)row * n_v + j] = v;
}
}
namespace {
struct MapCache {
std::unordered_map<uint64_t, CUtensorMap> maps;
CUtensorMap get(const void* base, uint64_t inner, uint64_t outer, uint32_t box_inner,
uint32_t box_outer, CUtensorMapSwizzle swz, CUtensorMapDataType dt,
int elembytes) {
uint64_t key = ((uint64_t)base) ^ (inner << 1) ^ (outer << 13) ^ ((uint64_t)box_inner << 40) ^
((uint64_t)box_outer << 48) ^ ((uint64_t)swz << 56) ^ ((uint64_t)dt << 61);
auto it = maps.find(key);
if (it != maps.end()) return it->second;
CUtensorMap map{};
uint64_t dims[2] = {inner, outer};
uint64_t strides[1] = {inner * (uint64_t)elembytes};
uint32_t box[2] = {box_inner, box_outer};
uint32_t estrides[2] = {1, 1};
cuTensorMapEncodeTiled(&map, dt, 2, const_cast<void*>(base), dims, strides, box,
estrides, CU_TENSOR_MAP_INTERLEAVE_NONE, swz,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
return map;
}
};
constexpr CUtensorMapSwizzle CuSw(int bk) {
return bk == 128 ? CU_TENSOR_MAP_SWIZZLE_128B : CU_TENSOR_MAP_SWIZZLE_64B;
}
MapCache g_maps;
int g_num_sms = 0;
} // namespace
template <int BM, int BN, int BK, int STAGES, int GM>
void launch_cfg(const CUtensorMap* a_map, const CUtensorMap* b_map, const CUtensorMap* c_map,
const float* scale, int M, int N, int K, cudaStream_t stream) {
constexpr int SMEM = STAGES * (BM * BK + BN * BK) + 2 * (BM * 64 * 2) + BN * 4 + 1024;
auto kern = k_gemm<BM, BN, BK, STAGES, GM>;
static bool attr_set = false;
if (!attr_set) {
cudaFuncSetAttribute(kern, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
attr_set = true;
}
int tiles_m = (M + BM - 1) / BM, tiles_n = (N + BN - 1) / BN;
int grid = std::min(g_num_sms, tiles_m * tiles_n);
kern<<<grid, 192, SMEM, stream>>>(*a_map, *b_map, *c_map, scale, M, N, K);
}
void fp8_run(int64_t x_ptr, int64_t w_ptr, int64_t s_ptr, int64_t y_ptr,
int64_t M, int64_t N, int64_t K, int64_t cfg, int64_t stream_ptr) {
void* x = (void*)x_ptr;
void* w = (void*)w_ptr;
if (!g_num_sms) {
cudaDeviceProp prop;
cudaGetDeviceProperties(&prop, 0);
g_num_sms = prop.multiProcessorCount;
}
cudaStream_t stream = (cudaStream_t)stream_ptr;
if (cfg == 0) {
CUtensorMap a_map = g_maps.get(x, K, M, 128, 64, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
CUtensorMap b_map = g_maps.get(w, K, N, 128, 64, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
CUtensorMap c_map = g_maps.get((void*)y_ptr, N, M, 64, 64, CU_TENSOR_MAP_SWIZZLE_NONE,
CU_TENSOR_MAP_DATA_TYPE_UINT16, 2);
launch_cfg<64, 64, 128, 12, 1>(&a_map, &b_map, &c_map, (const float*)s_ptr, (int)M, (int)N, (int)K, stream);
} else if (cfg == 1) {
CUtensorMap a_map = g_maps.get(x, K, M, 128, 128, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
CUtensorMap b_map = g_maps.get(w, K, N, 128, 256, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
CUtensorMap c_map = g_maps.get((void*)y_ptr, N, M, 64, 128, CU_TENSOR_MAP_SWIZZLE_NONE,
CU_TENSOR_MAP_DATA_TYPE_UINT16, 2);
launch_cfg<128, 256, 128, 4, 4>(&a_map, &b_map, &c_map, (const float*)s_ptr, (int)M, (int)N, (int)K, stream);
} else {
CUtensorMap a_map = g_maps.get(x, K, M, 128, 128, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
CUtensorMap b_map = g_maps.get(w, K, N, 128, 256, CuSw(128), CU_TENSOR_MAP_DATA_TYPE_UINT8, 1);
CUtensorMap c_map = g_maps.get((void*)y_ptr, N, M, 64, 128, CU_TENSOR_MAP_SWIZZLE_NONE,
CU_TENSOR_MAP_DATA_TYPE_UINT16, 2);
launch_cfg<128, 256, 128, 4, 16>(&a_map, &b_map, &c_map, (const float*)s_ptr, (int)M, (int)N, (int)K, stream);
}
}
void fp8_pad2(int64_t xa_ptr, int64_t xb_ptr, int64_t wa_ptr, int64_t wb_ptr,
int64_t Ma, int64_t Mb, int64_t K, int64_t K_pad, int64_t stream_ptr) {
const int dst_words = (int)(K_pad / 4);
const int sa = (int)(Ma * K / 4);
const int sb = (int)(Mb * K / 4);
k_pad2_funnel<<<(int)(Ma + Mb), 256, 0, (cudaStream_t)stream_ptr>>>(
(const uint32_t*)xa_ptr, (uint32_t*)xb_ptr, (const uint32_t*)wa_ptr, (uint32_t*)wb_ptr,
(int)K, dst_words, sa, sb, (int)Ma);
}
void fp8_pad(int64_t src_ptr, int64_t dst_ptr, int64_t M, int64_t K, int64_t K_pad,
int64_t stream_ptr) {
const int dst_words = (int)(K_pad / 4);
const int src_words = (int)(M * K / 4);
k_pad_funnel<<<(int)M, 256, 0, (cudaStream_t)stream_ptr>>>(
(const uint32_t*)src_ptr, (uint4*)dst_ptr, (int)K, dst_words, src_words);
}
"""
_ext = load_inline(
name="fp8_gemm_b200_v1",
cpp_sources=[r"""
void fp8_run(int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t);
void fp8_pad(int64_t, int64_t, int64_t, int64_t, int64_t, int64_t);
void fp8_pad2(int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t);
"""],
cuda_sources=[_CUDA_SRC],
functions=["fp8_run", "fp8_pad", "fp8_pad2"],
extra_cuda_cflags=["-O3", "-std=c++17", "--use_fast_math",
"-gencode=arch=compute_100a,code=sm_100a", "-lcuda"],
extra_ldflags=["-lcuda"],
verbose=False,
)
def _raw_stream():
return torch._C._cuda_getCurrentRawStream(torch.cuda.current_device())
class Model(nn.Module):
"""y = ((x_fp8 @ w_fp8.T) * weight_scale).to(bf16) on tcgen05 tensor cores."""
def __init__(self, M: int, N: int, K: int):
super().__init__()
self.M, self.N, self.K = M, N, K
w = torch.empty(N, K, dtype=torch.bfloat16)
nn.init.normal_(w, std=0.02)
s = (w.float().abs().amax(dim=1, keepdim=True) / E4M3_MAX).clamp(min=1e-12)
w_fp8 = (w.float() / s).to(torch.float8_e4m3fn)
self.register_buffer("weight", w_fp8) # (N, K) fp8
self.register_buffer("weight_scale", s.squeeze(1).to(torch.float32)) # (N,)
self._run = _ext.fp8_run
if M <= 64:
self._cfg = 0
elif N >= 8192:
self._cfg = 2
else:
self._cfg = 1
# pad state for unaligned K
self._w_pad = None # (version, buf)
self._x_pad = None # buf for x
def _pad(self, t, buf):
_ext.fp8_pad(t.data_ptr(), buf.data_ptr(), t.shape[0], t.shape[1], buf.shape[1],
_raw_stream())
return buf
def forward(self, x: torch.Tensor) -> torch.Tensor:
M, K = x.shape
y = torch.empty((M, self.N), dtype=torch.bfloat16, device=x.device)
stream = torch._C._cuda_getCurrentRawStream(0)
if K % 16 == 0:
w = self.weight
K0 = K
else:
K0 = (K + 15) // 16 * 16
if self._x_pad is None:
self._x_pad = torch.zeros((M, K0), dtype=x.dtype, device=x.device)
w = self.weight
ver = w._version
need_w = self._w_pad is None or self._w_pad[0] != ver
if need_w:
wbuf = torch.zeros((self.N, K0), dtype=w.dtype, device=w.device)
_ext.fp8_pad2(x.data_ptr(), self._x_pad.data_ptr(), w.data_ptr(), wbuf.data_ptr(),
M, self.N, K, K0, stream)
self._w_pad = (ver, wbuf)
w = wbuf
x = self._x_pad
else:
x = self._pad(x, self._x_pad)
w = self._w_pad[1]
self._run(x.data_ptr(), w.data_ptr(), self.weight_scale.data_ptr(), y.data_ptr(),
M, self.N, K0, self._cfg, stream)
return y
Model.__call__ = Model.forward
M = 4096
N = 4096
K = 4096
def get_inputs():
x = (torch.rand(M, K) * 8 - 4).to(torch.float8_e4m3fn)
return [x]
def get_init_inputs():
return [M, N, K]
20260715_220649_kinetic-claude_kinetic-0715_01_fp8_gemm