KernelBench cuda · RTX PRO 6000
GLM-5.2 Fused MoE Grok 4.6
9.39%geomean peak fraction across shapes
Real fused-MoE plumbing in moe_kernels.cu (gather, SiLU*mul, scatter-add, plus an mma.sync grouped GEMM). Default timed GEMM backend is cuBLAS (`GLM52_MOE_GEMM` defaults to "cublas"; MMA is the fallback). Same class as published grok-4.5 / kinetic cells: cuBLAS is not on the forbidden list.
harnessgrokagent session45mtotal wall56mcheck4mbenchmark4moutput tokens—regimecompute
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
shape 010.534 ms35.2%1.23 TB/s · 69% of 1.8 TB/s HBM · also 176 TFLOPS (35% of compute)
shape 110.608 ms35.3%1.23 TB/s · 68% of 1.8 TB/s HBM · also 176 TFLOPS (35% of compute)
shape 20.388 ms0.2%33.33 TB/s · 1852% of 1.8 TB/s HBM · also 1 TFLOPS (0% of compute)
shape 318.547 ms40.0%200 TFLOPS · 40% of 500 TF bf16 peak · also 0.70 TB/s (39% of HBM)
shape 48.291 ms5.6%1.56 TB/s · 87% of 1.8 TB/s HBM · also 28 TFLOPS (6% of compute)
shape 58.564 ms10.6%1.51 TB/s · 84% of 1.8 TB/s HBM · also 53 TFLOPS (11% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(35.2% · 35.3% · 0.2% · 40.0% · 5.6% · 10.6%) = 9.4%
Kernel source (redacted)
"""GLM-5.2 fused MoE: grouped BF16 GEMM + SiLU*mul + weighted reduce.
CUDA kernels live in moe_kernels.cu (mma.sync grouped GEMM, cuBLAS grouped GEMM,
SiLU-and-mul, gather, scatter-add). Routing is provided.
"""
from __future__ import annotations
import os
from pathlib import Path
import torch
import torch.nn as nn
from torch.utils.cpp_extension import load
os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0")
_SRC = Path(__file__).resolve().parent / "moe_kernels.cu"
_BUILD = Path(__file__).resolve().parent / ".moe_build"
_BUILD.mkdir(exist_ok=True)
_ext = load(
name="glm52_moe_kernels",
sources=[str(_SRC)],
extra_cuda_cflags=[
"-O3",
"--use_fast_math",
"-std=c++17",
"-U__CUDA_NO_BFLOAT16_OPERATORS__",
"-U__CUDA_NO_BFLOAT16_CONVERSIONS__",
"--expt-relaxed-constexpr",
],
extra_ldflags=["-lcublas"],
build_directory=str(_BUILD),
verbose=False,
)
# "cublas" is typically faster for medium/large expert M; WMMA is the
# portable Tensor-Core path and the correctness fallback.
_GEMM_BACKEND = os.environ.get("GLM52_MOE_GEMM", "cublas").lower()
class _Buf:
def __init__(self) -> None:
self._t: dict[str, torch.Tensor] = {}
def get(self, key: str, shape: tuple[int, ...], dtype: torch.dtype,
device: torch.device) -> torch.Tensor:
n = 1
for s in shape:
n *= int(s)
t = self._t.get(key)
if t is None or t.numel() < n or t.dtype != dtype or t.device != device:
t = torch.empty(n, dtype=dtype, device=device)
self._t[key] = t
return t[:n].view(shape)
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)
self._buf = _Buf()
self._backend = _GEMM_BACKEND
self._streams = None
def _grouped_gemm(
self,
A: torch.Tensor,
W: torch.Tensor,
counts: torch.Tensor,
counts_cpu: torch.Tensor,
offsets: torch.Tensor,
tile_prefix: torch.Tensor,
n_tiles: int,
out_n: int,
) -> torch.Tensor:
total_m = A.size(0)
C = self._buf.get(f"C_{out_n}", (total_m, out_n), torch.bfloat16, A.device)
if self._backend == "mma":
_ext.grouped_gemm_mma(A, W, C, offsets, counts, tile_prefix, n_tiles)
else:
_ext.grouped_gemm_cublas(A, W, C, counts_cpu)
return C
def forward(
self,
x: torch.Tensor,
expert_ids: torch.Tensor,
expert_weights: torch.Tensor,
) -> torch.Tensor:
T, H = x.shape
E, top_k, I = self.E, self.top_k, self.I
device = x.device
x = x.contiguous()
expert_ids = expert_ids.contiguous()
expert_weights = expert_weights.contiguous()
out_f = self._buf.get("out_f", (T, H), torch.float32, device)
out_f.zero_()
# Shared experts: always-on, one large GEMM pair each.
for s in range(self.n_shared):
gu = self._buf.get("s_gu", (T, 2 * I), torch.bfloat16, device)
h = self._buf.get("s_h", (T, I), torch.bfloat16, device)
ys = self._buf.get("s_y", (T, H), torch.bfloat16, device)
_ext.gemm_nt(x, self.w1_shared[s], gu)
_ext.silu_and_mul(gu, h)
_ext.gemm_nt(h, self.w2_shared[s], ys)
out_f.add_(ys.float())
# Decode: per-assignment GEMVs on a few streams (better L2 than grouped M=1).
if T <= 2:
ids_host = expert_ids.view(-1).tolist()
wts_host = expert_weights.reshape(-1).float().tolist()
n_assign = T * top_k
if self._streams is None:
self._streams = [torch.cuda.Stream() for _ in range(4)]
gu_s = self._buf.get("d_gu", (n_assign, 2 * I), torch.bfloat16, device)
h_s = self._buf.get("d_h", (n_assign, I), torch.bfloat16, device)
y_s = self._buf.get("d_y", (n_assign, H), torch.bfloat16, device)
main = torch.cuda.current_stream()
for a in range(n_assign):
e = int(ids_host[a])
tok = a // top_k
st = self._streams[a % len(self._streams)]
st.wait_stream(main)
with torch.cuda.stream(st):
_ext.gemm_nt(x[tok : tok + 1], self.w1_routed[e], gu_s[a : a + 1])
_ext.silu_and_mul(gu_s[a : a + 1], h_s[a : a + 1])
_ext.gemm_nt(h_s[a : a + 1], self.w2_routed[e], y_s[a : a + 1])
for st in self._streams:
main.wait_stream(st)
for a in range(n_assign):
tok = a // top_k
out_f[tok].add_(y_s[a].float(), alpha=float(wts_host[a]))
return out_f.to(torch.bfloat16)
# Prefill: sort tokens by expert, grouped GEMM, SiLU, down, scatter.
n_assign = T * top_k
flat_e = expert_ids.reshape(-1)
token_ids = torch.arange(T, device=device, dtype=torch.int32).repeat_interleave(
top_k
)
perm = torch.argsort(flat_e)
sorted_e = flat_e.index_select(0, perm).to(torch.int32)
sorted_tokens = token_ids.index_select(0, perm)
sorted_w = expert_weights.reshape(-1).index_select(0, perm)
counts = torch.bincount(sorted_e, minlength=E).to(torch.int32)
counts_cpu = counts.cpu()
offsets = self._buf.get("off", (E + 1,), torch.int32, device)
offsets.zero_()
offsets[1:].copy_(torch.cumsum(counts, 0))
gathered = self._buf.get("gx", (n_assign, H), torch.bfloat16, device)
_ext.gather_rows(x, sorted_tokens, gathered)
if self._backend == "mma":
BM = _ext.tile_size_m()
BN = _ext.tile_size_n()
m_tiles = (counts + (BM - 1)) // BM
pref1 = self._buf.get("tp1", (E + 1,), torch.int32, device)
pref2 = self._buf.get("tp2", (E + 1,), torch.int32, device)
pref1.zero_()
pref2.zero_()
pref1[1:].copy_(torch.cumsum(m_tiles * ((2 * I + BN - 1) // BN), 0))
pref2[1:].copy_(torch.cumsum(m_tiles * ((H + BN - 1) // BN), 0))
n1 = int(pref1[-1].item())
n2 = int(pref2[-1].item())
else:
pref1 = offsets
pref2 = offsets
n1 = n2 = 0
gu_r = self._grouped_gemm(
gathered, self.w1_routed, counts, counts_cpu, offsets, pref1, n1, 2 * I
)
h_r = self._buf.get("h_r", (n_assign, I), torch.bfloat16, device)
_ext.silu_and_mul(gu_r, h_r)
y_r = self._grouped_gemm(
h_r, self.w2_routed, counts, counts_cpu, offsets, pref2, n2, H
)
_ext.scatter_add(y_r, sorted_tokens, sorted_w, out_f)
return out_f.to(torch.bfloat16)
# ==================================================================
# ===== sidecar: moe_kernels.cu (19892 bytes, loaded by solution.py) =====
# ==================================================================
// GLM-5.2 fused MoE: grouped BF16 GEMM (cuBLAS + MMA), SiLU*mul, gather, scatter.
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAStream.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>
#include <torch/extension.h>
#include <algorithm>
#include <cstdint>
#include <stdexcept>
#include <vector>
using bf16 = __nv_bfloat16;
#define CUDA_CHECK(expr) \
do { \
cudaError_t err = (expr); \
if (err != cudaSuccess) { \
throw std::runtime_error(std::string("CUDA error: ") + \
cudaGetErrorString(err)); \
} \
} while (0)
#define CUBLAS_CHECK(expr) \
do { \
cublasStatus_t st = (expr); \
if (st != CUBLAS_STATUS_SUCCESS) { \
throw std::runtime_error("cuBLAS error status " + std::to_string((int)st)); \
} \
} while (0)
// ---------------------------------------------------------------------------
// SiLU(gate) * up. gate_up is (rows, 2*I) packed [gate | up].
// ---------------------------------------------------------------------------
__global__ void silu_and_mul_kernel(const bf16* __restrict__ gu,
bf16* __restrict__ out, int rows, int I) {
const int64_t total = (int64_t)rows * I;
for (int64_t idx = (int64_t)blockIdx.x * blockDim.x + threadIdx.x; idx < total;
idx += (int64_t)blockDim.x * gridDim.x) {
const int row = (int)(idx / I);
const int col = (int)(idx - (int64_t)row * I);
const bf16* rowp = gu + (int64_t)row * (I + I);
const float g = __bfloat162float(rowp[col]);
const float u = __bfloat162float(rowp[I + col]);
const float s = g * (1.f / (1.f + expf(-g)));
out[idx] = __float2bfloat16_rn(s * u);
}
}
void silu_and_mul(torch::Tensor gate_up, torch::Tensor out) {
TORCH_CHECK(gate_up.is_cuda() && out.is_cuda(), "cuda tensors");
const int rows = (int)gate_up.size(0);
const int I = (int)gate_up.size(1) / 2;
if (rows == 0) return;
const int threads = 256;
const int64_t total = (int64_t)rows * I;
int blocks = (int)std::min<int64_t>((total + threads - 1) / threads, 2048);
auto stream = at::cuda::getCurrentCUDAStream();
silu_and_mul_kernel<<<blocks, threads, 0, stream>>>(
reinterpret_cast<const bf16*>(gate_up.data_ptr()),
reinterpret_cast<bf16*>(out.data_ptr()), rows, I);
}
// ---------------------------------------------------------------------------
// Weighted scatter-add: out[token[i], :] += w[i] * y[i, :] (out fp32)
// ---------------------------------------------------------------------------
__global__ void scatter_add_kernel(const bf16* __restrict__ y,
const int32_t* __restrict__ tokens,
const bf16* __restrict__ weights,
float* __restrict__ out, int n, int H) {
const int i = blockIdx.x;
if (i >= n) return;
const int t = tokens[i];
const float w = __bfloat162float(weights[i]);
const bf16* yi = y + (int64_t)i * H;
float* o = out + (int64_t)t * H;
for (int h = threadIdx.x; h < H; h += blockDim.x) {
atomicAdd(o + h, w * __bfloat162float(yi[h]));
}
}
void scatter_add(torch::Tensor y, torch::Tensor tokens, torch::Tensor weights,
torch::Tensor out) {
const int n = (int)y.size(0);
const int H = (int)y.size(1);
if (n == 0) return;
auto stream = at::cuda::getCurrentCUDAStream();
scatter_add_kernel<<<n, 256, 0, stream>>>(
reinterpret_cast<const bf16*>(y.data_ptr()), tokens.data_ptr<int32_t>(),
reinterpret_cast<const bf16*>(weights.data_ptr()), out.data_ptr<float>(),
n, H);
}
// ---------------------------------------------------------------------------
// Gather rows of x by int32 indices (16-byte vectorized).
// ---------------------------------------------------------------------------
__global__ void gather_rows_kernel(const bf16* __restrict__ x,
const int32_t* __restrict__ idx,
bf16* __restrict__ out, int n, int H) {
const int i = blockIdx.x;
if (i >= n) return;
const bf16* src = x + (int64_t)idx[i] * H;
bf16* dst = out + (int64_t)i * H;
const int vec = H >> 3;
const int4* src4 = reinterpret_cast<const int4*>(src);
int4* dst4 = reinterpret_cast<int4*>(dst);
for (int j = threadIdx.x; j < vec; j += blockDim.x) dst4[j] = src4[j];
}
void gather_rows(torch::Tensor x, torch::Tensor idx, torch::Tensor out) {
const int n = (int)idx.size(0);
const int H = (int)x.size(1);
if (n == 0) return;
auto stream = at::cuda::getCurrentCUDAStream();
gather_rows_kernel<<<n, 128, 0, stream>>>(
reinterpret_cast<const bf16*>(x.data_ptr()), idx.data_ptr<int32_t>(),
reinterpret_cast<bf16*>(out.data_ptr()), n, H);
}
// ---------------------------------------------------------------------------
// MMA helpers: C += A @ B with A row-major 16x16, B col-major 16x8
// (B col-major KxN == W row-major NxK).
// ---------------------------------------------------------------------------
__device__ __forceinline__ void mma_m16n8k16(float& d0, float& d1, float& d2,
float& d3, unsigned a0, unsigned a1,
unsigned a2, unsigned a3,
unsigned b0, unsigned 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"(d0), "+f"(d1), "+f"(d2), "+f"(d3)
: "r"(a0), "r"(a1), "r"(a2), "r"(a3), "r"(b0), "r"(b1));
}
// Tile: 64x64, 4 warps (2x2), each warp 32x32 = 2x4 m16n8.
constexpr int BM = 64;
constexpr int BN = 64;
constexpr int BK = 16;
constexpr int THREADS = 128;
__device__ __forceinline__ int find_expert(const int* __restrict__ prefix, int E,
int tile) {
int lo = 0, hi = E;
while (lo + 1 < hi) {
const int mid = (lo + hi) >> 1;
if (prefix[mid] <= tile)
lo = mid;
else
hi = mid;
}
return lo;
}
__global__ void grouped_gemm_mma_kernel(const bf16* __restrict__ A,
const bf16* __restrict__ W,
bf16* __restrict__ C,
const int* __restrict__ offsets,
const int* __restrict__ counts,
const int* __restrict__ tile_prefix,
int E, int N, int K) {
const int tile = (int)blockIdx.x;
const int e = find_expert(tile_prefix, E, tile);
const int local = tile - tile_prefix[e];
const int n_ntiles = (N + BN - 1) / BN;
const int tm = local / n_ntiles;
const int tn = local - tm * n_ntiles;
const int m0 = tm * BM;
const int n0 = tn * BN;
const int M = counts[e];
if (m0 >= M) return;
const int row0 = offsets[e];
const int Mtile = min(BM, M - m0);
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int warp_m = warp >> 1; // 0..1
const int warp_n = warp & 1; // 0..1
const int group = lane >> 2; // 0..7
const int tidg = lane & 3; // 0..3
extern __shared__ bf16 smem[];
bf16* As = smem; // [BM][BK]
bf16* Bs = smem + BM * BK; // [BN][BK]
// Each warp holds 2x4 = 8 accum fragments of 4 floats (32x32).
float acc[2][4][4];
#pragma unroll
for (int i = 0; i < 2; i++)
#pragma unroll
for (int j = 0; j < 4; j++)
#pragma unroll
for (int r = 0; r < 4; r++) acc[i][j][r] = 0.f;
const bf16* A_base = A + (int64_t)(row0 + m0) * K;
const bf16* W_base = W + ((int64_t)e * N + n0) * K;
for (int k0 = 0; k0 < K; k0 += BK) {
for (int idx = threadIdx.x; idx < BM * BK; idx += THREADS) {
const int r = idx >> 4; // /BK
const int c = idx & 15;
bf16 v = __float2bfloat16(0.f);
if (r < Mtile) v = A_base[(int64_t)r * K + k0 + c];
As[idx] = v;
}
for (int idx = threadIdx.x; idx < BN * BK; idx += THREADS) {
const int r = idx >> 4;
const int c = idx & 15;
Bs[idx] = W_base[(int64_t)r * K + k0 + c];
}
__syncthreads();
// Load A fragments for this warp's 32 rows, B fragments for 32 cols.
#pragma unroll
for (int i = 0; i < 2; i++) {
const int row0a = warp_m * 32 + i * 16 + group;
const int colk = tidg * 2;
const bf16* a0p = As + row0a * BK + colk;
const bf16* a8p = As + (row0a + 8) * BK + colk;
unsigned a0 = *reinterpret_cast<const unsigned*>(a0p);
unsigned a1 = *reinterpret_cast<const unsigned*>(a8p);
unsigned a2 = *reinterpret_cast<const unsigned*>(a0p + 8);
unsigned a3 = *reinterpret_cast<const unsigned*>(a8p + 8);
#pragma unroll
for (int j = 0; j < 4; j++) {
const int ncol = warp_n * 32 + j * 8 + group;
const bf16* bp = Bs + ncol * BK + colk;
unsigned b0 = *reinterpret_cast<const unsigned*>(bp);
unsigned b1 = *reinterpret_cast<const unsigned*>(bp + 8);
mma_m16n8k16(acc[i][j][0], acc[i][j][1], acc[i][j][2], acc[i][j][3], a0,
a1, a2, a3, b0, b1);
}
}
__syncthreads();
}
// Store C fragments.
// c0 = C[group][tidg], c1 = C[group][tidg+4],
// c2 = C[group+8][tidg], c3 = C[group+8][tidg+4]
bf16* C_base = C + (int64_t)(row0 + m0) * N + n0;
#pragma unroll
for (int i = 0; i < 2; i++) {
#pragma unroll
for (int j = 0; j < 4; j++) {
const int rm = warp_m * 32 + i * 16;
const int cn = warp_n * 32 + j * 8;
const int r0 = rm + group;
const int r1 = rm + group + 8;
const int c0 = cn + tidg;
const int c1 = cn + tidg + 4;
if (r0 < Mtile) C_base[(int64_t)r0 * N + c0] = __float2bfloat16_rn(acc[i][j][0]);
if (r0 < Mtile) C_base[(int64_t)r0 * N + c1] = __float2bfloat16_rn(acc[i][j][1]);
if (r1 < Mtile) C_base[(int64_t)r1 * N + c0] = __float2bfloat16_rn(acc[i][j][2]);
if (r1 < Mtile) C_base[(int64_t)r1 * N + c1] = __float2bfloat16_rn(acc[i][j][3]);
}
}
}
void grouped_gemm_mma(torch::Tensor A, torch::Tensor W, torch::Tensor C,
torch::Tensor offsets, torch::Tensor counts,
torch::Tensor tile_prefix, int64_t n_tiles) {
const int E = (int)W.size(0);
const int N = (int)W.size(1);
const int K = (int)W.size(2);
if (n_tiles <= 0) return;
auto stream = at::cuda::getCurrentCUDAStream();
const size_t smem = sizeof(bf16) * (BM * BK + BN * BK);
grouped_gemm_mma_kernel<<<(int)n_tiles, THREADS, smem, stream>>>(
reinterpret_cast<const bf16*>(A.data_ptr()),
reinterpret_cast<const bf16*>(W.data_ptr()),
reinterpret_cast<bf16*>(C.data_ptr()), offsets.data_ptr<int32_t>(),
counts.data_ptr<int32_t>(), tile_prefix.data_ptr<int32_t>(), E, N, K);
}
// ---------------------------------------------------------------------------
// cuBLAS grouped GEMM. C = A @ W.T per expert.
// ---------------------------------------------------------------------------
static cublasHandle_t get_own_cublas() {
static cublasHandle_t h = nullptr;
if (!h) {
CUBLAS_CHECK(cublasCreate(&h));
CUBLAS_CHECK(cublasSetMathMode(h, CUBLAS_DEFAULT_MATH));
}
return h;
}
struct PtrScratch {
const void** dA = nullptr;
const void** dB = nullptr;
void** dC = nullptr;
int cap = 0;
void ensure(int n, cudaStream_t stream) {
if (n <= cap) return;
if (dA) cudaFree(dA);
if (dB) cudaFree(dB);
if (dC) cudaFree(dC);
cap = std::max(n, 256);
CUDA_CHECK(cudaMalloc(&dA, sizeof(void*) * cap));
CUDA_CHECK(cudaMalloc(&dB, sizeof(void*) * cap));
CUDA_CHECK(cudaMalloc(&dC, sizeof(void*) * cap));
(void)stream;
}
};
static PtrScratch& ptr_scratch() {
static PtrScratch s;
return s;
}
void grouped_gemm_cublas(torch::Tensor A, torch::Tensor W, torch::Tensor C,
torch::Tensor counts_cpu) {
TORCH_CHECK(counts_cpu.device().is_cpu(), "counts must be CPU");
const int E = (int)W.size(0);
const int N = (int)W.size(1);
const int K = (int)W.size(2);
const int32_t* cnt = counts_cpu.data_ptr<int32_t>();
std::vector<cublasOperation_t> ta, tb;
std::vector<int> m_arr, n_arr, k_arr, lda, ldb, ldc, gsz;
std::vector<const void*> Ap, Bp;
std::vector<void*> Cp;
std::vector<float> alpha, beta;
ta.reserve(E);
int off = 0;
const bf16* Aptr = reinterpret_cast<const bf16*>(A.data_ptr());
const bf16* Wptr = reinterpret_cast<const bf16*>(W.data_ptr());
bf16* Cptr = reinterpret_cast<bf16*>(C.data_ptr());
for (int e = 0; e < E; e++) {
const int Me = cnt[e];
if (Me <= 0) continue;
ta.push_back(CUBLAS_OP_T);
tb.push_back(CUBLAS_OP_N);
m_arr.push_back(N);
n_arr.push_back(Me);
k_arr.push_back(K);
lda.push_back(K);
ldb.push_back(K);
ldc.push_back(N);
gsz.push_back(1);
Ap.push_back(Wptr + (int64_t)e * N * K);
Bp.push_back(Aptr + (int64_t)off * K);
Cp.push_back(Cptr + (int64_t)off * N);
alpha.push_back(1.f);
beta.push_back(0.f);
off += Me;
}
const int group_count = (int)ta.size();
if (group_count == 0) return;
auto handle = get_own_cublas();
auto stream = at::cuda::getCurrentCUDAStream();
CUBLAS_CHECK(cublasSetStream(handle, stream));
auto& scratch = ptr_scratch();
scratch.ensure(group_count, stream);
CUDA_CHECK(cudaMemcpyAsync(scratch.dA, Ap.data(), sizeof(void*) * group_count,
cudaMemcpyHostToDevice, stream));
CUDA_CHECK(cudaMemcpyAsync(scratch.dB, Bp.data(), sizeof(void*) * group_count,
cudaMemcpyHostToDevice, stream));
CUDA_CHECK(cudaMemcpyAsync(scratch.dC, Cp.data(), sizeof(void*) * group_count,
cudaMemcpyHostToDevice, stream));
cublasStatus_t st = cublasGemmGroupedBatchedEx(
handle, ta.data(), tb.data(), m_arr.data(), n_arr.data(), k_arr.data(),
alpha.data(), reinterpret_cast<const void* const*>(scratch.dA),
CUDA_R_16BF, lda.data(),
reinterpret_cast<const void* const*>(scratch.dB), CUDA_R_16BF, ldb.data(),
beta.data(), reinterpret_cast<void* const*>(scratch.dC), CUDA_R_16BF,
ldc.data(), group_count, gsz.data(), CUBLAS_COMPUTE_32F);
if (st != CUBLAS_STATUS_SUCCESS) {
for (int i = 0; i < group_count; i++) {
CUBLAS_CHECK(cublasGemmEx(
handle, ta[i], tb[i], m_arr[i], n_arr[i], k_arr[i], &alpha[i], Ap[i],
CUDA_R_16BF, lda[i], Bp[i], CUDA_R_16BF, ldb[i], &beta[i], Cp[i],
CUDA_R_16BF, ldc[i], CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
}
}
}
void gemm_nt(torch::Tensor A, torch::Tensor W, torch::Tensor C) {
const int M = (int)A.size(0);
const int K = (int)A.size(1);
const int N = (int)W.size(0);
if (M == 0) return;
auto handle = get_own_cublas();
CUBLAS_CHECK(cublasSetStream(handle, at::cuda::getCurrentCUDAStream()));
const float alpha = 1.f, beta = 0.f;
CUBLAS_CHECK(cublasGemmEx(
handle, CUBLAS_OP_T, CUBLAS_OP_N, N, M, K, &alpha, W.data_ptr(),
CUDA_R_16BF, K, A.data_ptr(), CUDA_R_16BF, K, &beta, C.data_ptr(),
CUDA_R_16BF, N, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
}
// ---------------------------------------------------------------------------
// Small-T fused decode: one block per (token, expert) assignment.
// Streams that expert's w1/w2 once, writes weighted y into fp32 out.
// ---------------------------------------------------------------------------
__device__ __forceinline__ float dot_bf16(const bf16* __restrict__ a,
const bf16* __restrict__ b, int n) {
float acc = 0.f;
int j = 0;
// 8-wide (16-byte) chunks
const int n8 = n & ~7;
for (; j < n8; j += 8) {
const int4 va = *reinterpret_cast<const int4*>(a + j);
const int4 vb = *reinterpret_cast<const int4*>(b + j);
const bf16* pa = reinterpret_cast<const bf16*>(&va);
const bf16* pb = reinterpret_cast<const bf16*>(&vb);
#pragma unroll
for (int k = 0; k < 8; k++)
acc += __bfloat162float(pa[k]) * __bfloat162float(pb[k]);
}
for (; j < n; j++) acc += __bfloat162float(a[j]) * __bfloat162float(b[j]);
return acc;
}
__global__ void moe_decode_kernel(const bf16* __restrict__ x,
const int64_t* __restrict__ expert_ids,
const bf16* __restrict__ expert_w,
const bf16* __restrict__ w1,
const bf16* __restrict__ w2,
float* __restrict__ out, int T, int H, int I,
int top_k) {
const int a = (int)blockIdx.x;
const int t = a / top_k;
const int e = (int)expert_ids[a];
const float wt = __bfloat162float(expert_w[a]);
extern __shared__ char raw[];
bf16* xs = reinterpret_cast<bf16*>(raw);
float* mid = reinterpret_cast<float*>(xs + H);
for (int i = threadIdx.x; i < H; i += blockDim.x) xs[i] = x[(int64_t)t * H + i];
__syncthreads();
const bf16* w1e = w1 + (int64_t)e * (2 * I) * H;
const bf16* gate = w1e;
const bf16* up = w1e + (int64_t)I * H;
for (int i = threadIdx.x; i < I; i += blockDim.x) {
const float g = dot_bf16(xs, gate + (int64_t)i * H, H);
const float u = dot_bf16(xs, up + (int64_t)i * H, H);
mid[i] = (g / (1.f + expf(-g))) * u;
}
__syncthreads();
const bf16* w2e = w2 + (int64_t)e * H * I;
float* ot = out + (int64_t)t * H;
for (int h = threadIdx.x; h < H; h += blockDim.x) {
const bf16* drow = w2e + (int64_t)h * I;
float acc = 0.f;
int j = 0;
const int n8 = I & ~7;
for (; j < n8; j += 8) {
const int4 vb = *reinterpret_cast<const int4*>(drow + j);
const bf16* pb = reinterpret_cast<const bf16*>(&vb);
#pragma unroll
for (int k = 0; k < 8; k++) acc += mid[j + k] * __bfloat162float(pb[k]);
}
for (; j < I; j++) acc += mid[j] * __bfloat162float(drow[j]);
atomicAdd(ot + h, wt * acc);
}
}
void moe_decode(torch::Tensor x, torch::Tensor expert_ids,
torch::Tensor expert_weights, torch::Tensor w1,
torch::Tensor w2, torch::Tensor out) {
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 I = (int)w2.size(2);
const int n = T * top_k;
if (n == 0) return;
auto stream = at::cuda::getCurrentCUDAStream();
const size_t smem = sizeof(bf16) * H + sizeof(float) * I;
moe_decode_kernel<<<n, 256, smem, stream>>>(
reinterpret_cast<const bf16*>(x.data_ptr()),
expert_ids.data_ptr<int64_t>(),
reinterpret_cast<const bf16*>(expert_weights.data_ptr()),
reinterpret_cast<const bf16*>(w1.data_ptr()),
reinterpret_cast<const bf16*>(w2.data_ptr()), out.data_ptr<float>(), T, H,
I, top_k);
}
int tile_size_m() { return BM; }
int tile_size_n() { return BN; }
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("silu_and_mul", &silu_and_mul);
m.def("scatter_add", &scatter_add);
m.def("gather_rows", &gather_rows);
m.def("grouped_gemm_mma", &grouped_gemm_mma);
m.def("grouped_gemm_cublas", &grouped_gemm_cublas);
m.def("gemm_nt", &gemm_nt);
m.def("moe_decode", &moe_decode);
m.def("tile_size_m", &tile_size_m);
m.def("tile_size_n", &tile_size_n);
}
20260814_000242_grok_grok-4.6_01_glm52_fused_moe