KernelBench hard · RTX PRO 6000
Sonic MoE DeepSeek V4.1 Flash
10.8%geomean peak fraction across shapes
manually audited: clean
Triton variable-length grouped GEMM with SwiGLU genuinely fused into the epilogue. W_gate and W_up are interleaved column-wise into one (E, H, 2I) tensor so both projections come out of a single tl.dot into one accumulator, then reshape/split/silu is a register rename rather than a shared-memory round trip. 387.8 TFLOPS real, 77% of this part's bf16 dense peak.
harnessdeepseek-claudeagent session1h 51mtotal wall1h 59mcheck2mbenchmark2moutput tokens149,924cost$8.64gpu-lock wait46mgpu-lock held1h 2mregimecompute
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
32768×4096×1536×128×817.010 ms9.7%0.36 TB/s · 20% of 1.8 TB/s HBM · also 48 TFLOPS (10% of compute)
4096×2048×1024×64×40.511 ms13.5%1.25 TB/s · 69% of 1.8 TB/s HBM · also 67 TFLOPS (13% of compute)
16384×2048×4096×64×811.356 ms9.7%0.33 TB/s · 18% of 1.8 TB/s HBM · also 48 TFLOPS (10% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(9.7% · 13.5% · 9.7%) = 10.8%
Kernel source (redacted)
"""Triton grouped GEMM (variable-length, expert-major) with fused SwiGLU epilogue.
The MoE up-projection
h_e = silu(x_e @ W_gate[e]) * (x_e @ W_up[e])
is computed by a single fused kernel. Both projections share the activation
tile, so they are evaluated as one wide GEMM over a *column-interleaved* weight
layout
W_cat[e][k][2*j] = W_gate[e][k][j]
W_cat[e][k][2*j+1] = W_up[e][k][j]
which turns "two dots into two accumulators" into "one dot into one
accumulator". That matters on this part: the fused kernel is register limited,
and a single accumulator of the same footprint keeps the mainloop at one mma
pipeline with a single B stream instead of two interleaved ones (measured ~6%
faster than the two-accumulator formulation on the headline shape, and it
matches the throughput of a plain single-output grouped GEMM).
After the K loop the accumulator holds the gate and up values of the same
output channel in adjacent columns. The m16n8k16 accumulator layout keeps
consecutive column pairs in the same thread, so ``reshape -> split`` is a
register rename, not a shared-memory round trip; the epilogue is then a plain
elementwise ``silu(g) * u`` that is written straight to global memory. Neither
GEMM result ever touches DRAM separately.
Scheduling / variable-length M
------------------------------
The grid is ``(MP, I/BN, E)`` -- the expert index is the slowest-varying axis so
the programs of one expert are launched back to back and keep that expert's
weights hot in L2 while its activation rows stream through.
M is not uniform across experts, so the kernel derives it from the routing
metadata itself: every program reads ``expert_offsets[e]`` / ``expert_offsets[e+1]``
and walks ``mt = pid_m, pid_m + MP, ...`` up to ``ceil(nrows / BM)``. No host
side knowledge of the routing is needed, hence no device synchronisation, and
ragged/skewed routing (including empty experts, which are skipped) is handled
exactly. Rows past the end of an expert are masked on the activation load and
on the store, so no work is done for rows that belong to another expert.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import triton
import triton.language as tl
@triton.jit
def _moe_up_swiglu_kernel(
x_ptr, # (T_perm, H) bf16
w_ptr, # (E, H, 2*I) bf16, gate/up interleaved on the last axis
out_ptr, # (T_perm, I) bf16
off_ptr, # (E + 1,) int32 prefix sums
s_xm, s_xk,
s_we, s_wk, s_wn,
s_om, s_on,
H: tl.constexpr,
I: tl.constexpr,
MP: tl.constexpr,
BM: tl.constexpr,
BN: tl.constexpr, # output (gate) channels per tile; weight tile is 2*BN wide
BK: tl.constexpr,
MASK_N: tl.constexpr,
MASK_K: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
e = tl.program_id(2)
start = tl.load(off_ptr + e)
nrows = tl.load(off_ptr + e + 1) - start
num_mt = tl.cdiv(nrows, BM)
offs_m = tl.arange(0, BM)
offs_w = pid_n * 2 * BN + tl.arange(0, 2 * BN)
offs_k = tl.arange(0, BK)
wb = w_ptr + e * s_we
w_off = offs_k[:, None] * s_wk + offs_w[None, :] * s_wn
for mt in range(pid_m, num_mt, MP):
rm = mt * BM + offs_m
row_ok = rm < nrows
a_ptrs = x_ptr + (start + rm)[:, None] * s_xm + offs_k[None, :] * s_xk
b_ptrs = wb + w_off
acc = tl.zeros((BM, 2 * BN), dtype=tl.float32)
for _k in range(0, tl.cdiv(H, BK)):
if MASK_K:
km = offs_k[None, :] < H - _k * BK
a = tl.load(a_ptrs, mask=row_ok[:, None] & km, other=0.0)
else:
a = tl.load(a_ptrs, mask=row_ok[:, None], other=0.0)
if MASK_N:
b = tl.load(b_ptrs, mask=offs_w[None, :] < 2 * I, other=0.0)
else:
b = tl.load(b_ptrs)
acc = tl.dot(a, b, acc)
a_ptrs += BK * s_xk
b_ptrs += BK * s_wk
pair = tl.reshape(acc, (BM, BN, 2))
gate, up = tl.split(pair)
y = (gate * tl.sigmoid(gate) * up).to(tl.bfloat16)
col = pid_n * BN + tl.arange(0, BN)
c_ptrs = out_ptr + (start + rm)[:, None] * s_om + col[None, :] * s_on
if MASK_N:
tl.store(c_ptrs, y, mask=row_ok[:, None] & (col[None, :] < I))
else:
tl.store(c_ptrs, y, mask=row_ok[:, None])
# (H, I, E) -> (MP, BM, BN, BK, num_warps, num_stages)
#
# Every entry keeps BM * 2*BN == 16384 fp32 accumulator slots -- the largest
# footprint that stays off the spill path on SM120 (254 regs, 0 spills). Larger
# tiles spill and lose 30-60%; the alternative aspect ratio at the same
# footprint (BM=256, BN=32) is 7% slower, so M-heavy tiles are the wrong trade.
#
# Stage count is *not* monotonic: 3 beats 4 by ~4% on the headline shape while
# the ordering flips on the wide one. Deep pipelines only pay while the K loop
# is long relative to the tile; past that they just add L2 pressure.
_TUNED = {
(4096, 1536, 128): (8, 128, 64, 32, 4, 3),
(2048, 1024, 64): (2, 128, 64, 32, 4, 3),
(2048, 4096, 64): (16, 128, 64, 32, 4, 4),
}
def _default_config(H: int, I: int, E: int) -> tuple[int, int, int, int, int, int]:
for bn in (64, 32, 16, 8, 4, 2, 1):
if I % bn == 0:
break
for bk in (32, 16, 8, 4, 2, 1):
if H % bk == 0:
break
return (8, 128, bn, bk, 4, 3)
class Model(nn.Module):
"""Up-projection of a top-K MoE FFN with fused SwiGLU (grouped GEMM)."""
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._packed = None # (W_gate, W_up, ver_g, ver_u, W_cat)
# -- weight packing ---------------------------------------------------
def _packed_weight(self, wg: torch.Tensor, wu: torch.Tensor) -> torch.Tensor:
"""Interleave gate/up along the last axis.
Cached across calls: the packing is a pure function of the two weight
tensors, so it is rebuilt only when they are replaced or written to.
Holding references to the source tensors (rather than just their
data_ptrs) is what makes the identity check safe -- their storage can
never be freed and recycled underneath us.
"""
cacheable = wg.is_contiguous() and wu.is_contiguous()
cached = self._packed
if (cacheable and cached is not None
and cached[0] is wg and cached[1] is wu
and cached[2] == wg._version and cached[3] == wu._version):
return cached[4]
E, H, I = wg.shape
w_cat = torch.empty(E, H, 2 * I, dtype=wg.dtype, device=wg.device)
view = w_cat.view(E, H, I, 2)
view[..., 0].copy_(wg) # strided copy: handles non-contiguous sources
view[..., 1].copy_(wu)
if cacheable:
self._packed = (wg, wu, wg._version, wu._version, w_cat)
return w_cat
def forward(
self,
hidden_states: torch.Tensor,
expert_offsets: torch.Tensor,
) -> torch.Tensor:
T_perm, H = hidden_states.shape
I = self.I
if hidden_states.stride(1) != 1:
hidden_states = hidden_states.contiguous()
wg = self.W_gate
wu = self.W_up
if expert_offsets.dtype != torch.int32:
expert_offsets = expert_offsets.to(torch.int32)
w_cat = self._packed_weight(wg, wu)
out = torch.empty(T_perm, I, dtype=torch.bfloat16, device=hidden_states.device)
cfg = _TUNED.get((H, I, self.E))
if cfg is None:
cfg = _default_config(H, I, self.E)
MP, BM, BN, BK, warps, stages = cfg
grid = (MP, triton.cdiv(I, BN), self.E)
_moe_up_swiglu_kernel[grid](
hidden_states, w_cat, out, expert_offsets,
hidden_states.stride(0), hidden_states.stride(1),
w_cat.stride(0), w_cat.stride(1), w_cat.stride(2),
out.stride(0), out.stride(1),
H=H, I=I, MP=MP, BM=BM, BN=BN, BK=BK,
MASK_N=(I % BN) != 0,
MASK_K=(H % BK) != 0,
num_warps=warps, num_stages=stages,
)
return out
20260910_202134_deepseek-claude_deepseek-flash_06_sonic_moe_swiglu