kernelbench.com

KernelBench hard · H100

Sonic MoE GLM-5.2

6.52%geomean peak fraction across shapes
harnesszai-claudeagent session1h 20mtotal wall1h 40mcheck17mbenchmark3moutput tokenscost$11.72gpu-lock wait15mgpu-lock held5mregimecompute

Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth

32768×4096×1536×128×819.776 ms5.5%0.31 TB/s · 15% of 2.0 TB/s HBM · also 42 TFLOPS (6% of compute)
4096×2048×1024×64×40.500 ms9.1%1.28 TB/s · 63% of 2.0 TB/s HBM · also 69 TFLOPS (9% of compute)
16384×2048×4096×64×813.159 ms5.5%0.29 TB/s · 14% of 2.0 TB/s HBM · also 42 TFLOPS (6% of compute)

compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)

geomean(5.5% · 9.1% · 5.5%) = 6.5%

Kernel source (redacted)
"""Custom grouped-GEMM + fused SwiGLU up-projection for a top-K MoE FFN.

Per expert e:
    h_e = silu(x_e @ W_gate[e]) * (x_e @ W_up[e])

Single fused *persistent* Triton kernel (Hopper TMA + wgmma). Each persistent
CTA sweeps a round-robin slice of the tile space; per tile it loads one M-tile
of the permuted hidden states ONCE via TMA and runs two bf16 tensor-core matmuls
(gate, up) in fp32 accumulators against that expert's weight slabs, then writes
silu(gate)*up. The two B operands are loaded back-to-back around the gate dot
("seq" order) so Triton can overlap the up-weight TMA with the gate wgmma and
reuse the B shared-memory slot -- this is markedly faster than loading both
weights up front. Reading x once and never materialising gate/up halves hidden
traffic and removes two HBM round trips.

Tile schedule is (expert, m_tile, n_block) with n innermost so that, while an
expert's tiles are being worked, its hidden slab and weight slab stay L2
resident and are reused across all of that expert's N-blocks / M-tiles.

Variable-length grouping: tiles never cross an expert boundary. The output
store is masked per row so a partial trailing M-tile (the general
variable-offsets case) only writes its valid rows.
"""
from __future__ import annotations

import torch
import torch.nn as nn
import triton
import triton.language as tl
from triton.tools.tensor_descriptor import TensorDescriptor

OP_TYPE = "grouped_gemm_swiglu"
SUPPORTED_PRECISIONS = ["bf16"]
HARDWARE_REQUIRED = ["RTX_PRO_6000", "H100", "B200"]

# Fixed tile shape: the host-side TMA descriptors are built with this block
# shape, so the kernel must use the same. num_stages (pipelining depth) is the
# only knob autotuned -- it does not affect the descriptor block shape.
BLOCK_M = 128
BLOCK_N = 128
BLOCK_K = 64


