"""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 #include #include #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 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 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 __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(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(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(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(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(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(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 __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(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(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(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(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; static bool s1 = false; if (!s1) { cudaFuncSetAttribute(k1, cudaFuncAttributeMaxDynamicSharedMemorySize, sm1); s1 = true; } dim3 gr1((unsigned)(I / BN1), (unsigned)max_tiles); k1<<>>((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; static bool s2 = false; if (!s2) { cudaFuncSetAttribute(k2, cudaFuncAttributeMaxDynamicSharedMemorySize, sm2); s2 = true; } dim3 gr2((unsigned)(H / BN2), (unsigned)max_tiles); k2<<>>(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 #include 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, )