KernelBench hard · RTX PRO 6000
Sonic MoE GLM-5.3
9.97%geomean peak fraction across shapes
manually audited: clean
Isolated regrade 0.0997. load_inline grouped GEMM with mma.sync m16n8k16 and fused silu(g)*u. Tile-start map is keyed on (expert_offsets.data_ptr, _version); out is fresh every call. Same-buffer overwrite of hidden and of offsets both match reference. Lint CLEAN. Numeric stress on.
harnesszai-claudeagent session2h 34mtotal wall2h 45mcheck2mbenchmark89soutput tokens265,744cost$18.50regimecompute
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
32768×4096×1536×128×820.393 ms8.1%0.30 TB/s · 17% of 1.8 TB/s HBM · also 40 TFLOPS (8% of compute)
4096×2048×1024×64×40.462 ms14.9%1.38 TB/s · 77% of 1.8 TB/s HBM · also 74 TFLOPS (15% of compute)
16384×2048×4096×64×813.352 ms8.2%0.28 TB/s · 16% of 1.8 TB/s HBM · also 41 TFLOPS (8% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(8.1% · 14.9% · 8.2%) = 10.0%
Kernel source (redacted)
"""Sonic-MoE up-projection: variable-length grouped GEMM + fused SwiGLU.
Custom CUDA kernel (mma.sync m16n8k16 bf16, cp.async multi-stage pipeline,
SMEM XOR swizzle, device-side tile->expert mapping, fused silu(g)*u epilogue).
A Triton kernel covers shapes outside the tuned dispatch table.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import triton
import triton.language as tl
# ---------------------------------------------------------------------------
# CUDA source (kept free of torch headers; the wrapper below is host-only)
# ---------------------------------------------------------------------------
_CUDA_SRC = r"""
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#include <cstdint>
#define DEVINL __device__ __forceinline__
DEVINL uint32_t smem_u32(const void* p) { return (uint32_t)__cvta_generic_to_shared(p); }
DEVINL void cp_async16(void* dst, const void* src, bool pred) {
uint32_t d = smem_u32(dst);
asm volatile("cp.async.cg.shared.global [%0], [%1], 16, %2;\n" ::"r"(d), "l"(src), "r"(pred ? 16 : 0));
}
DEVINL void cp_commit() { asm volatile("cp.async.commit_group;\n"); }
template <int N> DEVINL void cp_wait() { asm volatile("cp.async.wait_group %0;\n" ::"n"(N)); }
DEVINL void ldmatrix_x4(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));
}
DEVINL void mma_bf16(float* c, uint32_t a0, uint32_t a1, uint32_t a2, uint32_t a3, 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"(a0), "r"(a1), "r"(a2), "r"(a3), "r"(b0), "r"(b1));
}
DEVINL float fast_silu(float x) {
float e = exp2f(-1.4426950408889634f * x);
return x / (1.0f + e);
}
DEVINL __nv_bfloat162 pack2(float a, float b) {
__nv_bfloat162 r;
r.x = __float2bfloat16(a); r.y = __float2bfloat16(b);
return r;
}
// Exclusive prefix of per-expert M-tile counts (tile -> expert map).
template <int BM>
__global__ void build_tile_map(const int* __restrict__ offs, int* __restrict__ tile_start, int E) {
if (threadIdx.x == 0) {
int acc = 0;
for (int e = 0; e < E; ++e) {
tile_start[e] = acc;
acc += (offs[e + 1] - offs[e] + BM - 1) / BM;
}
tile_start[E] = acc;
}
}
// Grouped GEMM + fused SwiGLU. X (T,H); WG/WU (E,I,H) pre-transposed; OUT (T,I).
template <int BM, int BN, int BK, int STAGES, int WM, int WN, bool EVEN_K, bool EVEN_N>
__global__ __launch_bounds__(32 * WM * WN, 1) void moe_swiglu_kernel(
const __nv_bfloat16* __restrict__ X,
const __nv_bfloat16* __restrict__ WG,
const __nv_bfloat16* __restrict__ WU,
const int* __restrict__ offs,
const int* __restrict__ tile_start,
__nv_bfloat16* __restrict__ OUT,
int E, int H, int I, int num_n)
{
constexpr int ROWBYTES = BK * 2;
constexpr int CHUNKS = ROWBYTES / 16;
constexpr int A_ELEMS = BM * BK;
constexpr int B_ELEMS = BN * BK;
constexpr int THREADS = 32 * WM * WN;
constexpr int TM = (BM / WM) / 16;
constexpr int TN = (BN / WN) / 8;
extern __shared__ __align__(16) __nv_bfloat16 smem[];
__nv_bfloat16* A_s = smem;
__nv_bfloat16* Bg_s = A_s + STAGES * A_ELEMS;
__nv_bfloat16* Bu_s = Bg_s + STAGES * B_ELEMS;
const int tid = threadIdx.x;
const int lane = tid & 31;
const int warp = tid >> 5;
const int pid = blockIdx.x;
const int gm = pid / num_n;
const int gn = pid - gm * num_n;
if (gm >= tile_start[E]) return;
int lo = 0, hi = E;
while (lo + 1 < hi) {
int mid = (lo + hi) >> 1;
if (tile_start[mid] <= gm) lo = mid; else hi = mid;
}
const int e = lo;
const int row0 = offs[e] + (gm - tile_start[e]) * BM;
const int mend = offs[e + 1];
const int n0 = gn * BN;
const int k_tiles = EVEN_K ? (H / BK) : ((H + BK - 1) / BK);
const __nv_bfloat16* xg = X + (long long)row0 * H;
const __nv_bfloat16* wgg = WG + (long long)e * I * H + (long long)n0 * H;
const __nv_bfloat16* wug = WU + (long long)e * I * H + (long long)n0 * H;
auto issue_A = [&](int stage, int kt) {
const __nv_bfloat16* src = xg + (long long)kt * BK;
#pragma unroll
for (int idx = tid; idx < BM * CHUNKS; idx += THREADS) {
const int r = idx / CHUNKS;
const int c = idx - r * CHUNKS;
const bool ok = ((row0 + r) < mend) && (EVEN_K || ((kt * BK + c * 8) < H));
void* dst = (char*)(A_s + stage * A_ELEMS) + (long)r * ROWBYTES +
((c ^ ((r >> 1) & (CHUNKS - 1))) * 16);
cp_async16(dst, src + (long long)r * H + c * 8, ok);
}
};
auto issue_B = [&](const __nv_bfloat16* base, __nv_bfloat16* sm, int stage, int kt) {
const __nv_bfloat16* src = base + (long long)kt * BK;
#pragma unroll
for (int idx = tid; idx < BN * CHUNKS; idx += THREADS) {
const int r = idx / CHUNKS;
const int c = idx - r * CHUNKS;
const bool ok = (EVEN_N || ((n0 + r) < I)) && (EVEN_K || ((kt * BK + c * 8) < H));
void* dst = (char*)(sm + stage * B_ELEMS) + (long)r * ROWBYTES +
((c ^ ((r >> 1) & (CHUNKS - 1))) * 16);
cp_async16(dst, src + (long long)r * H + c * 8, ok);
}
};
auto issue_all = [&](int stage, int kt) {
issue_A(stage, kt);
issue_B(wgg, Bg_s, stage, kt);
issue_B(wug, Bu_s, stage, kt);
cp_commit();
};
#pragma unroll
for (int s = 0; s < STAGES - 1; ++s) {
if (s < k_tiles) issue_all(s, s); else cp_commit();
}
float acc_g[TM][TN][4], acc_u[TM][TN][4];
#pragma unroll
for (int i = 0; i < TM; ++i)
#pragma unroll
for (int j = 0; j < TN; ++j)
#pragma unroll
for (int q = 0; q < 4; ++q) { acc_g[i][j][q] = 0.f; acc_u[i][j][q] = 0.f; }
const int warp_m = warp / WN;
const int warp_n = warp % WN;
int stage = 0;
for (int kt = 0; kt < k_tiles; ++kt) {
cp_wait<STAGES - 2>();
__syncthreads();
const __nv_bfloat16* As = A_s + stage * A_ELEMS;
const __nv_bfloat16* Bgs = Bg_s + stage * B_ELEMS;
const __nv_bfloat16* Bus = Bu_s + stage * B_ELEMS;
#pragma unroll
for (int kk = 0; kk < BK / 16; ++kk) {
uint32_t a[TM][4];
uint32_t bg[TN / 2][4], bu[TN / 2][4];
#pragma unroll
for (int mi = 0; mi < TM; ++mi) {
const int arow = warp_m * (BM / WM) + mi * 16 + (lane & 7) + (((lane >> 3) & 1) << 3);
const int achunk = kk * 2 + (lane >> 4);
const uint32_t a_addr = smem_u32((char*)(As + arow * BK) +
((achunk ^ ((arow >> 1) & (CHUNKS - 1))) * 16));
ldmatrix_x4(a[mi][0], a[mi][1], a[mi][2], a[mi][3], a_addr);
}
#pragma unroll
for (int jc = 0; jc < TN / 2; ++jc) {
const int brow = warp_n * (BN / WN) + jc * 16 + (lane & 7) + (((lane >> 4) & 1) << 3);
const int bchunk = kk * 2 + ((lane >> 3) & 1);
const int sw = (brow >> 1) & (CHUNKS - 1);
const uint32_t bg_addr = smem_u32((char*)(Bgs + brow * BK) + ((bchunk ^ sw) * 16));
const uint32_t bu_addr = smem_u32((char*)(Bus + brow * BK) + ((bchunk ^ sw) * 16));
ldmatrix_x4(bg[jc][0], bg[jc][1], bg[jc][2], bg[jc][3], bg_addr);
ldmatrix_x4(bu[jc][0], bu[jc][1], bu[jc][2], bu[jc][3], bu_addr);
}
#pragma unroll
for (int jc = 0; jc < TN / 2; ++jc)
#pragma unroll
for (int mi = 0; mi < TM; ++mi) {
mma_bf16(acc_g[mi][jc * 2], a[mi][0], a[mi][1], a[mi][2], a[mi][3], bg[jc][0], bg[jc][1]);
mma_bf16(acc_g[mi][jc * 2 + 1], a[mi][0], a[mi][1], a[mi][2], a[mi][3], bg[jc][2], bg[jc][3]);
mma_bf16(acc_u[mi][jc * 2], a[mi][0], a[mi][1], a[mi][2], a[mi][3], bu[jc][0], bu[jc][1]);
mma_bf16(acc_u[mi][jc * 2 + 1], a[mi][0], a[mi][1], a[mi][2], a[mi][3], bu[jc][2], bu[jc][3]);
}
}
const int next = kt + STAGES - 1;
if (next < k_tiles) issue_all(next % STAGES, next); else cp_commit();
stage = (stage + 1) % STAGES;
}
const int g = lane >> 2;
const int t2 = (lane & 3) * 2;
#pragma unroll
for (int mi = 0; mi < TM; ++mi) {
const int r_lo = row0 + warp_m * (BM / WM) + mi * 16 + g;
const int r_hi = r_lo + 8;
const bool ok_lo = r_lo < mend;
const bool ok_hi = r_hi < mend;
__nv_bfloat16* o_lo = OUT + (long long)r_lo * I + n0 + warp_n * (BN / WN) + t2;
__nv_bfloat16* o_hi = OUT + (long long)r_hi * I + n0 + warp_n * (BN / WN) + t2;
#pragma unroll
for (int jc = 0; jc < TN; ++jc) {
__nv_bfloat162 v0 = pack2(fast_silu(acc_g[mi][jc][0]) * acc_u[mi][jc][0],
fast_silu(acc_g[mi][jc][1]) * acc_u[mi][jc][1]);
__nv_bfloat162 v1 = pack2(fast_silu(acc_g[mi][jc][2]) * acc_u[mi][jc][2],
fast_silu(acc_g[mi][jc][3]) * acc_u[mi][jc][3]);
const int cbase = n0 + warp_n * (BN / WN) + jc * 8;
if (ok_lo && (EVEN_N || (cbase + t2 + 1) < I)) *(__nv_bfloat162*)o_lo = v0;
if (ok_hi && (EVEN_N || (cbase + t2 + 1) < I)) *(__nv_bfloat162*)o_hi = v1;
o_lo += 8; o_hi += 8;
}
}
}
template <int BM, int BN, int BK, int STAGES, int WM, int WN>
void run_cfg(const void* X, const void* WG, const void* WU, const int* offs,
int* tile_start, void* OUT, int T_perm, int E, int H, int I,
bool build_map, cudaStream_t stream) {
constexpr int KEY = BM * 1000000 + BN * 1000 + BK * 10 + STAGES;
const int num_n = (I + BN - 1) / BN;
const int ub_m = (T_perm + BM - 1) / BM + E;
const int grid = ub_m * num_n;
const size_t smem = (size_t)STAGES * (BM + 2 * BN) * BK * 2;
const bool even_k = (H % BK == 0);
const bool even_n = (I % BN == 0);
auto k = moe_swiglu_kernel<BM, BN, BK, STAGES, WM, WN, true, true>;
auto k_ek = moe_swiglu_kernel<BM, BN, BK, STAGES, WM, WN, false, true>;
auto k_en = moe_swiglu_kernel<BM, BN, BK, STAGES, WM, WN, true, false>;
auto k_nn = moe_swiglu_kernel<BM, BN, BK, STAGES, WM, WN, false, false>;
static int smem_set = -1;
if (smem_set != KEY) {
cudaFuncSetAttribute(k, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
cudaFuncSetAttribute(k_ek, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
cudaFuncSetAttribute(k_en, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
cudaFuncSetAttribute(k_nn, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
smem_set = KEY;
}
if (build_map)
build_tile_map<BM><<<1, 32, 0, stream>>>(offs, tile_start, E);
auto* xp = (const __nv_bfloat16*)X;
auto* wg = (const __nv_bfloat16*)WG;
auto* wu = (const __nv_bfloat16*)WU;
auto* op = (__nv_bfloat16*)OUT;
if (even_k && even_n) k<<<grid, 32 * WM * WN, smem, stream>>>(xp, wg, wu, offs, tile_start, op, E, H, I, num_n);
else if (even_k) k_en<<<grid, 32 * WM * WN, smem, stream>>>(xp, wg, wu, offs, tile_start, op, E, H, I, num_n);
else if (even_n) k_ek<<<grid, 32 * WM * WN, smem, stream>>>(xp, wg, wu, offs, tile_start, op, E, H, I, num_n);
else k_nn<<<grid, 32 * WM * WN, smem, stream>>>(xp, wg, wu, offs, tile_start, op, E, H, I, num_n);
}
extern "C" void moe_forward_c(const void* X, const void* WG, const void* WU, const int* offs,
int* tile_start, void* OUT, int T_perm, int E, int H, int I,
int64_t build_map,
int64_t BM, int64_t BN, int64_t BK, int64_t STAGES,
int64_t WM, int64_t WN, cudaStream_t stream) {
const bool bm = build_map != 0;
if (BM == 128 && BN == 128 && BK == 32 && STAGES == 3 && WM == 2 && WN == 4) run_cfg<128,128,32,3,2,4>(X, WG, WU, offs, tile_start, OUT, T_perm, E, H, I, bm, stream);
else if (BM == 128 && BN == 64 && BK == 64 && STAGES == 3 && WM == 2 && WN == 2) run_cfg<128,64,64,3,2,2>(X, WG, WU, offs, tile_start, OUT, T_perm, E, H, I, bm, stream);
else { }
}
"""
_CPP_SRC = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAStream.h>
#include <cstdint>
extern "C" void moe_forward_c(const void* X, const void* WG, const void* WU, const int* offs,
int* tile_start, void* OUT, int T_perm, int E, int H, int I,
int64_t build_map,
int64_t BM, int64_t BN, int64_t BK, int64_t STAGES,
int64_t WM, int64_t WN, cudaStream_t stream);
void moe_forward(torch::Tensor X, torch::Tensor WG, torch::Tensor WU, torch::Tensor offs,
torch::Tensor tile_start, torch::Tensor OUT, int64_t build_map,
int64_t BM, int64_t BN, int64_t BK, int64_t STAGES,
int64_t WM, int64_t WN) {
const int T_perm = (int)X.size(0);
const int E = (int)WG.size(0), I = (int)WG.size(1), H = (int)WG.size(2);
cudaStream_t stream = c10::cuda::getCurrentCUDAStream().stream();
moe_forward_c(X.data_ptr(), WG.data_ptr(), WU.data_ptr(),
offs.data_ptr<int>(), tile_start.data_ptr<int>(), OUT.data_ptr(),
T_perm, E, H, I, build_map, BM, BN, BK, STAGES, WM, WN, stream);
}
"""
_MOD_CACHE: dict = {}
def _get_cuda_module():
key = "moe_swiglu_v3"
if key not in _MOD_CACHE:
import os
from torch.utils.cpp_extension import load_inline
_MOD_CACHE[key] = load_inline(
name=key,
cpp_sources=[_CPP_SRC + '\nPYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("moe_forward", &moe_forward); }'],
cuda_sources=[_CUDA_SRC],
extra_cuda_cflags=["-O3", "--use_fast_math"],
no_implicit_headers=True,
verbose=False,
)
return _MOD_CACHE[key]
# Tuned (BM, BN, BK, STAGES, WM, WN) keyed by (H, I, E).
_CFG_TABLE = {
(2048, 1024, 64): (128, 64, 64, 3, 2, 2),
(2048, 4096, 64): (128, 128, 32, 3, 2, 4),
(4096, 1536, 128): (128, 128, 32, 3, 2, 4),
}
# ---------------------------------------------------------------------------
# Triton fallback for shapes outside the tuned table
# ---------------------------------------------------------------------------
@triton.jit
def _moe_swiglu_triton(
X, WG, WU, OUT, OFFS,
H: tl.constexpr, I: tl.constexpr, E: tl.constexpr,
stride_x: tl.constexpr, stride_we: tl.constexpr, stride_wk: tl.constexpr,
stride_o: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_M: tl.constexpr,
):
pid = tl.program_id(0)
num_n = I // BLOCK_N if I % BLOCK_N == 0 else (I + BLOCK_N - 1) // BLOCK_N
acc = 0
e_found = -1
row0 = 0
local = 0
ncount = 0
for e in range(E):
s = tl.load(OFFS + e)
t = tl.load(OFFS + e + 1)
c = t - s
tiles = ((c + BLOCK_M - 1) // BLOCK_M) * num_n
hit = (e_found < 0) & (pid >= acc) & (pid < acc + tiles)
e_found = tl.where(hit, e, e_found)
row0 = tl.where(hit, s, row0)
ncount = tl.where(hit, c, ncount)
local = tl.where(hit, pid - acc, local)
acc += tiles
if pid >= acc:
return
num_m = (ncount + BLOCK_M - 1) // BLOCK_M
group_id = local // (GROUP_M * num_n)
first_m = group_id * GROUP_M
gsize = tl.minimum(num_m - first_m, GROUP_M)
pid_m = first_m + ((local % (GROUP_M * num_n)) % gsize)
pid_n = (local % (GROUP_M * num_n)) // gsize
rm = row0 + pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
rk = tl.arange(0, BLOCK_K)
m_mask = rm < row0 + ncount
xp = X + rm[:, None] * stride_x + rk[None, :]
wgp = WG + e_found * stride_we + rk[:, None] * stride_wk + rn[None, :]
wup = WU + e_found * stride_we + rk[:, None] * stride_wk + rn[None, :]
acc_g = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
acc_u = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(0, H, BLOCK_K):
a = tl.load(xp, mask=m_mask[:, None], other=0.0)
bg = tl.load(wgp)
bu = tl.load(wup)
acc_g = tl.dot(a, bg, acc_g)
acc_u = tl.dot(a, bu, acc_u)
xp += BLOCK_K
wgp += BLOCK_K * stride_wk
wup += BLOCK_K * stride_wk
g = acc_g
u = acc_u
out = (g * tl.sigmoid(g)) * u
op = OUT + rm[:, None] * stride_o + rn[None, :]
tl.store(op, out.to(tl.bfloat16), mask=m_mask[:, None])
class Model(nn.Module):
"""Up-projection of a top-K MoE FFN with fused SwiGLU."""
def __init__(self, T_total: int, H: int, I: int, E: int, K: int): # noqa: E741
super().__init__()
self.T_total = T_total
self.H = H
self.I = I
self.E = E
self.K = K
self.W_gate = nn.Parameter(torch.empty(E, H, I, dtype=torch.bfloat16))
self.W_up = nn.Parameter(torch.empty(E, H, I, dtype=torch.bfloat16))
nn.init.normal_(self.W_gate, std=0.02)
nn.init.normal_(self.W_up, std=0.02)
self._wt_cache = None # (version, W_gate_t, W_up_t)
self._ts_buf = None
self._ts_key = None # (data_ptr, _version) of the offsets tensor
def _weights_t(self):
v = self.W_gate._version
cache = self._wt_cache
if cache is None or cache[0] != v or cache[1].device != self.W_gate.device:
wgt = self.W_gate.permute(0, 2, 1).contiguous()
wut = self.W_up.permute(0, 2, 1).contiguous()
self._wt_cache = (v, wgt, wut)
return self._wt_cache[1], self._wt_cache[2]
def forward(
self,
hidden_states: torch.Tensor, # (T_perm, H) bf16
expert_offsets: torch.Tensor, # (E+1,) int32
) -> torch.Tensor:
T_perm, H = hidden_states.shape
I = self.I
E = self.E
out = torch.empty(T_perm, I, dtype=torch.bfloat16, device=hidden_states.device)
cfg = _CFG_TABLE.get((H, I, E))
if cfg is not None:
mod = _get_cuda_module()
wgt, wut = self._weights_t()
key = (expert_offsets.data_ptr(), expert_offsets._version)
if self._ts_buf is None or self._ts_buf.device != out.device:
self._ts_buf = torch.zeros(E + 1, dtype=torch.int32, device=out.device)
self._ts_key = None
# tile_start depends only on the offsets contents: rebuild the
# map only when the tensor identity or version changes.
build = self._ts_key != key
mod.moe_forward(hidden_states, wgt, wut, expert_offsets,
self._ts_buf, out, 1 if build else 0, *cfg)
if build:
self._ts_key = key
return out
# Triton fallback
BM, BN, BK, GM, warps, stages = _config(T_perm, H, I, E)
num_n = (I + BN - 1) // BN
max_tiles = ((T_perm + BM - 1) // BM + E) * num_n
_moe_swiglu_triton[(max_tiles,)](
hidden_states, self.W_gate, self.W_up, out, expert_offsets,
H=H, I=I, E=E,
stride_x=hidden_states.stride(0), stride_we=self.W_gate.stride(0),
stride_wk=self.W_gate.stride(1), stride_o=out.stride(0),
BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, GROUP_M=GM,
num_warps=warps, num_stages=stages,
)
return out
def _config(T_perm: int, H: int, I: int, E: int):
return (128, 128, 64, 8, 8, 3)
# Module-level shape shims rewritten by check.py / benchmark.py per shape.
T_total = 32768
H = 4096
I = 1536 # noqa: E741
E = 128
K = 8
def _build_routing(T_total: int, E: int, K: int, device: str = "cpu") -> torch.Tensor:
T_perm = T_total * K
base = T_perm // E
rem = T_perm - base * E
counts = torch.full((E,), base, dtype=torch.int32, device=device)
counts[:rem] += 1
offsets = torch.zeros(E + 1, dtype=torch.int32, device=device)
offsets[1:] = torch.cumsum(counts, dim=0)
return offsets
def get_inputs():
T_perm = T_total * K
hidden_states = torch.randn(T_perm, H, dtype=torch.bfloat16) * 0.1
expert_offsets = _build_routing(T_total, E, K)
return [hidden_states, expert_offsets]
def get_init_inputs():
return [T_total, H, I, E, K]
20260822_075101_zai-claude_glm-5.3_06_sonic_moe_swiglu