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