KernelBench cuda · RTX PRO 6000
GLM-5.2 Fused MoE Grok 4.7
Grok 4.7 wrote real fused-MoE plumbing in CUDA - histogram, expert-grouped gather with per-expert padding, SiLU-and-multiply, weighted scatter-reduce, and a hand-written decode GEMV - but every routed and shared GEMM is `cublasGemmStridedBatchedEx` (solution.py:297-301), not authored MMA. Same class as the published grok-4.6 (0.0939) and deepseek-v4-pro (0.0968) cells on this problem, both annotated `interesting` for the same reason: cuBLAS is not on the forbidden list, but the peak measures a library GEMM. Routing is taken as given, the shared expert always fires, every GEMM runs bf16-in / fp32-accum at full K and N on the real weights, nothing is precomputed outside the timed window, and the correctness margin is 13x inside the gate.
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(33.3% · 32.7% · 0.3% · 41.2% · 5.5% · 10.4%) = 9.5%
Kernel source (redacted)
"""GLM-5.2 fused MoE: CUDA gather + cuBLAS tensor-core GEMM + CUDA epilogue.
Per expert: h = silu(x @ gate.T) * (x @ up.T); y = h @ down.T
out = sum_shared(y_s) + sum_k weight[t, k] * y_routed[e_k]
Large T: tokens grouped by expert, padded to a cuBLAS-friendly M, one
strided-batched GEMM per projection (weights read once). T=1: fused GEMV
over the 8 routed experts + the shared expert only.
"""
from __future__ import annotations
import os
os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0")
os.environ.setdefault("CUDA_HOME", "/usr/local/cuda")
os.environ.setdefault("MAX_JOBS", "4")
import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline
_CUDA = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cublas_v2.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#include <algorithm>
#include <cstdint>
#include <vector>
#define CHECK_CUBLAS(expr) \
do { \
cublasStatus_t _st = (expr); \
TORCH_CHECK(_st == CUBLAS_STATUS_SUCCESS, "cublas ", int(_st), " @ ", \
__LINE__); \
} while (0)
struct Ctx {
cublasHandle_t handle = nullptr;
int device = -1;
torch::Tensor workspace;
torch::Tensor counts_h; // pinned int32
torch::Tensor Xg, gate_up, hidden, Y;
torch::Tensor shared_gate, shared_h, shared_y;
torch::Tensor counts_d, cursor_d, row_of_pair;
int cap_E = 0;
int cap_pairs = 0;
int64_t cap_rows = 0;
int cap_T = 0;
int cap_I = 0;
int cap_H = 0;
};
static Ctx& ctx() {
static Ctx c;
return c;
}
static void ensure_handle() {
Ctx& c = ctx();
int dev = 0;
cudaGetDevice(&dev);
if (c.handle && c.device == dev) {
CHECK_CUBLAS(cublasSetStream(c.handle, at::cuda::getCurrentCUDAStream()));
return;
}
if (c.handle) {
cublasDestroy(c.handle);
c.handle = nullptr;
}
c.device = dev;
CHECK_CUBLAS(cublasCreate(&c.handle));
CHECK_CUBLAS(cublasSetStream(c.handle, at::cuda::getCurrentCUDAStream()));
cublasSetMathMode(c.handle, CUBLAS_DEFAULT_MATH);
auto opt = torch::TensorOptions().device(torch::kCUDA).dtype(torch::kUInt8);
c.workspace = torch::empty({256 << 20}, opt);
CHECK_CUBLAS(cublasSetWorkspace_v2(c.handle, c.workspace.data_ptr(),
c.workspace.numel()));
}
// Measured on RTX PRO 6000 sm_120: strided-batched bf16 GEMM falls off a cliff
// just above M=256 (7.5 ms -> 10.5 ms) and again above 384. Stay on the fast tiles.
static int choose_mpad(int maxc) {
if (maxc <= 0) return 0;
const int good[] = {32, 48, 64, 96, 112, 128, 160, 176, 192,
208, 224, 240, 256, 384};
for (int g : good) {
if (g >= maxc) return g;
}
return (maxc + 63) & ~63;
}
__device__ __forceinline__ float silu_f(float x) {
float y = expf(-fabsf(x));
float sig = (x >= 0.f) ? (1.f / (1.f + y)) : (y / (1.f + y));
return x * sig;
}
__device__ __forceinline__ float dot8(uint4 a, uint4 b) {
const __nv_bfloat16* ap = reinterpret_cast<const __nv_bfloat16*>(&a);
const __nv_bfloat16* bp = reinterpret_cast<const __nv_bfloat16*>(&b);
float acc = 0.f;
#pragma unroll
for (int j = 0; j < 8; ++j) {
acc = fmaf(__bfloat162float(ap[j]), __bfloat162float(bp[j]), acc);
}
return acc;
}
__device__ __forceinline__ uint4 ldcs_u4(const uint4* p) {
uint4 v;
asm volatile("ld.global.cs.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(v.x), "=r"(v.y), "=r"(v.z), "=r"(v.w)
: "l"(p));
return v;
}
__global__ void histogram_i64(const int64_t* __restrict__ ids, int n, int E,
int* __restrict__ counts) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < n) {
int64_t e = ids[i];
if ((unsigned)e < (unsigned)E) atomicAdd(counts + e, 1);
}
}
// One block per (token, slot). Assigns a row inside the expert's M_pad tile and copies x.
__global__ void assign_and_gather(const int64_t* __restrict__ ids,
const __nv_bfloat16* __restrict__ x,
__nv_bfloat16* __restrict__ Xg,
int* __restrict__ cursor,
int* __restrict__ row_of_pair,
int n_pairs, int top_k, int H, int M_pad) {
int pair = blockIdx.x;
if (pair >= n_pairs) return;
int t = pair / top_k;
int64_t e64 = ids[pair];
int e = (int)e64;
int local = 0;
if (threadIdx.x == 0) {
local = atomicAdd(cursor + e, 1);
row_of_pair[pair] = e * M_pad + local;
}
local = row_of_pair[pair] - e * M_pad;
// row_of_pair is written by thread 0; publish before the copy.
__syncthreads();
int row = row_of_pair[pair];
const uint4* src =
reinterpret_cast<const uint4*>(x + (int64_t)t * H);
uint4* dst = reinterpret_cast<uint4*>(Xg + (int64_t)row * H);
int nvec = H >> 3;
for (int i = threadIdx.x; i < nvec; i += blockDim.x) {
dst[i] = src[i];
}
(void)local;
}
__global__ void zero_pad_rows(const int* __restrict__ counts,
__nv_bfloat16* __restrict__ Xg, int E, int M_pad,
int H) {
int e = blockIdx.y;
if (e >= E) return;
int local = counts[e] + blockIdx.x;
if (local >= M_pad) return;
uint4* dst =
reinterpret_cast<uint4*>(Xg + ((int64_t)e * M_pad + local) * H);
int nvec = H >> 3;
for (int i = threadIdx.x; i < nvec; i += blockDim.x) {
uint4 z = make_uint4(0, 0, 0, 0);
dst[i] = z;
}
}
__global__ void add_bf16_kernel(const __nv_bfloat16* __restrict__ src,
float* __restrict__ dst, int n, int accumulate) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= n) return;
float v = __bfloat162float(src[i]);
dst[i] = accumulate ? (dst[i] + v) : v;
}
__global__ void cast_fp32_bf16_kernel(const float* __restrict__ src,
__nv_bfloat16* __restrict__ dst, int n) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < n) dst[i] = __float2bfloat16(src[i]);
}
__global__ void silu_mul_kernel(const __nv_bfloat16* __restrict__ gate_up,
__nv_bfloat16* __restrict__ hidden, int rows,
int I) {
int row = blockIdx.x;
if (row >= rows) return;
const __nv_bfloat16* g = gate_up + (int64_t)row * (I * 2);
const __nv_bfloat16* u = g + I;
__nv_bfloat16* h = hidden + (int64_t)row * I;
for (int i = threadIdx.x; i < I; i += blockDim.x) {
float hv = silu_f(__bfloat162float(g[i])) * __bfloat162float(u[i]);
h[i] = __float2bfloat16(hv);
}
}
// out[t] = shared[t] + sum_k weight[t,k] * Y[row_of_pair[t,k]]
template <int TOPK>
__global__ void reduce_kernel(const __nv_bfloat16* __restrict__ shared,
const __nv_bfloat16* __restrict__ Y,
const int* __restrict__ row_of_pair,
const __nv_bfloat16* __restrict__ weights,
__nv_bfloat16* __restrict__ out, int T, int H,
int top_k) {
int t = blockIdx.x;
if (t >= T) return;
const int kmax = TOPK > 0 ? TOPK : top_k;
for (int h = threadIdx.x; h < H; h += blockDim.x) {
float acc = __bfloat162float(shared[(int64_t)t * H + h]);
const int* rows = row_of_pair + (int64_t)t * kmax;
const __nv_bfloat16* w = weights + (int64_t)t * kmax;
#pragma unroll 1
for (int k = 0; k < kmax; ++k) {
int row = rows[k];
float yv = __bfloat162float(Y[(int64_t)row * H + h]);
acc = fmaf(__bfloat162float(w[k]), yv, acc);
}
out[(int64_t)t * H + h] = __float2bfloat16(acc);
}
}
// T=1 GEMV. One warp = one output. Streaming weight loads; x stays in L2.
// grid: (ceil(N / warps_per_block), n_exp)
constexpr int GEMV_THREADS = 128;
__global__ void gemv_nt_kernel(const __nv_bfloat16* __restrict__ x_base,
int64_t x_stride,
const __nv_bfloat16* __restrict__ W_base,
const int* __restrict__ expert_index,
int64_t expert_stride,
__nv_bfloat16* __restrict__ y, int N, int K) {
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
const int warps = blockDim.x >> 5;
const int n_out = blockIdx.x * warps + warp;
const int exp = blockIdx.y;
if (n_out >= N) return;
const __nv_bfloat16* x = x_base + (int64_t)exp * x_stride;
int e = expert_index[exp];
const __nv_bfloat16* row =
W_base + ((int64_t)e * expert_stride) + (int64_t)n_out * K;
const uint4* row4 = reinterpret_cast<const uint4*>(row);
const uint4* x4 = reinterpret_cast<const uint4*>(x);
const int nvec = K >> 3;
float a0 = 0.f, a1 = 0.f, a2 = 0.f, a3 = 0.f;
int i = lane;
for (; i + 96 < nvec; i += 128) {
uint4 w0 = ldcs_u4(row4 + i);
uint4 w1 = ldcs_u4(row4 + i + 32);
uint4 w2 = ldcs_u4(row4 + i + 64);
uint4 w3 = ldcs_u4(row4 + i + 96);
a0 += dot8(w0, x4[i]);
a1 += dot8(w1, x4[i + 32]);
a2 += dot8(w2, x4[i + 64]);
a3 += dot8(w3, x4[i + 96]);
}
float acc = (a0 + a1) + (a2 + a3);
for (; i < nvec; i += 32) acc += dot8(ldcs_u4(row4 + i), x4[i]);
#pragma unroll
for (int off = 16; off > 0; off >>= 1)
acc += __shfl_down_sync(0xffffffff, acc, off);
if (lane == 0) y[(int64_t)exp * N + n_out] = __float2bfloat16(acc);
}
__global__ void reduce_t1_kernel(const __nv_bfloat16* __restrict__ y_routed,
const __nv_bfloat16* __restrict__ y_shared,
const __nv_bfloat16* __restrict__ weights,
__nv_bfloat16* __restrict__ out, int top_k,
int n_shared, int H) {
for (int h = blockIdx.x * blockDim.x + threadIdx.x; h < H;
h += blockDim.x * gridDim.x) {
float acc = 0.f;
for (int s = 0; s < n_shared; ++s) {
acc += __bfloat162float(y_shared[(int64_t)s * H + h]);
}
for (int k = 0; k < top_k; ++k) {
float yv = __bfloat162float(y_routed[(int64_t)k * H + h]);
acc = fmaf(__bfloat162float(weights[k]), yv, acc);
}
out[h] = __float2bfloat16(acc);
}
}
// C = X @ W^T
// W: (batch, N, K) bf16, X: (batch, M, K) bf16, C: (batch, M, N) bf16
static void gemm_x_wt(int M, int N, int K, int batch, const void* W,
const void* X, void* C) {
float alpha = 1.f, beta = 0.f;
CHECK_CUBLAS(cublasGemmStridedBatchedEx(
ctx().handle, CUBLAS_OP_T, CUBLAS_OP_N, N, M, K, &alpha, W, CUDA_R_16BF,
K, (long long)N * K, X, CUDA_R_16BF, K, (long long)M * K, &beta, C,
CUDA_R_16BF, N, (long long)M * N, batch, CUBLAS_COMPUTE_32F,
CUBLAS_GEMM_DEFAULT));
}
static void launch_gemv(const __nv_bfloat16* x_base, int64_t x_stride,
const __nv_bfloat16* W_base, const int* expert_index,
int64_t expert_stride, __nv_bfloat16* y, int n_exp,
int N, int K, cudaStream_t stream) {
constexpr int warps = GEMV_THREADS / 32;
dim3 grid((N + warps - 1) / warps, n_exp);
gemv_nt_kernel<<<grid, GEMV_THREADS, 0, stream>>>(
x_base, x_stride, W_base, expert_index, expert_stride, y, N, K);
}
static torch::Tensor forward_t1(torch::Tensor x, torch::Tensor expert_ids,
torch::Tensor expert_weights,
torch::Tensor w1_routed, torch::Tensor w2_routed,
torch::Tensor w1_shared, torch::Tensor w2_shared) {
const int H = (int)x.size(1);
const int top_k = (int)expert_ids.size(1);
const int E = (int)w1_routed.size(0);
const int twoI = (int)w1_routed.size(1);
const int I = twoI / 2;
const int n_shared = (int)w1_shared.size(0);
TORCH_CHECK(H % 8 == 0 && I % 8 == 0, "H and I must be multiples of 8");
TORCH_CHECK(top_k > 0 && n_shared >= 0, "bad top_k/n_shared");
auto opt_i = torch::TensorOptions().dtype(torch::kInt32).device(x.device());
auto opt_b = x.options();
auto exp_idx = torch::empty({top_k + std::max(n_shared, 1)}, opt_i);
// ids are 8 int64s; pull them so the GEMV knows which rows of W to read.
auto ids_cpu = expert_ids.contiguous().view({-1}).to(torch::kCPU);
const int64_t* ids = ids_cpu.data_ptr<int64_t>();
std::vector<int> hidx(top_k + n_shared);
for (int k = 0; k < top_k; ++k) {
int e = (int)ids[k];
TORCH_CHECK(e >= 0 && e < E, "expert id out of range");
hidx[k] = e;
}
for (int s = 0; s < n_shared; ++s) hidx[top_k + s] = s;
cudaMemcpyAsync(exp_idx.data_ptr<int>(), hidx.data(),
(size_t)(top_k + n_shared) * sizeof(int),
cudaMemcpyHostToDevice, at::cuda::getCurrentCUDAStream());
auto y_routed = torch::empty({top_k, H}, opt_b);
auto y_shared = torch::empty({std::max(n_shared, 1), H}, opt_b);
auto gate_r = torch::empty({top_k, twoI}, opt_b);
auto hid_r = torch::empty({top_k, I}, opt_b);
auto gate_s = torch::empty({std::max(n_shared, 1), twoI}, opt_b);
auto hid_s = torch::empty({std::max(n_shared, 1), I}, opt_b);
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
const __nv_bfloat16* xptr =
reinterpret_cast<const __nv_bfloat16*>(x.data_ptr());
if (top_k > 0) {
launch_gemv(xptr, /*x_stride=*/0,
reinterpret_cast<const __nv_bfloat16*>(w1_routed.data_ptr()),
exp_idx.data_ptr<int>(), (int64_t)twoI * H,
reinterpret_cast<__nv_bfloat16*>(gate_r.data_ptr()), top_k,
twoI, H, stream);
silu_mul_kernel<<<top_k, 256, 0, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(gate_r.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(hid_r.data_ptr()), top_k, I);
launch_gemv(reinterpret_cast<const __nv_bfloat16*>(hid_r.data_ptr()),
/*x_stride=*/I,
reinterpret_cast<const __nv_bfloat16*>(w2_routed.data_ptr()),
exp_idx.data_ptr<int>(), (int64_t)H * I,
reinterpret_cast<__nv_bfloat16*>(y_routed.data_ptr()), top_k,
H, I, stream);
}
if (n_shared > 0) {
launch_gemv(xptr, 0,
reinterpret_cast<const __nv_bfloat16*>(w1_shared.data_ptr()),
exp_idx.data_ptr<int>() + top_k, (int64_t)twoI * H,
reinterpret_cast<__nv_bfloat16*>(gate_s.data_ptr()), n_shared,
twoI, H, stream);
silu_mul_kernel<<<n_shared, 256, 0, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(gate_s.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(hid_s.data_ptr()), n_shared, I);
// shared experts are indexed 0..n_shared-1 inside w1_shared, which we
// stored at exp_idx[top_k + s].
launch_gemv(reinterpret_cast<const __nv_bfloat16*>(hid_s.data_ptr()), I,
reinterpret_cast<const __nv_bfloat16*>(w2_shared.data_ptr()),
exp_idx.data_ptr<int>() + top_k, (int64_t)H * I,
reinterpret_cast<__nv_bfloat16*>(y_shared.data_ptr()), n_shared,
H, I, stream);
}
auto out = torch::empty({1, H}, opt_b);
reduce_t1_kernel<<<(H + 255) / 256, 256, 0, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(y_routed.data_ptr()),
reinterpret_cast<const __nv_bfloat16*>(y_shared.data_ptr()),
reinterpret_cast<const __nv_bfloat16*>(expert_weights.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), top_k, n_shared, H);
return out;
}
static torch::Tensor forward_grouped(torch::Tensor x, torch::Tensor expert_ids,
torch::Tensor expert_weights,
torch::Tensor w1_routed,
torch::Tensor w2_routed,
torch::Tensor w1_shared,
torch::Tensor w2_shared) {
ensure_handle();
const int T = (int)x.size(0);
const int H = (int)x.size(1);
const int top_k = (int)expert_ids.size(1);
const int E = (int)w1_routed.size(0);
const int twoI = (int)w1_routed.size(1);
const int I = twoI / 2;
const int n_shared = (int)w1_shared.size(0);
const int n_pairs = T * top_k;
TORCH_CHECK(H % 8 == 0 && I % 8 == 0, "H and I must be multiples of 8");
TORCH_CHECK(w1_routed.size(2) == H && w2_routed.size(1) == H, "H mismatch");
TORCH_CHECK(w2_routed.size(2) == I && w2_routed.size(0) == E, "I/E mismatch");
Ctx& c = ctx();
auto dev = x.device();
if (c.cap_E < E) {
c.counts_h = torch::empty(
{E}, torch::TensorOptions().dtype(torch::kInt32).pinned_memory(true));
c.counts_d = torch::empty(
{E}, torch::TensorOptions().dtype(torch::kInt32).device(dev));
c.cursor_d = torch::empty(
{E}, torch::TensorOptions().dtype(torch::kInt32).device(dev));
c.cap_E = E;
}
if (c.cap_pairs < n_pairs) {
c.row_of_pair = torch::empty(
{n_pairs}, torch::TensorOptions().dtype(torch::kInt32).device(dev));
c.cap_pairs = n_pairs;
}
if (c.cap_T < T || c.cap_H < H || c.cap_I < I) {
auto bf = x.options();
c.shared_gate = torch::empty({(int64_t)T * twoI}, bf);
c.shared_h = torch::empty({(int64_t)T * I}, bf);
c.shared_y = torch::empty({(int64_t)T * H}, bf);
c.cap_T = T;
c.cap_H = H;
c.cap_I = I;
}
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
cudaMemsetAsync(c.counts_d.data_ptr(), 0, (size_t)E * sizeof(int), stream);
histogram_i64<<<(n_pairs + 255) / 256, 256, 0, stream>>>(
expert_ids.data_ptr<int64_t>(), n_pairs, E, c.counts_d.data_ptr<int>());
cudaMemcpyAsync(c.counts_h.data_ptr<int>(), c.counts_d.data_ptr<int>(),
(size_t)E * sizeof(int), cudaMemcpyDeviceToHost, stream);
cudaStreamSynchronize(stream);
const int* hc = c.counts_h.data_ptr<int>();
int maxc = 0;
for (int e = 0; e < E; ++e) maxc = std::max(maxc, hc[e]);
int M_pad = choose_mpad(std::max(maxc, 1));
// Always batch all E experts. Empty experts have zero rows of x and are
// never scattered into the output. At serving T this set is full anyway;
// skipping the rare empty expert would force a weight compact copy.
int64_t rows = (int64_t)E * M_pad;
if (c.cap_rows < rows || c.cap_H < H || c.cap_I < I) {
auto bf = x.options();
c.Xg = torch::empty({rows, H}, bf);
c.gate_up = torch::empty({rows, twoI}, bf);
c.hidden = torch::empty({rows, I}, bf);
c.Y = torch::empty({rows, H}, bf);
c.cap_rows = rows;
}
// Zero only the padded tail of each expert (not the rows gather overwrites).
int pad_span = 0;
for (int e = 0; e < E; ++e) pad_span = std::max(pad_span, M_pad - hc[e]);
if (pad_span > 0) {
// counts_d still holds the real counts; cursor is separate.
zero_pad_rows<<<dim3(pad_span, E), 128, 0, stream>>>(
c.counts_d.data_ptr<int>(),
reinterpret_cast<__nv_bfloat16*>(c.Xg.data_ptr()), E, M_pad, H);
}
cudaMemsetAsync(c.cursor_d.data_ptr(), 0, (size_t)E * sizeof(int), stream);
assign_and_gather<<<n_pairs, 128, 0, stream>>>(
expert_ids.data_ptr<int64_t>(),
reinterpret_cast<const __nv_bfloat16*>(x.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(c.Xg.data_ptr()),
c.cursor_d.data_ptr<int>(), c.row_of_pair.data_ptr<int>(), n_pairs, top_k,
H, M_pad);
// Routed w1: (E, M_pad, H) @ (E, 2I, H)^T -> (E, M_pad, 2I)
gemm_x_wt(M_pad, twoI, H, E, w1_routed.data_ptr(), c.Xg.data_ptr(),
c.gate_up.data_ptr());
silu_mul_kernel<<<(int)rows, 256, 0, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(c.gate_up.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(c.hidden.data_ptr()), (int)rows, I);
gemm_x_wt(M_pad, H, I, E, w2_routed.data_ptr(), c.hidden.data_ptr(),
c.Y.data_ptr());
// Shared experts. n_shared==1 writes bf16 directly; more than one sums in fp32.
torch::Tensor shared_bf;
if (n_shared <= 0) {
shared_bf = torch::zeros({T, H}, x.options());
} else if (n_shared == 1) {
const __nv_bfloat16* w1s =
reinterpret_cast<const __nv_bfloat16*>(w1_shared.data_ptr());
const __nv_bfloat16* w2s =
reinterpret_cast<const __nv_bfloat16*>(w2_shared.data_ptr());
gemm_x_wt(T, twoI, H, 1, w1s, x.data_ptr(), c.shared_gate.data_ptr());
silu_mul_kernel<<<T, 256, 0, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(c.shared_gate.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(c.shared_h.data_ptr()), T, I);
gemm_x_wt(T, H, I, 1, w2s, c.shared_h.data_ptr(), c.shared_y.data_ptr());
// shared_y is cached at the high-water T; only the first T rows are live.
shared_bf = c.shared_y.narrow(0, 0, (int64_t)T * H).view({T, H});
} else {
auto shared_fp = torch::empty({T, H}, x.options().dtype(torch::kFloat32));
const int n_out = T * H;
for (int s = 0; s < n_shared; ++s) {
const __nv_bfloat16* w1s =
reinterpret_cast<const __nv_bfloat16*>(w1_shared.data_ptr()) +
(int64_t)s * twoI * H;
const __nv_bfloat16* w2s =
reinterpret_cast<const __nv_bfloat16*>(w2_shared.data_ptr()) +
(int64_t)s * H * I;
gemm_x_wt(T, twoI, H, 1, w1s, x.data_ptr(), c.shared_gate.data_ptr());
silu_mul_kernel<<<T, 256, 0, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(c.shared_gate.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(c.shared_h.data_ptr()), T, I);
gemm_x_wt(T, H, I, 1, w2s, c.shared_h.data_ptr(), c.shared_y.data_ptr());
add_bf16_kernel<<<(n_out + 255) / 256, 256, 0, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(c.shared_y.data_ptr()),
shared_fp.data_ptr<float>(), n_out, s > 0);
}
shared_bf = torch::empty({T, H}, x.options());
cast_fp32_bf16_kernel<<<(n_out + 255) / 256, 256, 0, stream>>>(
shared_fp.data_ptr<float>(),
reinterpret_cast<__nv_bfloat16*>(shared_bf.data_ptr()), n_out);
}
auto out = torch::empty({T, H}, x.options());
auto shared_c = shared_bf.contiguous();
if (top_k == 8) {
reduce_kernel<8><<<T, 256, 0, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(shared_c.data_ptr()),
reinterpret_cast<const __nv_bfloat16*>(c.Y.data_ptr()),
c.row_of_pair.data_ptr<int>(),
reinterpret_cast<const __nv_bfloat16*>(expert_weights.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), T, H, top_k);
} else {
reduce_kernel<0><<<T, 256, 0, stream>>>(
reinterpret_cast<const __nv_bfloat16*>(shared_c.data_ptr()),
reinterpret_cast<const __nv_bfloat16*>(c.Y.data_ptr()),
c.row_of_pair.data_ptr<int>(),
reinterpret_cast<const __nv_bfloat16*>(expert_weights.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), T, H, top_k);
}
return out;
}
torch::Tensor fused_moe(torch::Tensor x, torch::Tensor expert_ids,
torch::Tensor expert_weights, torch::Tensor w1_routed,
torch::Tensor w2_routed, torch::Tensor w1_shared,
torch::Tensor w2_shared) {
TORCH_CHECK(x.is_cuda(), "x must be CUDA");
x = x.contiguous();
expert_ids = expert_ids.contiguous();
expert_weights = expert_weights.contiguous();
w1_routed = w1_routed.contiguous();
w2_routed = w2_routed.contiguous();
w1_shared = w1_shared.contiguous();
w2_shared = w2_shared.contiguous();
if (expert_ids.scalar_type() != torch::kLong) {
expert_ids = expert_ids.to(torch::kLong);
}
if (x.scalar_type() != torch::kBFloat16) {
x = x.to(torch::kBFloat16);
}
TORCH_CHECK(x.dim() == 2, "x is (T, H)");
const int T = (int)x.size(0);
if (T == 0) {
return torch::empty({0, x.size(1)}, x.options());
}
ensure_handle();
if (T == 1) {
return forward_t1(x, expert_ids, expert_weights, w1_routed, w2_routed,
w1_shared, w2_shared);
}
return forward_grouped(x, expert_ids, expert_weights, w1_routed, w2_routed,
w1_shared, w2_shared);
}
"""
_CPP = r"""
#include <torch/extension.h>
torch::Tensor fused_moe(torch::Tensor x, torch::Tensor expert_ids,
torch::Tensor expert_weights, torch::Tensor w1_routed,
torch::Tensor w2_routed, torch::Tensor w1_shared,
torch::Tensor w2_shared);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("fused_moe", &fused_moe, "GLM-5.2 fused MoE");
}
"""
_EXT = None
def _ext():
global _EXT
if _EXT is None:
_EXT = load_inline(
name="glm52_fused_moe_v1",
cpp_sources=[_CPP],
cuda_sources=[_CUDA],
extra_ldflags=["-lcublas"],
extra_cuda_cflags=["-O3", "-std=c++20"],
extra_cflags=["-O3"],
verbose=False,
)
return _EXT
class Model(nn.Module):
def __init__(self, T: int, E: int, top_k: int, n_shared: int, H: int, I: int):
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().fused_moe(
x, expert_ids, expert_weights,
self.w1_routed, self.w2_routed, self.w1_shared, self.w2_shared,
)
20260917_005805_grok_grok-4.7_01_glm52_fused_moe