kernelbench.com

KernelBench hard · H100

FP8 GEMM Grok 4.5

23.8%geomean peak fraction across shapes
harnessgrokagent session42mtotal wall43mcheck61sbenchmark12soutput tokensgpu-lock wait45sgpu-lock held28sregimecompute

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

4096×4096×40960.201 ms45.3%685 TFLOPS · 45% of 1,513 TF fp8 peak · also 0.33 TB/s (16% of HBM)
4096×4096×41270.240 ms38.1%576 TFLOPS · 38% of 1,513 TF fp8 peak · also 0.28 TB/s (14% of HBM)
32×8192×81920.083 ms3.4%0.82 TB/s · 40% of 2.0 TB/s HBM · also 52 TFLOPS (3% of compute)
4096×14336×40960.579 ms54.9%830 TFLOPS · 55% of 1,513 TF fp8 peak · also 0.33 TB/s (16% of HBM)

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

geomean(45.3% · 38.1% · 3.4% · 54.9%) = 23.8%

Kernel source (redacted)
"""FP8 e4m3 GEMM via Triton tl.dot (tensor-core fp8 x fp8, fp32 accumulate).

y = (x @ weight.T) * weight_scale  -> bf16
x: (M, K) fp8_e4m3, weight: (N, K) fp8_e4m3, weight_scale: (N,) fp32

Strategy:
  - Real fp8 x fp8 tensor-core MMA via tl.dot (fp32 accumulate), then per-channel scale.
  - Unaligned K is zero-padded to a multiple of BLOCK_K so row-stride stays tile-aligned
    (raw unaligned K is ~80x slower on Hopper).
  - CUDA graph over static workspaces removes Python/launch overhead (critical for skinny M).
  - Weight staging is version-tracked so numeric-stress in-place mutations stay correct.
"""
from __future__ import annotations

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

E4M3_MAX = 448.0


@triton.jit
def _fp8_gemm_kernel(
    a_ptr,
    b_ptr,
    scale_ptr,
    c_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bn, stride_bk,
    stride_cm, stride_cn,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    GROUP_M: tl.constexpr,
    EVEN_M: tl.constexpr,
    EVEN_N: tl.constexpr,
):
    pid = tl.program_id(axis=0)
    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    num_pid_in_group = GROUP_M * num_pid_n
    group_id = pid // num_pid_in_group
    first_pid_m = group_id * GROUP_M
    group_size_m = min(num_pid_m - first_pid_m, GROUP_M)
    pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)
    pid_n = (pid % num_pid_in_group) // group_size_m

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)

    a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
    # B.T tile (BLOCK_K, BLOCK_N): element (k, n) = B[n, k]
    b_ptrs = b_ptr + offs_n[None, :] * stride_bn + offs_k[:, None] * stride_bk

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    for _k in range(0, tl.cdiv(K, BLOCK_K)):
        if EVEN_M:
            a = tl.load(a_ptrs)
        else:
            a = tl.load(a_ptrs, mask=offs_m[:, None] < M, other=0.0)
        if EVEN_N:
            b = tl.load(b_ptrs)
        else:
            b = tl.load(b_ptrs, mask=offs_n[None, :] < N, other=0.0)
        acc = tl.dot(a, b, acc)
        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += BLOCK_K * stride_bk

    if EVEN_N:
        scales = tl.load(scale_ptr + offs_n)
    else:
        scales = tl.load(scale_ptr + offs_n, mask=offs_n < N, other=0.0)
    acc = acc * scales[None, :]

    c = acc.to(tl.bfloat16)
    c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    if EVEN_M and EVEN_N:
        tl.store(c_ptrs, c)
    else:
        tl.store(c_ptrs, c, mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))


def _config_for(M: int, N: int, K: int):
    """BLOCK_M, BLOCK_N, BLOCK_K, num_warps, num_stages, GROUP_M."""
    if M <= 32:
        return 32, 128, 128, 4, 4, 1
    if M <= 64:
        return 64, 128, 128, 4, 4, 4
    return 128, 128, 128, 4, 3, 8


