KernelBench hard · RTX PRO 6000

FP8 GEMM DeepSeek V4.1 Flash

33.1%geomean peak fraction across shapes

manually audited: clean

Hand-written Triton fp8 GEMM: tl.dot on two e4m3 operands lowering to mma.sync.m16n8k32 with an fp32 accumulator, per-output-channel dequant in the epilogue. 578.9 TFLOPS on 4096-cubed is above this part's 500 TFLOPS bf16 dense peak, which is independent physical proof the fp8 pipe is really running rather than a bf16-upcast fallback.

harnessdeepseek-claudeagent session1h 27mtotal wall1h 29mcheck8sbenchmark7soutput tokens239,981cost$14.71gpu-lock wait37mgpu-lock held3mregimecompute

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

4096×4096×40960.228 ms60.2%602 TFLOPS · 60% of 1,000 TF fp8 peak · also 0.29 TB/s (16% of HBM)
4096×4096×41270.254 ms54.6%546 TFLOPS · 55% of 1,000 TF fp8 peak · also 0.27 TB/s (15% of HBM)
32×8192×81920.080 ms5.4%0.85 TB/s · 47% of 1.8 TB/s HBM · also 54 TFLOPS (5% of compute)
4096×14336×40960.705 ms68.3%682 TFLOPS · 68% of 1,000 TF fp8 peak · also 0.27 TB/s (15% of HBM)

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

geomean(60.2% · 54.6% · 5.4% · 68.3%) = 33.1%

Kernel source (redacted)
"""FP8 e4m3 GEMM:  y = (x @ w.T) * weight_scale.

x      : (M, K) float8_e4m3fn
weight : (N, K) float8_e4m3fn, per-output-channel normalised into e4m3 range
scale  : (N,)   float32 per-output-channel dequant scale
y      : (M, N) bfloat16

The kernel runs a genuine fp8 x fp8 tensor-core MMA (the m16n8k32 e4m3 x e4m3
shape, fp32 accumulate) and applies the per-channel dequant scale in the
epilogue, then narrows to bf16.

Two things dominate the achievable rate on SM120 (GB202, 188 SMs):

* The fp8 tensor-core pipe retires 0.25 m16n8k32/clock/SM, i.e. ~890 TFLOPS at
  2.43 GHz -- that is the hard ceiling, not the 1000 TFLOPS headline.  But a
  *real* tile mainloop (cp.async -> ldmatrix -> tensor-core MMA) cannot
  approach it:
  running this exact loop with its operands pinned to two k-tiles (all loads
  L1-resident, same instruction stream, no DRAM traffic) tops out at ~600
  TFLOPS for the 4096^2 grid and ~700 for the 14336-wide one.  So the data
  path, not DRAM, is what caps these shapes -- 569/679 TFLOPS against real
  operand traffic is ~95-97% of that limit, and matches cuBLAS (582 TFLOPS on
  4096^3 measured on this part).
* With cp.async feeding the pipeline, an fp8 operand row must be 16-byte
  aligned or Triton falls back to byte-wise `ld.global.b8` and loses ~3x.  A K
  that is not a multiple of the block width makes every row misaligned, so
  those operands are copied into a K-padded staging buffer first.  The weight
  is constant across calls, so its padded copy is cached.

The skinny decode shape (M=32) is instead memory-bound: its 8192x8192 weight
is 67 MB that must cross DRAM once, and nothing about the GEMM changes that.
A trivially optimal *contiguous* read of that same 67 MB takes 50.3us on this
part (1334 GB/s, i.e. 74% of the 1.8 TB/s spec -- and unchanged whether or not
L2 is flushed); the GEMM, doing that same read plus its output writes, takes
53.4us (1267 GB/s).  So it runs at ~95% of the read ceiling, and thirteen
different tilings (grid 64..512 CTAs, BLOCK_N 32..128, BLOCK_K 128..512, 2..8
warps, 2..3 stages) all land within 53.3-55.6us -- the limit is the DRAM
stream, not the configuration.

Each distinct input/output-buffer combination is also captured into a CUDA
graph: on this GPU a single Triton dispatch is ~13-19us of CPU time, which is
1-20% of the kernel itself across the benchmark shapes.  The graph is only
replayed when the caller has dropped the previous result (tracked with a
weakref); otherwise we fall back to a normal launch into a fresh buffer, so the
module never hands back a tensor whose storage it is about to overwrite.
"""
import weakref

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

E4M3_MAX = 448.0
_GRAPH_MAX = 4


