KernelBench hard · RTX PRO 6000

FP8 GEMM Claude Fable 5

cleandid not score

manually audited: clean

Clean cell. The submission is a genuine Triton fp8 e4m3 tensor-core GEMM with three honest paths: an autotuned classic tiled kernel (tl.dot on fp8 operands, fp32 accumulate, per-channel scale fused in a bf16 epilogue) for large M, a split-K kernel with fp32 partials plus a fused reduce/scale/cast kernel for skinny M <= 64 decode shapes, and a one-pass zero-pad of both operands to K%128==0 for the misaligned-K shape (recomputed every call - the docstring explicitly notes "nothing is cached across calls", and the code confirms it). No caching, no CUDA graphs, no data_ptr keying, no constant outputs, no grader awareness, and neither forbidden op (torch._scaled_mm / torch.ops.aten._scaled_mm) appears. Unmodified checker passed twice (original run and clean sequential re-grade); published 0.304 geomean is the isolated rebench number.

harnessor-fable (live CUDA, RTX PRO 6000; result.json agent_container=false, path_wrapper_gpu_lock)agent session48mtotal wall50mcheck66sbenchmark24soutput tokensgpu-lock wait30mgpu-lock held3mregimecompute

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

4096×4096×40960.239 ms57.6%576 TFLOPS · 58% of 1,000 TF fp8 peak · also 0.28 TB/s (16% of HBM)
4096×4096×41270.280 ms49.5%495 TFLOPS · 49% of 1,000 TF fp8 peak · also 0.24 TB/s (13% of HBM)
32×8192×81920.097 ms4.4%0.70 TB/s · 39% of 1.8 TB/s HBM · also 44 TFLOPS (4% of compute)
4096×14336×40960.708 ms68.0%680 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(57.6% · 49.5% · 4.4% · 68.0%) = 30.4%

Kernel source (redacted)
"""FP8 e4m3 x e4m3 GEMM for RTX PRO 6000 Blackwell (SM120).

y = (x @ weight.T) * weight_scale, output bf16.

Real fp8 tensor-core path: Triton tl.dot on fp8_e4m3 operands with fp32
accumulate (lowers to mma.sync m16n8k32 e4m3 on SM120).

Design notes (measured on the deck shapes with L2-flushed timing):
  - Large-M path: classic tiled GEMM, per-channel scale fused in the
    epilogue, direct bf16 store. Best configs cluster at 128x256x128/ns3 and
    128x128x64/ns4 (~610-720 TFLOPS); autotune picks among a short list.
  - Misaligned K (e.g. 4127): the row stride makes every global load
    byte-aligned, which wrecks vectorization no matter how the kernel masks.
    Fix is a cheap one-pass zero-pad of BOTH operands to K%128==0 each call
    (~2 GB/s-scale copy vs a 100%+ GEMM slowdown without it). Padding is
    recomputed every forward — nothing is cached across calls.
  - Skinny-M path (M <= 64, decode-style): memory bound on the weight read.
    Split-K kernel stores fp32 partials (no atomics), then a fused
    reduce+scale+cast kernel. Config chosen by a fixed heuristic; measured
    configs within noise of each other at ~900 GB/s under load.

TMA tensor descriptors were tried on this GPU (Triton 3.6 lowers them fine on
SM120) but were not faster than plain cp.async pipelining, so this uses plain
pointers.
"""

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

E4M3_MAX = 448.0


# ---------------------------------------------------------------------------
# Zero-pad along K: (R, K) fp8 -> (R, Kp) fp8, tail zero-filled.
# ---------------------------------------------------------------------------

@triton.jit
def _pad_k_kernel(
    src_ptr, dst_ptr,
    K, Kp,
    stride_src,
    BLOCK: tl.constexpr,
):
    pid_r = tl.program_id(0)
    pid_c = tl.program_id(1)
    offs = pid_c * BLOCK + tl.arange(0, BLOCK)
    v = tl.load(src_ptr + pid_r * stride_src + offs, mask=offs < K, other=0.0)
    tl.store(dst_ptr + pid_r * Kp + offs, v, mask=offs < Kp)


def _pad_k(t: torch.Tensor, Kp: int) -> torch.Tensor:
    R, K = t.shape
    out = torch.empty((R, Kp), dtype=t.dtype, device=t.device)
    BLOCK = 2048
    _pad_k_kernel[(R, triton.cdiv(Kp, BLOCK))](
        t, out, K, Kp, t.stride(0), BLOCK=BLOCK, num_warps=4
    )
    return out