def _launch_gemm(
    x, weight, weight_scale, out,
    M, N, K,
    BLOCK_M, BLOCK_N, BLOCK_K,
    num_warps, num_stages, GROUP_M,
):
    EVEN_M = 1 if (M % BLOCK_M == 0) else 0
    EVEN_N = 1 if (N % BLOCK_N == 0) else 0
    grid = (triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N),)
    _fp8_gemm_kernel[grid](
        x, weight, weight_scale, out,
        M, N, K,
        x.stride(0), x.stride(1),
        weight.stride(0), weight.stride(1),
        out.stride(0), out.stride(1),
        BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
        GROUP_M=GROUP_M, EVEN_M=EVEN_M, EVEN_N=EVEN_N,
        num_warps=num_warps, num_stages=num_stages,
    )


class Model(nn.Module):
    """y = ((x @ w.T) * weight_scale).to(bf16)."""

    def __init__(self, M: int, N: int, K: int):
        super().__init__()
        self.M, self.N, self.K = M, N, K
        w = torch.empty(N, K, dtype=torch.bfloat16)
        nn.init.normal_(w, std=0.02)
        s = (w.float().abs().amax(dim=1, keepdim=True) / E4M3_MAX).clamp(min=1e-12)
        w_fp8 = (w.float() / s).to(torch.float8_e4m3fn)
        self.register_buffer("weight", w_fp8)
        self.register_buffer("weight_scale", s.squeeze(1).to(torch.float32))

        self._BM, self._BN, self._BK, self._nw, self._ns, self._gm = _config_for(M, N, K)
        self._K_work = (K + self._BK - 1) // self._BK * self._BK
        self._need_pad = self._K_work != K

        # Static device workspaces for CUDA-graph capture
        self.register_buffer(
            "_out", torch.empty(M, N, dtype=torch.bfloat16), persistent=False
        )
        self.register_buffer(
            "_x_work",
            torch.zeros(M, self._K_work, dtype=torch.float8_e4m3fn),
            persistent=False,
        )
        self.register_buffer(
            "_w_work",
            torch.zeros(N, self._K_work, dtype=torch.float8_e4m3fn),
            persistent=False,
        )
        self._w_version = -1
        self._graph: torch.cuda.CUDAGraph | None = None
        self._graph_ready = False

    def _sync_w_work(self) -> None:
        """Refresh staged weight when the live buffer is mutated."""
        if self.weight._version == self._w_version:
            return
        if self._need_pad:
            self._w_work[:, : self.K].copy_(self.weight)
        else:
            self._w_work.copy_(self.weight)
        self._w_version = self.weight._version

    def _capture_graph(self) -> None:
        self._sync_w_work()
        for _ in range(3):
            _launch_gemm(
                self._x_work, self._w_work, self.weight_scale, self._out,
                self.M, self.N, self._K_work,
                self._BM, self._BN, self._BK, self._nw, self._ns, self._gm,
            )
        torch.cuda.synchronize()
        g = torch.cuda.CUDAGraph()
        with torch.cuda.graph(g):
            _launch_gemm(
                self._x_work, self._w_work, self.weight_scale, self._out,
                self.M, self.N, self._K_work,
                self._BM, self._BN, self._BK, self._nw, self._ns, self._gm,
            )
        self._graph = g
        self._graph_ready = True

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        if not self._graph_ready:
            self._capture_graph()
        self._sync_w_work()
        if self._need_pad:
            self._x_work[:, : self.K].copy_(x)
        else:
            self._x_work.copy_(x)
        assert self._graph is not None
        self._graph.replay()
        return self._out


M = 4096
N = 4096
K = 4096


def get_inputs():
    x = (torch.rand(M, K) * 8 - 4).to(torch.float8_e4m3fn)
    return [x]


def get_init_inputs():
    return [M, N, K]

20260709_010052_grok_grok-4.5_01_fp8_gemm