@triton.jit
def _fp8_gemm_kernel(
    a_ptr, b_ptr, c_ptr, s_ptr,
    M, N, K,
    stride_am, stride_bn, stride_cm,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    GROUP_M: tl.constexpr, MASK_M: tl.constexpr, MASK_N: tl.constexpr,
):
    """Both operands are row-major with the K axis contiguous.

    K must be a multiple of BLOCK_K (guaranteed by the host-side pad); the
    caller passes only the row stride, since the K stride is implicitly 1.
    """
    pid = tl.program_id(0)
    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    # Grouped ordering keeps the active A/B panels resident in L2.
    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

    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)
    a_ptrs = a_ptr + rm[:, None] * stride_am + rk[None, :]
    b_ptrs = b_ptr + rn[:, None] * stride_bn + rk[None, :]

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    for _ in range(0, K, BLOCK_K):
        a = tl.load(a_ptrs)
        b = tl.load(b_ptrs)
        acc = tl.dot(a, tl.trans(b), acc)
        a_ptrs += BLOCK_K
        b_ptrs += BLOCK_K

    scale = tl.load(s_ptr + rn)
    c = (acc * scale[None, :]).to(tl.bfloat16)
    if MASK_M or MASK_N:
        cm = rm[:, None] < M
        cn = rn[None, :] < N
        tl.store(c_ptr + rm[:, None] * stride_cm + rn[None, :], c, mask=cm & cn)
    else:
        tl.store(c_ptr + rm[:, None] * stride_cm + rn[None, :], c)


@triton.jit
def _pad_k_kernel(src, dst, K, Kp, NFULL, BLOCK: tl.constexpr):
    """Copy (rows, K) -> (rows, Kp), zero-filling the padded tail columns.

    The bulk of each row goes through a fully unmasked block so the destination
    (whose row stride Kp is 16-byte aligned, unlike K) gets vector accesses.
    """
    row = tl.program_id(0)
    blk = tl.program_id(1)
    offs = blk * BLOCK + tl.arange(0, BLOCK)
    if blk < NFULL:
        v = tl.load(src + row * K + offs)
        tl.store(dst + row * Kp + offs, v)
    else:
        v = tl.load(src + row * K + offs, mask=offs < K, other=0.0)
        # Pad columns [K, Kp) with the zero `other`; the store still needs a
        # mask because BLOCK can overshoot the end of the (Kp-wide) row.
        tl.store(dst + row * Kp + offs, v, mask=offs < Kp)


# Tuned on RTX PRO 6000 (SM120). Compute-bound shapes want a 128x128x64 tile
# with 4 warps (largest per-warp MMA reuse); skinny-M decode shapes want a
# small tile with a deep k-loop so the grid still covers all 188 SMs.
_COMPUTE_CFG = dict(BLOCK_M=128, BLOCK_N=128, BLOCK_K=64, GROUP_M=8,
                    num_warps=4, num_stages=3)
_SKINNY_CFG = dict(BLOCK_M=32, BLOCK_N=64, BLOCK_K=128, GROUP_M=8,
                   num_warps=4, num_stages=3)

_PAD_BLOCK = 1024


