KernelBench hard · RTX PRO 6000

W4A16 GEMM Claude Fable 5

26.1%geomean peak fraction across shapes

manually audited: clean

Genuine fused W4A16 Triton kernel: int4 nibbles unpacked in-registers inside the GEMM loop via a bf16 bit-trick (OR the nibble into the 0x4300 exponent pattern, bitcast, subtract zeros+128 - bit-identical to the reference's dequant rounding), tensor-core tl.dot with the activation tile deinterleaved by tl.split, split-K over quant groups with fp32 partials + a tiny reduce kernel, weights repacked once into a block-column-major layout (pure permutation, values untouched, rebuilt on any in-place or identity change tracked via tensor._version). Per-shape latency is hidden behind captured CUDA graphs whose first node dereferences a one-entry device pointer table, so every replay stages the caller's CURRENT input buffer - empirically verified to recompute, not replay stale outputs. No forbidden ops, no pre-dequant-to-bf16 + cuBLAS shortcut, no output caching. 0.2608 geomean (best per-shape 0.38, worst 0.135 at M=256 compute-bound) is an honest bandwidth-bound number, nowhere near a lookup signature.

harnessor-fableagent session1h 24mtotal wall1h 24mcheck3sbenchmark3soutput tokensgpu-lock wait28mgpu-lock held5mregimememory

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

1×12288×40960.043 ms34.5%0.62 TB/s · 35% of 1.8 TB/s HBM · also 2 TFLOPS (0% of compute)
32×12288×40960.044 ms34.8%0.63 TB/s · 35% of 1.8 TB/s HBM · also 73 TFLOPS (14% of compute)
256×12288×40960.144 ms13.5%179 TFLOPS · 36% of 500 TF bf16 peak · also 0.24 TB/s (14% of HBM)
1×4096×40960.025 ms19.7%0.36 TB/s · 20% of 1.8 TB/s HBM · also 1 TFLOPS (0% of compute)
16×14336×40960.047 ms37.7%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(34.5% · 34.8% · 13.5% · 19.7% · 37.7%) = 26.1%

Kernel source (redacted)
"""Fused W4A16 weight-only quantized GEMM for RTX PRO 6000 (SM120).

Scheme (AWQ/GPTQ-style asymmetric int4, group_size=128 along K):
  x:      (M, K)          bf16
  w_q:    (K // 2, N)     uint8  (low nibble = even-K row, high nibble = odd-K row)
  scales: (K // 128, N)   bf16
  zeros:  (K // 128, N)   bf16
  out[m, n] = sum_k x[m, k] * (unpack(w_q)[k, n] - zeros[k // 128, n]) * scales[k // 128, n]

Design notes
------------
* The packed int4 stream (25 MB at N=12288, K=4096) is the whole game for the
  decode shapes, so weights are repacked ONCE per weight version into a
  block-column-major layout (N/64, K/2, 64): every program then reads a
  fully contiguous byte range instead of 64-byte rows at a 12 KB stride.
  Scales/zeros get the same treatment. The repack is a pure layout permutation
  of the canonical buffers (which stay untouched for state_dict round-trips)
  and is redone automatically if any weight buffer is modified in place
  (tracked via tensor._version).
* Dequant happens in registers with the reference's exact bf16 rounding:
  nibble minus integer zero-point is exact in bf16, the only rounding is the
  multiply by the bf16 scale. Instead of re-interleaving the weight tile back
  into K order, the activation tile is deinterleaved with tl.split, and each
  quant group contributes dot(x_even, deq(lo)) + dot(x_odd, deq(hi)) on
  tensor cores.
* Kernel launches from Python cost ~10 us of host dispatch, which would
  dominate a ~20 us decode call, so each (M, N, K) gets a captured CUDA
  graph. The input is reached through a one-entry device pointer table: the
  kernel dereferences the table at the start of every replay, so the graph
  always computes from the caller's CURRENT input memory (no staging copy,
  no stale-contents assumptions); the 8-byte entry is rewritten whenever the
  input's device address changes. Weights are read at their stable repacked
  storage locations.
* Split-K over quant groups (fp32 partials + a tiny reduce kernel, no
  atomics) fills all 188 SMs; 64-wide column panels with 4 warps keep
  register pressure low enough for multiple blocks per SM.
"""
from __future__ import annotations

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

OP_TYPE = "gemm_w4a16"
SUPPORTED_PRECISIONS = ["int4_bf16"]
HARDWARE_REQUIRED = ["RTX_PRO_6000", "H100", "B200"]

GROUP_SIZE = 128


def _pick_block_n(N: int) -> int:
    """Column-panel width of the repacked layout. 64-wide panels win across
    the deck: they halve the accumulator tile, which lifts the register-bound
    occupancy ceiling (ncu showed the 128-wide M=256 kernel limited to one
    block/SM at 16.7% occupancy), and they launch enough blocks to cover the
    188 SMs even at narrow N."""
    assert N % 64 == 0
    return 64


@triton.jit
def _stage_x_kernel(xtab_ptr, dst_ptr, NWORDS, BLOCK: tl.constexpr):
    """Copy the live input into the graph's static buffer, 8 bytes/lane.

    The source address is read from a one-entry device pointer table, so a
    captured CUDA graph re-reads the caller's CURRENT input every replay.
    Only this tiny copy pays the runtime-pointer penalty; the GEMM below
    keeps its input as a normal (noalias, aligned) kernel argument.
    """
    pid = tl.program_id(0)
    src = tl.load(xtab_ptr).to(tl.pointer_type(tl.int64))
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < NWORDS
    tl.store(dst_ptr + offs, tl.load(src + offs, mask=mask), mask=mask)


@triton.jit
def _w4a16_kernel(
    x_ptr,        # (M, K) bf16
    wr_ptr,       # (N//BN, K//2, BN) uint8, block-column-major repack
    sr_ptr,       # (N//BN, K//128, BN) bf16
    zr_ptr,       # (N//BN, K//128, BN) bf16
    out_ptr,      # (M, N) bf16, or fp32 workspace when SPLIT_K > 1
    M, N, K,
    N_GROUPS: tl.constexpr,
    GROUPS_PER_SPLIT: tl.constexpr,
    SPLIT_K: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BN: tl.constexpr,        # 64
    GROUP_K: tl.constexpr,   # 128
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    pid_k = tl.program_id(2)

    m_offs = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    n_lane = tl.arange(0, BN)
    m_mask = m_offs < M

    k_offs = tl.arange(0, GROUP_K)
    kh_offs = tl.arange(0, GROUP_K // 2)

    g_lo = pid_k * GROUPS_PER_SPLIT
    g_hi = tl.minimum(g_lo + GROUPS_PER_SPLIT, N_GROUPS)

    w_base = wr_ptr + pid_n * (K // 2) * BN
    s_base = sr_ptr + pid_n * N_GROUPS * BN
    z_base = zr_ptr + pid_n * N_GROUPS * BN

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

    for g in range(g_lo, g_hi):
        k_base = g * GROUP_K
        x_tile = tl.load(
            x_ptr + m_offs[:, None] * K + (k_base + k_offs)[None, :],
            mask=m_mask[:, None], other=0.0,
        )
        x_pair = tl.reshape(x_tile, (BLOCK_M, GROUP_K // 2, 2))
        x_even, x_odd = tl.split(x_pair)

        # contiguous (GROUP_K//2, BN) byte tile; weights/scales/zeros are
        # streamed exactly once, so mark them evict-first to keep them from
        # displacing reusable lines (and to soften the dirty-L2 writeback tax
        # from the benchmark's cache flush).
        p = tl.load(
            w_base + (k_base // 2 + kh_offs)[:, None] * BN + n_lane[None, :],
            eviction_policy="evict_first",
        ).to(tl.uint16)
        s = tl.load(s_base + g * BN + n_lane, eviction_policy="evict_first")
        z = tl.load(z_base + g * BN + n_lane, eviction_policy="evict_first")

        # int4 -> bf16 without an integer convert: 0x4300 | n is the bf16 bit
        # pattern of exactly 128 + n (n in 0..15 sits in the low mantissa
        # bits), so dequant becomes OR+bitcast and a subtract of (zeros+128).
        # 128 + z is exact in bf16, and (128+n) - (128+z) == n - z exactly,
        # so this is bit-identical to the reference's dequant rounding.
        z128 = z + 128.0
        lo = (((p & 0xF) | 0x4300).to(tl.bfloat16, bitcast=True) - z128[None, :]) * s[None, :]
        hi = (((p >> 4) | 0x4300).to(tl.bfloat16, bitcast=True) - z128[None, :]) * s[None, :]

        acc = tl.dot(x_even, lo, acc=acc, out_dtype=tl.float32)
        acc = tl.dot(x_odd, hi, acc=acc, out_dtype=tl.float32)

    if SPLIT_K == 1:
        out_ptrs = out_ptr + m_offs[:, None] * N + (pid_n * BN + n_lane)[None, :]
        tl.store(out_ptrs, acc.to(tl.bfloat16), mask=m_mask[:, None])
    else:
        # fp32 partials at (pid_k, m, n); a tiny reduce kernel sums them.
        # Plain stores avoid atomics and the workspace pre-zero pass.
        part_ptrs = (
            out_ptr + pid_k * M * N
            + m_offs[:, None] * N + (pid_n * BN + n_lane)[None, :]
        )
        tl.store(part_ptrs, acc, mask=m_mask[:, None])


@triton.jit
def _reduce_kernel(part_ptr, out_ptr, MN, SPLIT_K: tl.constexpr, BLOCK: tl.constexpr):
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < MN
    acc = tl.zeros((BLOCK,), dtype=tl.float32)
    for k in tl.static_range(SPLIT_K):
        acc += tl.load(part_ptr + k * MN + offs, mask=mask, other=0.0)
    tl.store(out_ptr + offs, acc.to(tl.bfloat16), mask=mask)


def _pick_config(M: int, N: int, bn: int):
    """(BLOCK_M, SPLIT_K, num_warps, num_stages), tuned on RTX PRO 6000.

    Narrow N (< 96 panels): deep split-K over quant groups to cover the SMs.
    Wide N: shallow split-K (2) keeps every SM fed while the fp32-partials
    epilogue stays cheap; 4 warps with BLOCK_M<=32 keeps register pressure
    low enough for multiple blocks per SM (the occupancy fix for M=256).
    """
    if N // bn < 96:
        return (16 if M <= 16 else 32), 4, 4, 5
    if M <= 1:
        return 16, 1, 4, 4
    if M <= 16:
        return 16, 2, 4, 4
    return 32, 2, 4, 3


class _GraphRunner:
    """CUDA-graph wrapper around the fused kernel for one (M, N, K)."""

    def __init__(self, M, N, K, bn, w_r, s_r, z_r):
        BLOCK_M, SPLIT_K, warps, stages = _pick_config(M, N, bn)
        n_groups = K // GROUP_SIZE
        SPLIT_K = min(SPLIT_K, n_groups)
        groups_per_split = triton.cdiv(n_groups, SPLIT_K)
        device = w_r.device
        # Every buffer the captured graph touches is kept alive on self: the
        # graph records raw pointers, so letting one get GC'd/reused would
        # corrupt unrelated allocations. The input staging copy runs INSIDE
        # the graph and pulls from the caller's live buffer via ptr_table,
        # so each replay recomputes from current input memory while the GEMM
        # keeps a static, fully-optimizable input argument.
        self.ptr_table = torch.zeros(1, dtype=torch.int64, device=device)
        self._last_ptr = -1
        self.static_x = torch.zeros((M, K), dtype=torch.bfloat16, device=device)
        self.static_out = torch.zeros((M, N), dtype=torch.bfloat16, device=device)
        nwords = (M * K * 2) // 8
        assert (M * K * 2) % 8 == 0
        static_x_words = self.static_x.reshape(-1).view(torch.int64)
        stage_block = 1024
        stage_grid = (triton.cdiv(nwords, stage_block),)
        self.ws = (
            torch.zeros((SPLIT_K, M, N), dtype=torch.float32, device=device)
            if SPLIT_K > 1 else None
        )
        self._hold = (w_r, s_r, z_r)

        grid = (triton.cdiv(M, BLOCK_M), N // bn, SPLIT_K)
        red_block = 1024
        red_grid = (triton.cdiv(M * N, red_block),)

        def run():
            _stage_x_kernel[stage_grid](
                self.ptr_table, static_x_words, nwords, BLOCK=stage_block,
                num_warps=4,
            )
            out = self.ws if self.ws is not None else self.static_out
            _w4a16_kernel[grid](
                self.static_x, w_r, s_r, z_r, out, M, N, K,
                N_GROUPS=n_groups,
                GROUPS_PER_SPLIT=groups_per_split,
                SPLIT_K=SPLIT_K,
                BLOCK_M=BLOCK_M, BN=bn, GROUP_K=GROUP_SIZE,
                num_warps=warps, num_stages=stages,
            )
            if self.ws is not None:
                _reduce_kernel[red_grid](
                    self.ws, self.static_out, M * N,
                    SPLIT_K=SPLIT_K, BLOCK=red_block, num_warps=4,
                )

        # Warm up (Triton JIT) outside capture, then capture the launch chain.
        # A scratch input backs the table during warmup/capture; real calls
        # publish the caller's input address before replaying.
        warm_x = torch.zeros((M, K), dtype=torch.bfloat16, device=device)
        self.ptr_table.fill_(warm_x.data_ptr())
        stream = torch.cuda.Stream()
        stream.wait_stream(torch.cuda.current_stream())
        with torch.cuda.stream(stream):
            run()
            run()
        torch.cuda.current_stream().wait_stream(stream)
        self.graph = torch.cuda.CUDAGraph()
        with torch.cuda.graph(self.graph):
            run()

    def __call__(self, x: torch.Tensor) -> torch.Tensor:
        p = x.data_ptr()
        if p != self._last_ptr:
            if p % 8:
                # stage kernel copies 8-byte words; realign (rare: bf16
                # tensors are normally at least 8-byte aligned)
                x = x.clone()
                p = x.data_ptr()
            self.ptr_table.fill_(p)
            self._last_ptr = p
        self.graph.replay()
        return self.static_out


class Model(nn.Module):
    """W4A16 GEMM: y = x @ dequant(w_q, scales, zeros), fused unpack+GEMM."""

    def __init__(self, M: int, N: int, K: int, group_size: int = GROUP_SIZE):
        super().__init__()
        assert K % group_size == 0
        assert K % 2 == 0
        self.M, self.N, self.K = M, N, K
        self.group_size = group_size
        n_groups = K // group_size
        # Same buffer names/shapes/dtypes as the reference; real values arrive
        # via load_state_dict.
        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))
        self._bn = _pick_block_n(N)
        self._runners: dict[int, _GraphRunner] = {}
        self._packed = None       # (w_r, s_r, z_r)
        self._packed_versions = None

    def _buffer_versions(self):
        # Tensor identity catches buffer replacement (.to(), reassignment);
        # _version catches in-place mutation (copy_, zero_, ...). Compared
        # field-by-field in __call__ -- never with tuple ==, which would
        # trigger elementwise tensor comparison.
        return (
            self.w_q, self.scales, self.zeros,
            self.w_q._version, self.scales._version, self.zeros._version,
        )

    def _repack(self):
        """Layout permutation: (K//2, N) -> (N//BN, K//2, BN) contiguous.

        Pure data movement -- values are untouched. Rebuilt from the live
        canonical buffers whenever they change, so the kernel always computes
        from the current weights.
        """
        K, N = self.K, self.N
        bn = self._bn
        n_groups = K // self.group_size
        w_r = (
            self.w_q.reshape(K // 2, N // bn, bn)
            .permute(1, 0, 2).contiguous()
        )
        s_r = (
            self.scales.reshape(n_groups, N // bn, bn)
            .permute(1, 0, 2).contiguous()
        )
        z_r = (
            self.zeros.reshape(n_groups, N // bn, bn)
            .permute(1, 0, 2).contiguous()
        )
        self._packed = (w_r, s_r, z_r)
        self._packed_versions = self._buffer_versions()
        self._runners = {}

    def __call__(self, x: torch.Tensor) -> torch.Tensor:
        # Overrides nn.Module.__call__ to skip the hook-dispatch machinery:
        # at ~20 us per GEMM the Python between call and kernel launch is
        # measurable. No hooks are used by this model. Every call recomputes
        # the GEMM from the live buffers via graph replay.
        pk = self._packed_versions
        if (
            pk is None
            or pk[0] is not self.w_q
            or pk[1] is not self.scales
            or pk[2] is not self.zeros
            or pk[3] != self.w_q._version
            or pk[4] != self.scales._version
            or pk[5] != self.zeros._version
        ):
            self._repack()
        if x.dtype is not torch.bfloat16 or not x.is_contiguous():
            x = x.to(torch.bfloat16).contiguous()
        runner = self._runners.get(x.shape[0])
        if runner is None:
            runner = _GraphRunner(x.shape[0], self.N, self.K, self._bn, *self._packed)
            self._runners[x.shape[0]] = runner
        return runner(x)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.__call__(x)


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]

20260719_070614_or-fable_anthropic_claude-fable-5_07_w4a16_gemm