kernelbench.com

KernelBench hard · H100

Sonic MoE Tencent Hy3

slowdid not score
agent session2h 22mtotal wall2h 26mcheck4mbenchmarkoutput tokensgpu-lock wait46sgpu-lock held3mregimecompute

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

No per-shape benchmark data archived for this run.

Kernel source (redacted)
"""Grouped GEMM + fused SwiGLU up-projection for MoE FFN (H100 / SM90, bf16).

Custom Triton kernel. We launch one fused kernel per expert e. Each expert owns
a variable-length slice x_e = hidden_states[offsets[e]:offsets[e+1]] of shape
(n_e, H); the kernel computes, in a single pass over those rows,

    gate = x_e @ W_gate[e]   (n_e, I)
    up   = x_e @ W_up[e]     (n_e, I)

accumulating in fp32 on the tensor core (bf16 mma.sync), applies silu to the
gate accumulation, multiplies by the up accumulation, and writes the (n_e, I)
result. Fusing the two GEMMs means x_e is read once and the output written once.

The variable-length grouped structure is naturally handled by one launch per
expert (E <= 128); launch overhead is negligible next to the per-expert compute
(2048x4096x1536 FMAs), and per-expert launches let each CTA use large,
tensor-core-friendly tiles without cross-expert coordination overhead.
"""
from __future__ import annotations

import torch
import torch.nn as nn
import triton
import triton.language as tl


@triton.autotune(
    configs=[
        triton.Config({"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=4, num_stages=3),
        triton.Config({"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 128}, num_warps=4, num_stages=2),
        triton.Config({"BLOCK_M": 128, "BLOCK_N": 64, "BLOCK_K": 128}, num_warps=4, num_stages=2),
        triton.Config({"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 128}, num_warps=4, num_stages=3),
        triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=4, num_stages=2),
        triton.Config({"BLOCK_M": 64, "BLOCK_N": 64, "BLOCK_K": 128}, num_warps=4, num_stages=3),
    ],
    key=["M", "N", "K"],
)
@triton.jit
def _fused_swiglu_gemm(
    X, Wg, Wu, OUT,
    stride_xm, stride_xk,
    stride_wk, stride_wn,
    stride_om, stride_ok,
    M, N, K,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)

    rm = 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)

    x_ptrs = X + (rm[:, None] * stride_xm + rk[None, :] * stride_xk)
    wg_ptrs = Wg + (rk[:, None] * stride_wk + rn[None, :] * stride_wn)
    wu_ptrs = Wu + (rk[:, None] * stride_wk + rn[None, :] * stride_wn)

    acc_g = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    acc_u = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    mask_m = rm < M
    mask_n = rn < N
    full_mask = mask_m[:, None] & mask_n[None, :]

    for k in range(0, K, BLOCK_K):
        kk = k + rk
        k_mask = kk < K
        xm = tl.load(x_ptrs, mask=k_mask[None, :] & mask_m[:, None], other=0.0)
        wg = tl.load(wg_ptrs, mask=k_mask[:, None] & mask_n[None, :], other=0.0)
        wu = tl.load(wu_ptrs, mask=k_mask[:, None] & mask_n[None, :], other=0.0)
        acc_g = tl.dot(xm, wg, acc=acc_g)
        acc_u = tl.dot(xm, wu, acc=acc_u)
        x_ptrs += BLOCK_K * stride_xk
        wg_ptrs += BLOCK_K * stride_wk
        wu_ptrs += BLOCK_K * stride_wk

    sig = acc_g * tl.sigmoid(acc_g)
    out = sig * acc_u

    out_ptrs = OUT + (rm[:, None] * stride_om + rn[None, :] * stride_ok)
    tl.store(out_ptrs, out.to(tl.bfloat16), mask=full_mask)


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))
        nn.init.normal_(self.W_gate, std=0.02)
        nn.init.normal_(self.W_up, std=0.02)

    def forward(self, hidden_states, expert_offsets):
        T_perm, H = hidden_states.shape
        device = hidden_states.device
        out = torch.empty(T_perm, self.I, dtype=torch.bfloat16, device=device)
        offsets = expert_offsets.to(torch.int32)

        for e in range(self.E):
            start = int(offsets[e].item())
            end = int(offsets[e + 1].item())
            n_e = end - start
            if n_e == 0:
                continue
            x_e = hidden_states[start:end]
            out_e = out[start:end]
            grid = lambda META: (
                triton.cdiv(n_e, META["BLOCK_M"]),
                triton.cdiv(self.I, META["BLOCK_N"]),
            )
            _fused_swiglu_gemm[grid](
                x_e, self.W_gate[e], self.W_up[e], out_e,
                x_e.stride(0), x_e.stride(1),
                self.W_gate[e].stride(0), self.W_gate[e].stride(1),
                out_e.stride(0), out_e.stride(1),
                M=n_e, N=self.I, K=self.H,
            )
        return out


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]

20260709_090251_hy3_hy3_06_sonic_moe_swiglu