def _round_up(v: int, m: int) -> int:
    return ((v + m - 1) // m) * m


def _pick_cfg(M: int, N: int, K: int) -> dict:
    if M <= 64:
        return dict(_SKINNY_CFG)
    return dict(_COMPUTE_CFG)


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)                         # (N, K) fp8
        self.register_buffer("weight_scale", s.squeeze(1).to(torch.float32))  # (N,)
        self._wpad = None
        self._wpad_key = None
        # key -> [graph, static_out, weakref-to-last-return]
        self._graphs: dict = {}
        self._graphs_off = False

    # -- weight padding -----------------------------------------------------

    @staticmethod
    def _pad_rows(t: torch.Tensor, Kp: int) -> torch.Tensor:
        rows, K = t.shape
        out = torch.empty(rows, Kp, dtype=t.dtype, device=t.device)
        grid = (rows, triton.cdiv(Kp, _PAD_BLOCK))
        _pad_k_kernel[grid](t, out, K, Kp, K // _PAD_BLOCK,
                            BLOCK=_PAD_BLOCK, num_warps=4)
        return out

    def _padded_weight(self, Kp: int) -> torch.Tensor:
        """K-padded weight, cached because the weight is constant across calls.

        Keyed on the buffer's version counter so an in-place weight update
        (e.g. the numeric-stress harness scaling `weight` by 1e-2) invalidates
        the cache instead of silently reusing a stale copy.
        """
        w = self.weight
        key = (Kp, w.data_ptr(), w._version, w.device, w.dtype, w.is_contiguous())
        if self._wpad_key != key:
            if Kp == w.shape[1] and w.is_contiguous():
                self._wpad = w
            else:
                self._wpad = self._pad_rows(w.contiguous(), Kp)
            self._wpad_key = key
        return self._wpad

    # -- raw launch ---------------------------------------------------------

    def _run(self, x: torch.Tensor) -> torch.Tensor:
        """One full launch.  Always returns a freshly allocated output."""
        M, K = x.shape
        if not x.is_contiguous():
            x = x.contiguous()
        w = self.weight
        N = w.shape[0]
        cfg = _pick_cfg(M, N, K)
        BK = cfg["BLOCK_K"]

        if K % BK != 0:
            # The row stride must be 16-byte aligned or cp.async degrades to
            # byte-wise loads; also pad to a whole block so the K loop needs no
            # mask.  `contiguous()` is a no-op on the benchmark inputs.
            Kp = _round_up(K, BK)
            x = self._pad_rows(x, Kp)
            w = self._padded_weight(Kp)
            K = Kp
        else:
            w = self._padded_weight(K)

        out = torch.empty((M, N), dtype=torch.bfloat16, device=x.device)
        grid = (triton.cdiv(M, cfg["BLOCK_M"]) * triton.cdiv(N, cfg["BLOCK_N"]),)
        _fp8_gemm_kernel[grid](
            x, w, out, self.weight_scale,
            M, N, K,
            x.stride(0), w.stride(0), out.stride(0),
            BLOCK_M=cfg["BLOCK_M"], BLOCK_N=cfg["BLOCK_N"], BLOCK_K=BK,
            GROUP_M=cfg["GROUP_M"],
            MASK_M=(M % cfg["BLOCK_M"]) != 0,
            MASK_N=(N % cfg["BLOCK_N"]) != 0,
            num_warps=cfg["num_warps"], num_stages=cfg["num_stages"],
        )
        return out

    # -- cuda graph fast path -----------------------------------------------

    def _graph_key(self, x: torch.Tensor):
        # `weight._version` is part of the key because the padded weight for an
        # unaligned K is resolved once at capture time; an in-place weight
        # update (numeric-stress scales it) must force a fresh capture rather
        # than replay a graph still pointing at the previous padded copy.
        w = self.weight
        return (x.data_ptr(), tuple(x.shape), tuple(x.stride()), x.dtype,
                w.data_ptr(), w._version, self.weight_scale.data_ptr(), w.device)

    @staticmethod
    def _held(entry) -> bool:
        """True if the caller still holds the tensor we handed out last time.

        The entry keeps a strong reference to the *base* buffer (the graph
        writes into it), so a weakref to the base would always look alive.
        Instead the caller is given a fresh view object each call and we watch
        that; only the view is ever handed out, so its lifetime tracks exactly
        the caller's reference.
        """
        return entry[2] is not None and entry[2]() is not None

    def _try_capture(self, key, x: torch.Tensor) -> None:
        """Capture `_run` for this exact input buffer, ready for replay.

        The warmup passes run with the ordinary allocator, so the cached padded
        weight lands in normal memory; only the output (and the padded copy of
        a misaligned `x`) is allocated inside the capture, and both are
        referenced solely by the graph.
        """
        if self._graphs_off:
            return
        while len(self._graphs) >= _GRAPH_MAX:
            self._graphs.pop(next(iter(self._graphs)))
        try:
            g = torch.cuda.CUDAGraph()
            side = torch.cuda.Stream()
            side.wait_stream(torch.cuda.current_stream())
            with torch.cuda.stream(side):
                for _ in range(2):
                    self._run(x)
            torch.cuda.current_stream().wait_stream(side)
            with torch.cuda.graph(g):
                static = self._run(x)
        except Exception:
            # A failed capture can leave the stream unusable; stop trying.
            self._graphs_off = True
            return
        self._graphs[key] = [g, static, None]

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        if not (x.is_cuda and x.is_contiguous() and x.dtype == torch.float8_e4m3fn
                and not torch.cuda.is_current_stream_capturing()):
            return self._run(x)

        try:
            key = self._graph_key(x)
            entry = self._graphs.get(key)
            if entry is None:
                out = self._run(x)
                self._try_capture(key, x)
                return out
            if self._held(entry):
                # The caller is still holding the tensor this graph writes
                # into; replay would silently corrupt it.  Take the slow path.
                return self._run(x)
            entry[0].replay()
            base = entry[1]
            out = base.as_strided(base.shape, base.stride(), base.storage_offset())
            entry[2] = weakref.ref(out)
            return out
        except Exception:
            # Graphs are an optimisation only: never let one break the module.
            self._graphs_off = True
            self._graphs.clear()
            return self._run(x)

20260910_202114_deepseek-claude_deepseek-flash_01_fp8_gemm