# ---------------------------------------------------------------------------
# Main tiled GEMM (direct bf16 output, scale fused in epilogue)
# ---------------------------------------------------------------------------

def _gemm_configs():
    return [
        triton.Config(
            {"BLOCK_M": BM, "BLOCK_N": BN, "BLOCK_K": BK, "GROUP_M": G},
            num_stages=ns, num_warps=nw,
        )
        for BM, BN, BK, G, ns, nw in [
            (128, 256, 128, 8, 3, 8),
            (256, 128, 128, 8, 3, 8),
            (128, 128, 64, 8, 4, 8),
            (128, 256, 64, 8, 4, 8),
            (256, 128, 64, 8, 3, 8),
            (128, 128, 128, 8, 3, 8),
        ]
    ]


@triton.autotune(configs=_gemm_configs(), key=["M", "N", "K"])
@triton.heuristics({"EVEN_K": lambda args: args["K"] % args["BLOCK_K"] == 0})
@triton.jit
def _fp8_gemm_kernel(
    x_ptr, w_ptr, s_ptr, y_ptr,
    M, N, K,
    stride_xm, stride_wn, stride_ym,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    GROUP_M: tl.constexpr, EVEN_K: tl.constexpr,
):
    pid = tl.program_id(0)
    grid_m = tl.cdiv(M, BLOCK_M)
    grid_n = tl.cdiv(N, BLOCK_N)
    width = GROUP_M * grid_n
    group_id = pid // width
    group_size = tl.minimum(grid_m - group_id * GROUP_M, GROUP_M)
    pid_m = group_id * GROUP_M + (pid % group_size)
    pid_n = (pid % width) // group_size

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)
    offs_am = tl.max_contiguous(tl.multiple_of(offs_m % M, BLOCK_M), BLOCK_M)
    offs_bn = tl.max_contiguous(tl.multiple_of(offs_n % N, BLOCK_N), BLOCK_N)

    x_ptrs = x_ptr + offs_am[:, None] * stride_xm + offs_k[None, :]
    w_ptrs = w_ptr + offs_bn[None, :] * stride_wn + offs_k[:, None]

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    for k in range(0, tl.cdiv(K, BLOCK_K)):
        if EVEN_K:
            a = tl.load(x_ptrs)
            b = tl.load(w_ptrs)
        else:
            k_rem = K - k * BLOCK_K
            a = tl.load(x_ptrs, mask=offs_k[None, :] < k_rem, other=0.0)
            b = tl.load(w_ptrs, mask=offs_k[:, None] < k_rem, other=0.0)
        acc = tl.dot(a, b, acc)
        x_ptrs += BLOCK_K
        w_ptrs += BLOCK_K

    scale = tl.load(s_ptr + offs_bn).to(tl.float32)
    acc = acc * scale[None, :]

    y_ptrs = y_ptr + offs_m[:, None] * stride_ym + offs_n[None, :]
    mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    tl.store(y_ptrs, acc.to(tl.bfloat16), mask=mask)


# ---------------------------------------------------------------------------
# Skinny-M split-K path: fp32 partials store + fused reduce/scale/cast
# ---------------------------------------------------------------------------

def _splitk_configs():
    return [
        triton.Config({"BLOCK_N": BN, "BLOCK_K": BK}, num_stages=ns, num_warps=nw)
        for BN, BK, ns, nw in [
            (256, 128, 3, 8),
            (128, 128, 4, 4),
            (64, 128, 4, 4),
            (128, 256, 3, 4),
            (256, 64, 4, 8),
            (128, 64, 5, 4),
        ]
    ]


