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