KernelBench hard · RTX PRO 6000
W4A16 GEMM DeepSeek V4.1 Flash
24.9%geomean peak fraction across shapes
manually audited: clean
Real fused W4A16 Triton kernel. Packed uint8 is read inside the K loop and both nibble planes are widened, zero-shifted and scaled in registers, so no dequantised weight ever reaches memory and no pre-dequant cache exists. The scale is folded into the weight operand so one fp32 accumulator spans the whole K, measured at 2.9x the two-accumulator form.
harnessdeepseek-claudeagent session2h 13mtotal wall2h 13mcheck5sbenchmark3soutput tokens302,459cost$18.56gpu-lock wait1h 13mgpu-lock held8mregimememory
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
1×12288×40960.040 ms37.1%0.67 TB/s · 37% of 1.8 TB/s HBM · also 3 TFLOPS (1% of compute)
32×12288×40960.046 ms33.9%0.61 TB/s · 34% of 1.8 TB/s HBM · also 71 TFLOPS (14% of compute)
256×12288×40960.163 ms12.0%158 TFLOPS · 32% of 500 TF bf16 peak · also 0.22 TB/s (12% of HBM)
1×4096×40960.029 ms16.8%0.30 TB/s · 17% of 1.8 TB/s HBM · also 1 TFLOPS (0% of compute)
16×14336×40960.047 ms37.8%0.68 TB/s · 38% of 1.8 TB/s HBM · also 40 TFLOPS (8% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(37.1% · 33.9% · 12.0% · 16.8% · 37.8%) = 24.9%
Kernel source (redacted)
"""Fused W4A16 (int4 weight-only) GEMM for RTX PRO 6000 (SM120 Blackwell).
Math
----
The reference dequantises per 128-wide group along K and then does a bf16
matmul:
w_bf[k, n] = (q[k, n] - z[k // G, n]) * s[k // G, n], q in [0, 15]
We never materialise ``w_bf``. Instead the *dequantised weight fragment is
built inside the GEMM loop*, straight from the packed int4 byte stream:
out[m, n] = sum_k x[m, k] * (((q[k, n] - z_g[n]) * s_g[n]))
Because ``(q - z)`` is an integer in [-15, 15] and ``s`` is bf16, the product
``(q - z) * s`` is *bit-identical* to what the reference computes for the
dequantised weight (both round a bf16 operand product to bf16). The only
difference from the reference is accumulation order, so the kernel lands
within a couple of bf16 ulps of ``x @ w_bf``.
Two things make this cheap:
* the scale is folded into the weight operand, so the MMA accumulator runs
over the whole K loop with no per-group reset and no second accumulator
tile -- that register saving is what lets the tensor-core path stay
resident (measured 2.9x faster than the two-accumulator formulation);
* the two nibble planes are de-interleaved from a *contiguous* x load with
``reshape``+``split`` instead of two stride-2 gathers, which keeps every
global load vectorisable.
Layout of one k-iteration (BK == group size == 128):
packed tile (KH=BK/2, BN) uint8 -> low nibbles are k-even, high nibbles
k-odd, so the low plane pairs with x[:, 0::2] and the high plane with
x[:, 1::2]. DRAM therefore sees exactly (K/2)*N weight bytes per call --
the int4 stream, nothing else.
Launch
------
A ~40 us fixed GPU dispatch latency dominates every shape here (measured with
a null kernel through the harness' own event timing), and a Triton python
launch adds ~10 us of host time on top. The forward path therefore records
the launch into a CUDA graph the first time it sees a given
(x pointer, weight pointer, shape) and replays it afterwards, falling back
to an ordinary launch if capture is unavailable. A replay is ~5 us cheaper
and, more importantly, removes the host-side gap between the timing event and
the kernel.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import triton
import triton.language as tl
GROUP_SIZE = 128
# (max M, BM, BN, num_warps, num_stages). BN=128 with 8 warps wins on every
# M, and the decode tier (M<=2) wants the deepest prefetch queue because it has
# almost no arithmetic to hide DRAM latency behind.
_TIERS = (
(2, (16, 128, 8, 5)),
(16, (16, 128, 8, 4)),
(32, (32, 128, 8, 4)),
(256, (64, 128, 8, 4)),
(None, (128, 128, 8, 3)),
)
def _pick_config(M: int, N: int):
for limit, cfg in _TIERS:
if limit is None or M <= limit:
BM, BN, nw, ns = cfg
break
# A 128-wide n-tile needs N >= 8192 to hand out at least 64 blocks; below
# that the part is block-starved (4096 wide gives only 32 blocks for 188
# SMs) and halving BN to 64, i.e. doubling the block count, is worth ~28%
# on the square decode shape even though it costs a narrower weight row.
while BN > 32 and BN > N // 64:
BN //= 2
return BM, BN, nw, ns
@triton.jit
def _w4a16_kernel(
x_ptr, wq_ptr, sc_ptr, zr_ptr, out_ptr,
M, N, K,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
):
KH: tl.constexpr = BK // 2 # packed rows covered by one k-group
pid_n = tl.program_id(0)
pid_m = tl.program_id(1)
offs_m = pid_m * BM + tl.arange(0, BM)
offs_n = pid_n * BN + tl.arange(0, BN)
offs_kh = tl.arange(0, KH)
offs_k = tl.arange(0, BK)
m_mask = offs_m < M
n_mask = offs_n < N
tot = tl.zeros((BM, BN), dtype=tl.float32)
x_row = x_ptr + offs_m[:, None] * K
w_row = wq_ptr + offs_n[None, :]
nrow = offs_kh[:, None] * N
for _g in range(0, K // BK):
# Contiguous x tile, de-interleaved into even/odd k planes.
xt = tl.load(x_row + offs_k[None, :], mask=m_mask[:, None], other=0.0)
x0, x1 = tl.split(tl.reshape(xt, (BM, KH, 2)))
wq = tl.load(w_row + nrow, mask=n_mask[None, :], other=0)
zbf = tl.load(zr_ptr + offs_n, mask=n_mask, other=0.0)
s = tl.load(sc_ptr + offs_n, mask=n_mask, other=0.0)
# (q - z) is exact in bf16; the product rounds exactly like the
# reference's bf16 dequant.
wlo = ((wq & 15).to(tl.bfloat16) - zbf[None, :]) * s[None, :]
whi = ((wq >> 4).to(tl.bfloat16) - zbf[None, :]) * s[None, :]
tot = tl.dot(x0, wlo, tot)
tot = tl.dot(x1, whi, tot)
x_row += BK
w_row += KH * N
zr_ptr += N
sc_ptr += N
om = out_ptr + offs_m[:, None] * N + offs_n[None, :]
tl.store(om, tot.to(tl.bfloat16), mask=m_mask[:, None] & n_mask[None, :])
# CUDA-graph replay cache. Keyed on the operand pointers *and* the shape, so
# a hit guarantees the recorded kernel reads exactly what this call would.
_STATE: dict = {"key": None, "graph": None, "out": None, "pool": None, "ok": True}
def _launch(x, w_q, scales, zeros, out, M, N, K, BM, BN, nw, ns):
grid = (triton.cdiv(N, BN), triton.cdiv(M, BM))
_w4a16_kernel[grid](
x, w_q, scales, zeros, out, M, N, K,
BM=BM, BN=BN, BK=GROUP_SIZE,
num_warps=nw, num_stages=ns,
)
class Model(nn.Module):
def __init__(self, M: int, N: int, K: int, group_size: int = GROUP_SIZE):
super().__init__()
assert K % group_size == 0, "K must be divisible by group_size"
assert K % 2 == 0, "K must be even (int4 packing)"
self.M, self.N, self.K = M, N, K
self.group_size = group_size
n_groups = K // group_size
self.register_buffer("w_q", torch.zeros(K // 2, N, dtype=torch.uint8))
self.register_buffer("scales", torch.zeros(n_groups, N, dtype=torch.bfloat16))
self.register_buffer("zeros", torch.zeros(n_groups, N, dtype=torch.bfloat16))
def forward(self, x: torch.Tensor) -> torch.Tensor:
if x.dtype != torch.bfloat16:
x = x.to(torch.bfloat16)
if not x.is_contiguous():
x = x.contiguous()
M, K = x.shape
N = self.N
BM, BN, nw, ns = _pick_config(M, N)
key = (x.data_ptr(), self.w_q.data_ptr(), M, N, K)
st = _STATE
if st["ok"] and st["key"] == key:
st["graph"].replay()
return st["out"]
out = torch.empty((M, N), dtype=torch.bfloat16, device=x.device)
_launch(x, self.w_q, self.scales, self.zeros, out, M, N, K, BM, BN, nw, ns)
if st["ok"]:
try:
if st["pool"] is None:
st["pool"] = torch.cuda.graphs.graph_pool_handle()
g = torch.cuda.CUDAGraph()
# Warm capture: the eager launch above already compiled the
# kernel and claimed any global scratch it needs. Capture
# itself does not execute, so `out` still holds a valid result.
with torch.cuda.graph(g, pool=st["pool"]):
_launch(x, self.w_q, self.scales, self.zeros, out,
M, N, K, BM, BN, nw, ns)
st["key"], st["graph"], st["out"] = key, g, out
except Exception:
# Any capture failure (unsupported driver, stream conflict)
# is permanent for this process: stay on the eager path.
st["ok"] = False
st["key"], st["graph"], st["out"] = None, None, None
return out
M = 1
N = 12288
K = 4096
def get_inputs():
x = torch.randn(M, K, dtype=torch.bfloat16)
return [x]
def get_init_inputs():
return [M, N, K]
20260910_202139_deepseek-claude_deepseek-flash_07_w4a16_gemm