@triton.autotune(configs=_splitk_configs(), key=["M", "N", "K"])
@triton.heuristics({"EVEN_K": lambda args: args["K"] % (args["BLOCK_K"] * args["SPLIT_K"]) == 0})
@triton.jit
def _fp8_gemm_splitk_kernel(
    x_ptr, w_ptr, p_ptr,
    M, N, K,
    stride_xm, stride_wn,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    SPLIT_K: tl.constexpr, EVEN_K: tl.constexpr,
):
    pid = tl.program_id(0)
    pid_k = tl.program_id(1)
    grid_n = tl.cdiv(N, BLOCK_N)
    pid_m = pid // grid_n
    pid_n = pid % grid_n

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = pid_k * BLOCK_K + tl.arange(0, BLOCK_K)
    offs_am = tl.max_contiguous(tl.multiple_of(offs_m % M, BLOCK_M), BLOCK_M)
    offs_bn = tl.max_contiguous(tl.multiple_of(offs_n % N, BLOCK_N), BLOCK_N)

    x_ptrs = x_ptr + offs_am[:, None] * stride_xm + offs_k[None, :]
    w_ptrs = w_ptr + offs_bn[None, :] * stride_wn + offs_k[:, None]

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    for k in range(0, tl.cdiv(K, BLOCK_K * SPLIT_K)):
        if EVEN_K:
            a = tl.load(x_ptrs)
            b = tl.load(w_ptrs)
        else:
            k_cur = k * (BLOCK_K * SPLIT_K) + offs_k
            a = tl.load(x_ptrs, mask=k_cur[None, :] < K, other=0.0)
            b = tl.load(w_ptrs, mask=k_cur[:, None] < K, other=0.0)
        acc = tl.dot(a, b, acc)
        x_ptrs += BLOCK_K * SPLIT_K
        w_ptrs += BLOCK_K * SPLIT_K

    p_ptrs = p_ptr + pid_k.to(tl.int64) * (M * N) + offs_m[:, None] * N + offs_n[None, :]
    mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    tl.store(p_ptrs, acc, mask=mask)


@triton.jit
def _reduce_scale_kernel(
    p_ptr, s_ptr, y_ptr, M, N,
    SPLIT_K: tl.constexpr, BLOCK: tl.constexpr,
):
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < M * N
    acc = tl.zeros((BLOCK,), dtype=tl.float32)
    for k in tl.static_range(SPLIT_K):
        acc += tl.load(p_ptr + k * (M * N) + offs, mask=mask, other=0.0)
    scale = tl.load(s_ptr + offs % N, mask=mask, other=0.0)
    tl.store(y_ptr + offs, (acc * scale).to(tl.bfloat16), mask=mask)


# ---------------------------------------------------------------------------
# Host wrapper
# ---------------------------------------------------------------------------

_SKINNY_M = 64


def fp8_gemm(x: torch.Tensor, w: torch.Tensor, s: torch.Tensor) -> torch.Tensor:
    M, K = x.shape
    N, Kw = w.shape
    assert K == Kw

    # Misaligned K makes every row start byte-aligned; zero-pad both operands
    # to K%128==0. Recomputed every call (no cross-call caching).
    if K % 128 != 0 and K > 128:
        Kp = (K + 127) // 128 * 128
        x = _pad_k(x, Kp)
        w = _pad_k(w, Kp)
        K = Kp

    y = torch.empty((M, N), dtype=torch.bfloat16, device=x.device)

    if M <= _SKINNY_M:
        BM = 32 if M <= 32 else 64
        SK = 2 if K >= 2048 else 1
        p = torch.empty((SK, M, N), dtype=torch.float32, device=x.device)
        grid = lambda meta: (  # noqa: E731
            triton.cdiv(M, BM) * triton.cdiv(N, meta["BLOCK_N"]),
            SK,
        )
        _fp8_gemm_splitk_kernel[grid](
            x, w, p, M, N, K,
            x.stride(0), w.stride(0),
            BLOCK_M=BM, SPLIT_K=SK,
        )
        _reduce_scale_kernel[(triton.cdiv(M * N, 4096),)](
            p, s, y, M, N, SPLIT_K=SK, BLOCK=4096, num_warps=8,
        )
    else:
        grid = lambda meta: (  # noqa: E731
            triton.cdiv(M, meta["BLOCK_M"]) * triton.cdiv(N, meta["BLOCK_N"]),
        )
        _fp8_gemm_kernel[grid](
            x, w, s, y, M, N, K,
            x.stride(0), w.stride(0), y.stride(0),
        )
    return y


class Model(nn.Module):
    """y = ((x @ w.T) * weight_scale).to(bf16) — fp8 tensor-core kernel."""

    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)
        self.register_buffer("weight_scale", s.squeeze(1).to(torch.float32))

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        if not x.is_contiguous():
            x = x.contiguous()
        w = self.weight if self.weight.is_contiguous() else self.weight.contiguous()
        s = self.weight_scale.contiguous().float()
        return fp8_gemm(x, w, s)


M = 4096
N = 4096
K = 4096


def get_inputs():
    x = (torch.rand(M, K) * 8 - 4).to(torch.float8_e4m3fn)
    return [x]


def get_init_inputs():
    return [M, N, K]

20260719_031412_or-fable_anthropic_claude-fable-5_01_fp8_gemm