KernelBench cuda · RTX PRO 6000
GLM-5.2 Fused MoE DeepSeek V4.1 Flash
9.46%geomean peak fraction across shapes
manually audited: clean
DeepSeek V4.1 Flash hand-rolled the whole GLM-5.2 MoE layer in inline PTX: mma.sync bf16 tensor cores, ldmatrix, xor-swizzled cp.async, hist/scan/scatter token packing, one code path for every T. Every expert GEMM runs on real weights at full K and N; row padding makes it burn 25-50% more FLOPs than the roofline charges. Error margin 2.5x inside the gate.
harnessdeepseek-claudeagent session2h 54mtotal wall3h 6mcheck4mbenchmark5moutput tokens437,222cost$23.60gpu-lock wait52sgpu-lock held1h 55mregimecompute
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
shape 010.781 ms34.4%1.21 TB/s · 67% of 1.8 TB/s HBM · also 172 TFLOPS (34% of compute)
shape 110.751 ms34.8%1.21 TB/s · 67% of 1.8 TB/s HBM · also 174 TFLOPS (35% of compute)
shape 20.396 ms0.2%32.65 TB/s · 100% of 1.8 TB/s HBM · also 1 TFLOPS (0% of compute)
shape 314.236 ms52.1%261 TFLOPS · 52% of 500 TF bf16 peak · also 0.92 TB/s (51% of HBM)
shape 49.021 ms5.1%1.43 TB/s · 80% of 1.8 TB/s HBM · also 26 TFLOPS (5% of compute)
shape 59.293 ms9.8%1.39 TB/s · 77% of 1.8 TB/s HBM · also 49 TFLOPS (10% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(34.4% · 34.8% · 0.2% · 52.1% · 5.1% · 9.8%) = 9.5%
Kernel source (redacted)
"""GLM-5.2 fused MoE layer in raw CUDA (SM120) — KernelBench-CUDA.
Structure (per problem statement):
E=256 routed experts, top_k=8, n_shared=1 shared expert, H=4096, I=2048.
w1_routed (E,2I,H) packed gate|up, w2_routed (E,H,I)
w1_shared (S,2I,H), w2_shared (S,H,I)
out = sum_s silu(x W1s^T) . (x W2s^T) + sum_k w_k * silu(x W1[e_k]^T) . W2[e_k]^T
Implementation (one CUDA extension; no torch compute ops):
1. histogram of expert ids + exclusive scan -> BM-aligned padded row offsets
2. scatter -> sorted_token / sorted_weight, grouped by expert (vLLM-style
moe_align_block_size). The shared expert is appended as group index E.
3. grouped GEMM #1 (mma.sync m16n8k16 bf16, cp.async pipeline, xor-swizzled
smem): one CTA owns (expert, BM rows, BN columns of I); computes the gate
and up tiles together and writes h = silu(gate)*up as bf16.
4. grouped GEMM #2: h @ W2^T -> weighted atomic-add into fp32 out
5. cast fp32 -> bf16
"""
from __future__ import annotations
import os
import torch
import torch.nn as nn
os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0")
# make the venv's ninja visible to torch's cpp_extension bootstrap
import sys as _sys # noqa: E402
_bin = os.path.dirname(_sys.executable)
if os.path.isdir(_bin) and _bin not in os.environ.get("PATH", ""):
os.environ["PATH"] = _bin + os.pathsep + os.environ.get("PATH", "")
from torch.utils.cpp_extension import load_inline # noqa: E402
_CUDA_SRC = r"""
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <cstdint>
#define DEV __device__ __forceinline__
using bf16 = __nv_bfloat16;
// tile geometry (BK is fixed at 64 halves = 128 byte smem rows)
constexpr int BM1 = 128; // gate|up rows per CTA
constexpr int BN1 = 128; // gate|up columns-of-I per CTA
constexpr int ST1 = 2;
constexpr int BM2 = 128; // down rows per CTA
constexpr int BN2 = 256; // down columns of H per CTA
constexpr int ST2 = 2;
// The activation tile is re-read once per column tile, so its L2 cost scales
// as 1/BN; the weight tile is re-read once per row tile, so its cost scales as
// 1/BM. Making gate|up twice as wide halves the dominant A traffic. That
// costs BN*BK*2 = 16 KB more smem per stage, which the 99 KB budget only
// affords at two stages -- still a real pipeline, because the wait leaves one
// group in flight. 16 warps keep the accumulator count at 128 registers.
constexpr int WM1 = 4, WN1 = 4; // gate|up warp grid (16 warps = 512 threads)
constexpr int WM2 = 4, WN2 = 4; // down warp grid (16 warps = 512 threads)
// ---------------------------------------------------------------------------
// pipeline / mma helpers
// ---------------------------------------------------------------------------
DEV uint32_t smem_u32(const void* p) { return (uint32_t)__cvta_generic_to_shared(p); }
DEV void cp_async16(void* dst, const void* src) {
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;\n" ::"r"(smem_u32(dst)), "l"(src));
}
DEV void cp_commit() { asm volatile("cp.async.commit_group;\n"); }
template <int N>
DEV void cp_wait() { asm volatile("cp.async.wait_group %0;\n" ::"n"(N)); }
DEV void ldsm4(uint32_t& r0, uint32_t& r1, uint32_t& r2, uint32_t& r3, uint32_t addr) {
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];\n"
: "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3)
: "r"(addr));
}
DEV void mma_16816(float* c, const uint32_t* a, uint32_t b0, uint32_t b1) {
asm volatile(
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};\n"
: "+f"(c[0]), "+f"(c[1]), "+f"(c[2]), "+f"(c[3])
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b0), "r"(b1));
}
DEV void red_add2(float* p, float a, float b) {
asm volatile("red.global.add.v2.f32 [%0], {%1,%2};\n" ::"l"(p), "f"(a), "f"(b) : "memory");
}
DEV float silu_f(float x) { return x / (1.0f + __expf(-x)); }
DEV int count_nonempty(const int* counts, int e) { return __ldg(&counts[e]); }
// xor-swizzled byte offset inside a row-major smem tile with 128-byte rows
// (BK = 64 halves -> 8 chunks of 16 bytes). chunk in [0,8).
template <int BK>
DEV uint32_t swz(int row, int chunk) {
return (uint32_t)(row * (BK * 2) + ((chunk ^ (row & 7)) << 4));
}
// ---------------------------------------------------------------------------
// prep kernels
// ---------------------------------------------------------------------------
__global__ void k_count(const long long* __restrict__ ids, int n, int E, int* __restrict__ counts) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < n) atomicAdd(&counts[(int)ids[i]], 1);
}
// one block: exclusive scan of per-expert counts on BM-padded boundaries.
// counts[E] is forced to T (the shared expert sees every token).
// NOTE: counts_in and counts_out intentionally alias the same buffer; the scan
// snapshot lives in shared memory, so no __restrict__ here.
__global__ void k_scan(const int* counts_in, int E, int T, int BM,
int* __restrict__ row_off, int* __restrict__ tile_expert,
int* counts_out, int* __restrict__ num_tiles) {
// smem: two ping-pong rows of E+2 ints.
extern __shared__ int sm[];
const int n = E + 1; // experts 0..E-1 plus the shared pseudo-expert at E
const int i = threadIdx.x;
if (n > (int)blockDim.x) {
// serialize for expert counts that do not fit in one block
if (i == 0) {
int acc = 0, t = 0;
for (int e = 0; e <= E; ++e) {
int v = (e < E) ? __ldg(&counts_in[e]) : T;
row_off[e] = acc;
counts_out[e] = v;
int nt = (v + BM - 1) / BM;
for (int j = 0; j < nt; ++j) tile_expert[t++] = e;
acc += nt * BM;
}
row_off[E + 1] = acc;
*num_tiles = t;
}
return;
}
int* s = sm;
int* d = sm + (E + 2);
int tc = 0;
if (i < n) {
int c = (i < E) ? __ldg(&counts_in[i]) : T;
counts_out[i] = c;
tc = (c + BM - 1) / BM; // tiles this group occupies
s[i] = tc;
}
__syncthreads();
// Hillis-Steele inclusive scan of the tile counts
for (int off = 1; off < n; off <<= 1) {
int v = (i < n) ? s[i] : 0;
int p = (i < n && i >= off) ? s[i - off] : 0;
if (i < n) d[i] = v + p;
__syncthreads();
int* tmp = s; s = d; d = tmp;
__syncthreads();
}
if (i < n) {
const int base = s[i] - tc; // exclusive tile prefix
row_off[i] = base * BM;
if (i == E) {
row_off[E + 1] = s[i] * BM;
*num_tiles = s[i];
}
#pragma unroll 4
for (int j = 0; j < tc; ++j) tile_expert[base + j] = i;
}
}
// group the (token,k) assignment list by expert
__global__ void k_scatter(const long long* __restrict__ ids, const bf16* __restrict__ wts,
int n_routed, int top_k, const int* __restrict__ row_off,
int* __restrict__ cursor, int* __restrict__ sorted_token,
float* __restrict__ sorted_weight) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < n_routed) {
int e = (int)ids[i];
int pos = atomicAdd(&cursor[e], 1);
int row = row_off[e] + pos;
sorted_token[row] = i / top_k;
sorted_weight[row] = __bfloat162float(wts[i]);
}
}
// shared expert rows: every token, weight 1
__global__ void k_shared_rows(int T, const int* __restrict__ row_off, int E,
int* __restrict__ sorted_token, float* __restrict__ sorted_weight) {
int t = blockIdx.x * blockDim.x + threadIdx.x;
if (t < T) {
int row = row_off[E] + t;
sorted_token[row] = t;
sorted_weight[row] = 1.0f;
}
}
// ---------------------------------------------------------------------------
// grouped GEMM 1: gate|up
// ---------------------------------------------------------------------------
template <int BM, int BN, int STAGES, int WM, int WN>
__global__ void __launch_bounds__(WM* WN * 32, 1) k_gate_up(
const bf16* __restrict__ x, int H, const int* __restrict__ tile_expert,
const int* __restrict__ row_off, const int* __restrict__ counts,
const int* __restrict__ ntiles, const int* __restrict__ sorted_token,
const bf16* __restrict__ w1r, const bf16* __restrict__ w1s, int E, int I,
bf16* __restrict__ hbuf) {
constexpr int BK = 64;
constexpr int NT = WM * WN * 32;
constexpr int AU = BM * 8;
constexpr int BU = BN * 8;
extern __shared__ char smem[];
bf16* sA = (bf16*)smem; // STAGES * BM * BK
bf16* sBg = sA + (size_t)STAGES * BM * BK; // STAGES * BN * BK
bf16* sBu = sBg + (size_t)STAGES * BN * BK; // STAGES * BN * BK
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
const int wm = warp % WM;
const int wn = warp / WM;
constexpr int WTM = BM / WM; // warp tile rows
constexpr int WTN = BN / WN; // warp tile cols
const int pid_n = blockIdx.x;
const int pid_m = blockIdx.y;
if (pid_m >= *ntiles) return;
const int e = __ldg(&tile_expert[pid_m]);
const int row0 = pid_m * BM;
const int cnt = count_nonempty(counts, e);
const int mbase = row0 - __ldg(&row_off[e]);
const bf16* W1 = (e < E) ? (w1r + (size_t)e * 2 * I * H) : w1s;
const bf16* Wg = W1;
const bf16* Wu = W1 + (size_t)I * H;
const int n0 = pid_n * BN;
const int NK = H / BK;
// The gather index and the smem destination of every A chunk are invariant
// across the K loop, so resolve them once instead of re-loading the token
// inside the pipeline (that load otherwise sits in the issue path of every
// cp.async).
constexpr int AIT = (AU + NT - 1) / NT;
const bf16* abase[AIT];
int asox[AIT];
bool arun[AIT], aok[AIT];
#pragma unroll
for (int it = 0; it < AIT; ++it) {
int u = tid + it * NT;
bool run = u < AU;
bool valid = run && ((mbase + (u >> 3)) < cnt);
int tok = valid ? __ldg(&sorted_token[row0 + (u >> 3)]) : 0;
abase[it] = x + (size_t)tok * H + (size_t)((u & 7) * 8);
asox[it] = run ? swz<BK>(u >> 3, u & 7) / 2 : 0;
arun[it] = run;
aok[it] = valid;
}
constexpr int BIT = (2 * BU + NT - 1) / NT;
const char* bbase[BIT];
int bsox[BIT];
bool bok[BIT];
#pragma unroll
for (int it = 0; it < BIT; ++it) {
int u = tid + it * NT;
bool valid = u < 2 * BU;
int half = valid && (u >= BU);
int v = half ? (u - BU) : u;
bbase[it] = (const char*)(half ? Wu : Wg) + (size_t)(n0 + (v >> 3)) * H * 2 +
(size_t)((v & 7) * 8) * 2;
bsox[it] = valid ? swz<BK>(v >> 3, v & 7) / 2 : 0;
bok[it] = valid;
}
auto load_a = [&](int stage, int k0) {
bf16* sb = sA + (size_t)stage * BM * BK;
#pragma unroll
for (int it = 0; it < AIT; ++it) {
if (!arun[it]) continue;
bf16* sp = sb + asox[it];
if (aok[it])
cp_async16(sp, (const char*)(abase[it] + k0));
else
*reinterpret_cast<uint4*>(sp) = make_uint4(0, 0, 0, 0);
}
};
auto load_b = [&](int stage, int k0) {
#pragma unroll
for (int it = 0; it < BIT; ++it) {
if (!bok[it]) continue;
int u = tid + it * NT;
int half = u >= BU;
bf16* sp = (half ? sBu : sBg) + (size_t)stage * BN * BK + bsox[it];
cp_async16(sp, bbase[it] + (size_t)k0 * 2);
}
};
float accg[WTM / 16][WTN / 8][4];
float accu[WTM / 16][WTN / 8][4];
#pragma unroll
for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
for (int ni = 0; ni < WTN / 8; ++ni)
#pragma unroll
for (int q = 0; q < 4; ++q) {
accg[mi][ni][q] = 0.f;
accu[mi][ni][q] = 0.f;
}
#pragma unroll 1
for (int s = 0; s < STAGES - 1; ++s) {
if (s < NK) {
load_a(s, s * BK);
load_b(s, s * BK);
}
cp_commit();
}
#pragma unroll 1
for (int ks = 0; ks < NK; ++ks) {
int lk = ks + STAGES - 1;
if (lk < NK) {
load_a(lk % STAGES, lk * BK);
load_b(lk % STAGES, lk * BK);
}
cp_commit();
cp_wait<1>();
__syncthreads();
const int st = ks % STAGES;
const bf16* sAb = sA + (size_t)st * BM * BK;
const bf16* sBgb = sBg + (size_t)st * BN * BK;
const bf16* sBub = sBu + (size_t)st * BN * BK;
#pragma unroll
for (int kk = 0; kk < BK / 16; ++kk) {
const int c0 = 2 * kk;
uint32_t aA[WTM / 16][4];
#pragma unroll
for (int mi = 0; mi < WTM / 16; ++mi) {
int row = wm * WTM + mi * 16 + ((lane >> 3) & 1) * 8 + (lane & 7);
int ch = c0 + (lane >> 4);
ldsm4(aA[mi][0], aA[mi][1], aA[mi][2], aA[mi][3], smem_u32(sAb) + swz<BK>(row, ch));
}
uint32_t bg[WTN / 8][2], bu[WTN / 8][2];
#pragma unroll
for (int h2 = 0; h2 < (WTN / 8) / 2; ++h2) {
int row = wn * WTN + h2 * 16 + ((lane >> 4) & 1) * 8 + (lane & 7);
int ch = c0 + ((lane >> 3) & 1);
ldsm4(bg[h2 * 2][0], bg[h2 * 2][1], bg[h2 * 2 + 1][0], bg[h2 * 2 + 1][1],
smem_u32(sBgb) + swz<BK>(row, ch));
ldsm4(bu[h2 * 2][0], bu[h2 * 2][1], bu[h2 * 2 + 1][0], bu[h2 * 2 + 1][1],
smem_u32(sBub) + swz<BK>(row, ch));
}
#pragma unroll
for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
for (int ni = 0; ni < WTN / 8; ++ni) {
mma_16816(accg[mi][ni], aA[mi], bg[ni][0], bg[ni][1]);
mma_16816(accu[mi][ni], aA[mi], bu[ni][0], bu[ni][1]);
}
}
__syncthreads();
}
// epilogue: h = silu(gate) * up
#pragma unroll
for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
for (int ni = 0; ni < WTN / 8; ++ni) {
int r = row0 + wm * WTM + mi * 16 + (lane >> 2);
int col = n0 + wn * WTN + ni * 8 + 2 * (lane & 3);
float h0 = silu_f(accg[mi][ni][0]) * accu[mi][ni][0];
float h1 = silu_f(accg[mi][ni][1]) * accu[mi][ni][1];
float h2 = silu_f(accg[mi][ni][2]) * accu[mi][ni][2];
float h3 = silu_f(accg[mi][ni][3]) * accu[mi][ni][3];
__nv_bfloat162 p0 = __floats2bfloat162_rn(h0, h1);
__nv_bfloat162 p1 = __floats2bfloat162_rn(h2, h3);
*(uint32_t*)(hbuf + (size_t)r * I + col) = *(uint32_t*)&p0;
*(uint32_t*)(hbuf + (size_t)(r + 8) * I + col) = *(uint32_t*)&p1;
}
}
// ---------------------------------------------------------------------------
// grouped GEMM 2: h @ W2^T -> weighted atomic add into fp32 out
// ---------------------------------------------------------------------------
template <int BM, int BN, int STAGES, int WM, int WN>
__global__ void __launch_bounds__(WM* WN * 32, 1) k_down(
const bf16* __restrict__ hbuf, int I, const int* __restrict__ tile_expert,
const int* __restrict__ row_off, const int* __restrict__ counts,
const int* __restrict__ ntiles, const int* __restrict__ sorted_token,
const float* __restrict__ sorted_weight, const bf16* __restrict__ w2r,
const bf16* __restrict__ w2s, int E, int H, float* __restrict__ out) {
constexpr int BK = 64;
constexpr int NT = WM * WN * 32;
constexpr int AU = BM * 8;
constexpr int BU = BN * 8;
extern __shared__ char smem[];
bf16* sA = (bf16*)smem;
bf16* sB = sA + (size_t)STAGES * BM * BK;
constexpr int WTM = BM / WM;
constexpr int WTN = BN / WN;
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
const int wm = warp % WM;
const int wn = warp / WM;
const int pid_n = blockIdx.x;
const int pid_m = blockIdx.y;
if (pid_m >= *ntiles) return;
const int e = __ldg(&tile_expert[pid_m]);
const int row0 = pid_m * BM;
const int cnt = count_nonempty(counts, e);
const int mbase = row0 - __ldg(&row_off[e]);
const int n0 = pid_n * BN;
const bf16* W2 = (e < E) ? (w2r + (size_t)e * H * I) : w2s;
const int NK = I / BK;
auto load_a = [&](int stage, int k0) {
#pragma unroll
for (int it = 0; it < (AU + NT - 1) / NT; ++it) {
int u = tid + it * NT;
if (u < AU) {
int r = u >> 3;
int ch = u & 7;
const char* src =
(const char*)hbuf + (size_t)(row0 + r) * I * 2 + (size_t)(k0 + ch * 8) * 2;
bf16* sp = sA + (size_t)stage * BM * BK + swz<BK>(r, ch) / 2;
cp_async16(sp, src);
}
}
};
auto load_b = [&](int stage, int k0) {
#pragma unroll
for (int it = 0; it < (BU + NT - 1) / NT; ++it) {
int u = tid + it * NT;
if (u < BU) {
int r = u >> 3;
int ch = u & 7;
const char* src =
(const char*)W2 + (size_t)(n0 + r) * I * 2 + (size_t)(k0 + ch * 8) * 2;
bf16* sp = sB + (size_t)stage * BN * BK + swz<BK>(r, ch) / 2;
cp_async16(sp, src);
}
}
};
float acc[WTM / 16][WTN / 8][4];
#pragma unroll
for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
for (int ni = 0; ni < WTN / 8; ++ni)
#pragma unroll
for (int q = 0; q < 4; ++q) acc[mi][ni][q] = 0.f;
#pragma unroll 1
for (int s = 0; s < STAGES - 1; ++s) {
if (s < NK) {
load_a(s, s * BK);
load_b(s, s * BK);
}
cp_commit();
}
#pragma unroll 1
for (int ks = 0; ks < NK; ++ks) {
int lk = ks + STAGES - 1;
if (lk < NK) {
load_a(lk % STAGES, lk * BK);
load_b(lk % STAGES, lk * BK);
}
cp_commit();
cp_wait<1>();
__syncthreads();
const int st = ks % STAGES;
const bf16* sAb = sA + (size_t)st * BM * BK;
const bf16* sBb = sB + (size_t)st * BN * BK;
#pragma unroll
for (int kk = 0; kk < BK / 16; ++kk) {
const int c0 = 2 * kk;
uint32_t aA[WTM / 16][4];
#pragma unroll
for (int mi = 0; mi < WTM / 16; ++mi) {
int row = wm * WTM + mi * 16 + ((lane >> 3) & 1) * 8 + (lane & 7);
int ch = c0 + (lane >> 4);
ldsm4(aA[mi][0], aA[mi][1], aA[mi][2], aA[mi][3], smem_u32(sAb) + swz<BK>(row, ch));
}
uint32_t bb[WTN / 8][2];
#pragma unroll
for (int h2 = 0; h2 < (WTN / 8) / 2; ++h2) {
int row = wn * WTN + h2 * 16 + ((lane >> 4) & 1) * 8 + (lane & 7);
int ch = c0 + ((lane >> 3) & 1);
ldsm4(bb[h2 * 2][0], bb[h2 * 2][1], bb[h2 * 2 + 1][0], bb[h2 * 2 + 1][1],
smem_u32(sBb) + swz<BK>(row, ch));
}
#pragma unroll
for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
for (int ni = 0; ni < WTN / 8; ++ni) mma_16816(acc[mi][ni], aA[mi], bb[ni][0], bb[ni][1]);
}
__syncthreads();
}
// epilogue: out[token] += w * y. The two rows of each m16 fragment
// (r and r+8) must be guarded independently: rows past the group's count
// carry no sorted_token entry, so writing them would scatter atoms of the
// output all over the device heap.
const int nvalid = cnt - mbase;
#pragma unroll
for (int mi = 0; mi < WTM / 16; ++mi)
#pragma unroll
for (int ni = 0; ni < WTN / 8; ++ni) {
int col = n0 + wn * WTN + ni * 8 + 2 * (lane & 3);
int l0 = wm * WTM + mi * 16 + (lane >> 2);
if (l0 < nvalid) {
int tok = __ldg(&sorted_token[row0 + l0]);
float w = __ldg(&sorted_weight[row0 + l0]);
red_add2(out + (size_t)tok * H + col, w * acc[mi][ni][0], w * acc[mi][ni][1]);
}
int l1 = l0 + 8;
if (l1 < nvalid) {
int tok = __ldg(&sorted_token[row0 + l1]);
float w = __ldg(&sorted_weight[row0 + l1]);
red_add2(out + (size_t)tok * H + col, w * acc[mi][ni][2], w * acc[mi][ni][3]);
}
}
}
// ---------------------------------------------------------------------------
// cast
// ---------------------------------------------------------------------------
// n is the element count; each thread converts a pair of floats.
__global__ void k_cast(const float* __restrict__ in, bf16* __restrict__ outp, long long n) {
long long i = ((long long)blockIdx.x * blockDim.x + threadIdx.x) * 2;
if (i >= n) return;
float2 v = *(const float2*)(in + i);
__nv_bfloat162 b = __floats2bfloat162_rn(v.x, v.y);
*(uint32_t*)(outp + i) = *(uint32_t*)&b;
}
// ---------------------------------------------------------------------------
// host side
// ---------------------------------------------------------------------------
struct Scratch {
int* counts = nullptr; // E+1
int* cursor = nullptr; // E
int* row_off = nullptr; // E+2
int* tile_expert = nullptr; // max_tiles
int* num_tiles = nullptr; // 1
int* [REDACTED credential assignment]; // max_rows
float* sorted_weight = nullptr;
bf16* hbuf = nullptr; // max_rows * I
float* out32 = nullptr; // T * H
long long cap_rows = 0, cap_T = 0, cap_tiles = 0;
};
static Scratch g;
static void ensure_scratch(long long max_rows, long long T, int H, int I, long long max_tiles) {
if (g.cap_rows >= max_rows && g.cap_T >= T && g.cap_tiles >= max_tiles) return;
if (g.sorted_token) cudaFree(g.sorted_token);
if (g.sorted_weight) cudaFree(g.sorted_weight);
if (g.hbuf) cudaFree(g.hbuf);
if (g.out32) cudaFree(g.out32);
if (g.tile_expert) cudaFree(g.tile_expert);
[REDACTED credential assignment];
g.sorted_weight = nullptr;
g.hbuf = nullptr;
g.out32 = nullptr;
g.tile_expert = nullptr;
size_t r = (size_t)max_rows;
cudaMalloc(&g.sorted_token, r * sizeof(int));
cudaMalloc(&g.sorted_weight, r * sizeof(float));
cudaMalloc(&g.hbuf, r * (size_t)I * sizeof(bf16));
cudaMalloc(&g.out32, (size_t)T * H * sizeof(float));
cudaMalloc(&g.tile_expert, (size_t)max_tiles * sizeof(int));
g.cap_rows = max_rows;
g.cap_T = T;
g.cap_tiles = max_tiles;
}
static void ensure_static(int E) {
static int cap_E = 0;
if (cap_E >= E + 2 && g.counts) return;
if (g.counts) cudaFree(g.counts);
if (g.cursor) cudaFree(g.cursor);
if (g.row_off) cudaFree(g.row_off);
if (g.num_tiles) cudaFree(g.num_tiles);
cudaMalloc(&g.counts, (E + 1) * sizeof(int));
cudaMalloc(&g.cursor, E * sizeof(int));
cudaMalloc(&g.row_off, (E + 2) * sizeof(int));
cudaMalloc(&g.num_tiles, sizeof(int));
cap_E = E + 2;
}
extern "C" void moe_launch(uintptr_t x, uintptr_t ids, uintptr_t wts, uintptr_t w1r,
uintptr_t w2r, uintptr_t w1s, uintptr_t w2s, uintptr_t out_bf,
long long T, long long E, long long top_k, long long H, long long I,
long long stream) {
cudaStream_t st = (cudaStream_t)stream;
long long max_rows = T * (top_k + 1) + (E + 1) * (BM1 - 1) + BM1;
// Tiles per expert are padded to BM1 rows, so the count is bounded by what a
// single fully packed expert would need, plus one partial tile per routed
// expert, plus the shared expert's tiles. A tight bound keeps the mostly
// empty grid rows out of the small-T shapes.
long long n_routed_tok = T * top_k;
long long partial = n_routed_tok < E ? n_routed_tok : E;
long long max_tiles = (n_routed_tok + BM1 - 1) / BM1 + partial + (T + BM1 - 1) / BM1;
ensure_static((int)E);
ensure_scratch(max_rows, T, (int)H, (int)I, max_tiles);
int n_routed = (int)(T * top_k);
cudaMemsetAsync(g.counts, 0, (E + 1) * sizeof(int), st);
cudaMemsetAsync(g.cursor, 0, E * sizeof(int), st);
cudaMemsetAsync(g.out32, 0, (size_t)T * H * sizeof(float), st);
k_count<<<(n_routed + 255) / 256, 256, 0, st>>>((const long long*)ids, n_routed, (int)E, g.counts);
k_scan<<<1, 512, (int)(2 * (E + 2) * sizeof(int)), st>>>(g.counts, (int)E, (int)T, BM1, g.row_off,
g.tile_expert, g.counts, g.num_tiles);
k_scatter<<<(n_routed + 255) / 256, 256, 0, st>>>((const long long*)ids, (const bf16*)wts,
n_routed, (int)top_k, g.row_off, g.cursor,
g.sorted_token, g.sorted_weight);
k_shared_rows<<<(int)((T + 255) / 256), 256, 0, st>>>((int)T, g.row_off, (int)E,
g.sorted_token, g.sorted_weight);
int sm1 = (ST1 * BM1 * 64 + 2 * ST1 * BN1 * 64) * (int)sizeof(bf16);
auto k1 = k_gate_up<BM1, BN1, ST1, WM1, WN1>;
static bool s1 = false;
if (!s1) {
cudaFuncSetAttribute(k1, cudaFuncAttributeMaxDynamicSharedMemorySize, sm1);
s1 = true;
}
dim3 gr1((unsigned)(I / BN1), (unsigned)max_tiles);
k1<<<gr1, WM1 * WN1 * 32, sm1, st>>>((const bf16*)x, (int)H, g.tile_expert, g.row_off, g.counts,
g.num_tiles, g.sorted_token, (const bf16*)w1r,
(const bf16*)w1s, (int)E, (int)I, g.hbuf);
int sm2 = (ST2 * BM2 * 64 + ST2 * BN2 * 64) * (int)sizeof(bf16);
auto k2 = k_down<BM2, BN2, ST2, WM2, WN2>;
static bool s2 = false;
if (!s2) {
cudaFuncSetAttribute(k2, cudaFuncAttributeMaxDynamicSharedMemorySize, sm2);
s2 = true;
}
dim3 gr2((unsigned)(H / BN2), (unsigned)max_tiles);
k2<<<gr2, WM2 * WN2 * 32, sm2, st>>>(g.hbuf, (int)I, g.tile_expert, g.row_off, g.counts,
g.num_tiles, g.sorted_token, g.sorted_weight,
(const bf16*)w2r, (const bf16*)w2s, (int)E, (int)H,
g.out32);
long long n2 = T * H;
k_cast<<<(unsigned)((n2 / 2 + 255) / 256), 256, 0, st>>>(g.out32, (bf16*)out_bf, n2);
}
"""
_CPP_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
extern "C" void moe_launch(uintptr_t x, uintptr_t ids, uintptr_t wts, uintptr_t w1r,
uintptr_t w2r, uintptr_t w1s, uintptr_t w2s, uintptr_t out_bf,
long long T, long long E, long long top_k, long long H, long long I,
long long stream);
torch::Tensor moe_forward(torch::Tensor x, torch::Tensor ids, torch::Tensor wts,
torch::Tensor w1r, torch::Tensor w2r, torch::Tensor w1s,
torch::Tensor w2s, int64_t top_k) {
TORCH_CHECK(x.is_cuda() && x.dtype() == torch::kBFloat16);
TORCH_CHECK(x.is_contiguous() && ids.is_contiguous() && wts.is_contiguous());
TORCH_CHECK(w1r.is_contiguous() && w2r.is_contiguous());
TORCH_CHECK(w1s.is_contiguous() && w2s.is_contiguous());
long long T = x.size(0);
long long H = x.size(1);
long long E = w1r.size(0);
long long I = w2r.size(2);
auto out = torch::empty({T, H}, x.options());
cudaStream_t st = at::cuda::getCurrentCUDAStream();
moe_launch((uintptr_t)x.data_ptr(), (uintptr_t)ids.data_ptr(), (uintptr_t)wts.data_ptr(),
(uintptr_t)w1r.data_ptr(), (uintptr_t)w2r.data_ptr(), (uintptr_t)w1s.data_ptr(),
(uintptr_t)w2s.data_ptr(), (uintptr_t)out.data_ptr(), T, E, top_k, H, I,
(long long)st);
return out;
}
"""
ext = load_inline(
name="glm52_moe_cuda",
cpp_sources=_CPP_SRC,
cuda_sources=_CUDA_SRC,
functions=["moe_forward"],
extra_cuda_cflags=["-O3", "-std=c++17", "-lineinfo"],
verbose=False,
)
class Model(nn.Module):
def __init__(self, T, E, top_k, n_shared, H, I):
super().__init__()
self.T, self.E, self.top_k = T, E, top_k
self.n_shared, self.H, self.I = n_shared, H, I
self.w1_routed = nn.Parameter(torch.empty(E, 2 * I, H, dtype=torch.bfloat16))
self.w2_routed = nn.Parameter(torch.empty(E, H, I, dtype=torch.bfloat16))
self.w1_shared = nn.Parameter(torch.empty(n_shared, 2 * I, H, dtype=torch.bfloat16))
self.w2_shared = nn.Parameter(torch.empty(n_shared, H, I, dtype=torch.bfloat16))
for p in self.parameters():
nn.init.normal_(p, std=0.02)
def forward(self, x, expert_ids, expert_weights):
return ext.moe_forward(
x.contiguous(),
expert_ids.contiguous(),
expert_weights.contiguous(),
self.w1_routed,
self.w2_routed,
self.w1_shared,
self.w2_shared,
self.top_k,
)
20260910_115701_deepseek-claude_deepseek-flash_01_glm52_fused_moe