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