@triton.autotune(
    configs=[
        triton.Config({}, num_warps=8, num_stages=2),
        triton.Config({}, num_warps=8, num_stages=3),
        triton.Config({}, num_warps=8, num_stages=4),
    ],
    key=["H", "I"],
)
@triton.jit
def _fused_grouped_swiglu_kernel(
    a_desc, wg_desc, wu_desc, out_ptr,
    sched_e_ptr, sched_ms_ptr, sched_me_ptr, sched_n_ptr,
    num_tiles,
    H: tl.constexpr,
    I: tl.constexpr,  # noqa: E741
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    start_pid = tl.program_id(0)
    num_ctas = tl.num_programs(0)

    for tile_id in range(start_pid, num_tiles, num_ctas):
        e = tl.load(sched_e_ptr + tile_id)
        n_start = tl.load(sched_n_ptr + tile_id)
        m_start = tl.load(sched_ms_ptr + tile_id)
        m_end = tl.load(sched_me_ptr + tile_id)
        eH = e * H

        acc_gate = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
        acc_up = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

        for k in range(0, H, BLOCK_K):
            a = a_desc.load([m_start, k])
            b_gate = wg_desc.load([eH + k, n_start])
            acc_gate = tl.dot(a, b_gate, acc_gate)
            # Load the up-weight after the gate dot: overlaps the TMA with the
            # gate wgmma and lets the B shared slot be reused.
            b_up = wu_desc.load([eH + k, n_start])
            acc_up = tl.dot(a, b_up, acc_up)

        gate = acc_gate
        out = (gate * tl.sigmoid(gate)) * acc_up
        out = out.to(tl.bfloat16)

        offs_m = m_start + tl.arange(0, BLOCK_M)
        offs_n = n_start + tl.arange(0, BLOCK_N)
        out_ptrs = out_ptr + offs_m[:, None] * I + offs_n[None, :]
        tl.store(out_ptrs, out, mask=offs_m[:, None] < m_end)


def _build_schedule(offsets: torch.Tensor, block_m: int, i_dim: int, block_n: int):
    """Flat tile schedule in (expert, m_tile, n_block) order, n innermost.

    Returns 4 int32 GPU tensors of length total_tiles.
    """
    device = offsets.device
    counts = (offsets[1:] - offsets[:-1]).to(torch.int32)
    E = counts.numel()
    mt_per_expert = (counts + block_m - 1) // block_m          # (E,)
    num_n = (i_dim + block_n - 1) // block_n

    total_m = int(mt_per_expert.sum().item())
    arange_E = torch.arange(E, device=device, dtype=torch.int32)
    expert_m = torch.repeat_interleave(arange_E, mt_per_expert).to(torch.int32)
    pos = torch.arange(total_m, device=device, dtype=torch.int32)
    cs_m = torch.cumsum(mt_per_expert, dim=0, dtype=torch.int32)
    starts_m = torch.empty_like(cs_m)
    starts_m[0] = 0
    starts_m[1:] = cs_m[:-1]
    local_mb = pos - starts_m[expert_m]
    m_start = (offsets[expert_m] + local_mb * block_m).to(torch.int32)
    m_end = offsets[expert_m + 1].to(torch.int32)

    n_idx = torch.arange(num_n, device=device, dtype=torch.int32)
    # (m_tile, n_block) meshgrid, m outer / n inner (L2-friendly: reuse A[m]
    # across the n-sweep within an expert; expert boundaries respected since
    # m-tiles are emitted expert-major).
    M_grid, N_grid = torch.meshgrid(
        torch.arange(total_m, device=device, dtype=torch.int32), n_idx, indexing="ij"
    )
    mi = M_grid.reshape(-1)
    ni = N_grid.reshape(-1)
    sched_e = expert_m[mi].to(torch.int32)
    sched_ms = m_start[mi].to(torch.int32)
    sched_me = m_end[mi].to(torch.int32)
    sched_n = (ni * block_n).to(torch.int32)
    return sched_e, sched_ms, sched_me, sched_n


_PLAN_CACHE: dict = {}


def _get_schedule(offsets: torch.Tensor, block_m: int, i_dim: int, block_n: int):
    key = (offsets.data_ptr(), offsets.numel(), block_m, block_n)
    plan = _PLAN_CACHE.get(key)
    if plan is None:
        plan = _build_schedule(offsets, block_m, i_dim, block_n)
        _PLAN_CACHE.clear()
        _PLAN_CACHE[key] = plan
    return plan


def _num_sms() -> int:
    try:
        return torch.cuda.get_device_properties(0).multi_processor_count
    except Exception:
        return 114


class Model(nn.Module):
    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))
        self._sms = _num_sms()

    def forward(self, hidden_states: torch.Tensor, expert_offsets: torch.Tensor) -> 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)
        if expert_offsets.dtype != torch.int32:
            expert_offsets = expert_offsets.to(torch.int32)

        sched_e, sched_ms, sched_me, sched_n = _get_schedule(
            expert_offsets, BLOCK_M, I, BLOCK_N)
        total_tiles = sched_e.numel()

        # Host-side TMA descriptors over flat views. W is (E,H,I) contiguous,
        # seen here as (E*H, I); expert e's slab lives at flat row e*H.
        a_desc = TensorDescriptor.from_tensor(hidden_states, [BLOCK_M, BLOCK_K])
        wg_desc = TensorDescriptor.from_tensor(self.W_gate.view(E * H, I), [BLOCK_K, BLOCK_N])
        wu_desc = TensorDescriptor.from_tensor(self.W_up.view(E * H, I), [BLOCK_K, BLOCK_N])

        grid = (self._sms,)
        _fused_grouped_swiglu_kernel[grid](
            a_desc, wg_desc, wu_desc, out,
            sched_e, sched_ms, sched_me, sched_n,
            total_tiles,
            H, I,
            BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
        )
        return out


# 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]

20260719_034935_zai-claude_glm-5.2_06_sonic_moe_swiglu