KernelBench hard · H100

KDA CUTLASS Qwen 3.8 Max

1.68%geomean peak fraction across shapes

manually audited: clean

Clean pass. Real custom three-kernel Triton KDA chunk-form implementation, independently written in-session (FLA sources consulted read-only under explicit PROMPT permission, nothing copied from prior runs), passing the official stress-enabled checker run by the harness under an exclusive H100 PCIe lock, with the graded benchmark measured sequentially in the isolated per-GPU queue. No reward hacking, no template mutation, no forbidden ops, no pointer-keyed caching or CUDA-graph replay risk, and no contamination. peak_fraction 0.0168 is low in roofline terms but is the honest graded metric and is publishable for this cell.

harnessor-fableagent session1h 28mtotal wall1h 28mcheck27sbenchmark6soutput tokens400,080gpu-lock wait0sgpu-lock held33sregimecompute

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

2×1024×8×128×128×640.157 ms1.8%0.16 TB/s · 8% of 2.0 TB/s HBM · also 14 TFLOPS (2% of compute)
2×2048×8×128×128×640.230 ms2.5%0.22 TB/s · 11% of 2.0 TB/s HBM · also 19 TFLOPS (2% of compute)
1×4096×8×128×128×640.284 ms2.0%0.18 TB/s · 9% of 2.0 TB/s HBM · also 15 TFLOPS (2% of compute)
1×2048×4×128×128×640.161 ms0.9%0.08 TB/s · 4% of 2.0 TB/s HBM · also 7 TFLOPS (1% of compute)

compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)

geomean(1.8% · 2.5% · 2.0% · 0.9%) = 1.7%

Kernel source (redacted)
"""Kimi Delta Attention (KDA) forward, chunk form -- custom Triton kernels for SM90.

Two-kernel chunk-parallel implementation:

  Kernel 1 ("prep"), one program per chunk (parallel over B*H*NT):
    - in-chunk cumsum of the fp32 log-decay g (converted to base-2)
    - Aqk = tril( (q*scale*exp2(G)) @ (k*exp2(-G))^T )
    - L   = strict_tril( -beta_row * (k*exp2(G)) @ (k*exp2(-G))^T )
    - X = (I - L)^-1 as the finite Neumann series sum_{j<64} L^j, built with
      repeated squaring:  prod_{i=0..5} (I + L^{2^i})  (L^64 = 0)
    - w = X @ (beta * k * exp2(G)),  u = X @ (beta * v)
    - stores qg = q*scale*exp2(G), kg = k*exp2(G_last - G), G_last
  Kernel 2 ("scan"), one program per (batch-head, value-block), sequential
  over chunks, state S kept in registers:
      v_n = u - w @ S
      o   = qg @ S + Aqk @ v_n
      S   = exp2(G_last) * S + kg^T @ v_n

Matches the math of the reference chunk formulation (see reference.py).
"""
from __future__ import annotations

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

RCP_LN2 = tl.constexpr(1.4426950408889634)


# ---------------------------------------------------------------------------
# Kernel 1: per-chunk preparation (fully parallel)
# ---------------------------------------------------------------------------
@triton.jit(do_not_specialize=["T", "NT"])
def _kda_prep_kernel(
    q_ptr, k_ptr, g_ptr, beta_ptr,
    a_ptr, qg_ptr, kg_ptr, gl_ptr, l_ptr, f_ptr,
    scale,
    T, NT,
    H: tl.constexpr,
    K: tl.constexpr,
    BT: tl.constexpr,
):
    """Streaming phase: reads q,k,g,beta once; writes qg, kg, gl, Aqk, L, f.

    f = beta * k * exp2(G) so the compute kernel can form w = X @ f without
    touching q/k/g again.
    """
    i_n = tl.program_id(0)
    i_bh = tl.program_id(1)
    i_b = i_bh // H
    i_h = i_bh % H

    BK: tl.constexpr = K // 2                          # 64 for K=128
    r64 = tl.arange(0, BT)
    tok = i_b * T + i_n * BT + r64                     # (BT,) absolute tokens
    row_off = ((tok * H + i_h) * K)                    # (BT,) row offsets
    slab = (i_bh * NT + i_n) * BT
    cols0 = tl.arange(0, BK)
    cols1 = BK + tl.arange(0, BK)
    b_beta = tl.load(beta_ptr + tok * H + i_h).to(tl.float32)

    acc_a = tl.zeros([BT, BT], dtype=tl.float32)
    acc_l = tl.zeros([BT, BT], dtype=tl.float32)

    # ---- K slice 0 ----
    b_g = tl.load(g_ptr + row_off[:, None] + cols0[None, :])
    G = tl.cumsum(b_g, 0) * RCP_LN2
    b_q = tl.load(q_ptr + row_off[:, None] + cols0[None, :]).to(tl.float32)
    b_k = tl.load(k_ptr + row_off[:, None] + cols0[None, :]).to(tl.float32)
    e_pos = tl.exp2(G)
    e_neg = tl.exp2(-G)
    qg0 = (b_q * e_pos * scale).to(tl.bfloat16)
    kp0 = (b_k * e_pos).to(tl.bfloat16)
    kn0 = (b_k * e_neg).to(tl.bfloat16)
    f0 = (b_beta[:, None] * b_k * e_pos).to(tl.bfloat16)
    tl.store(qg_ptr + (slab + r64[:, None]) * K + cols0[None, :], qg0)
    tl.store(f_ptr + (slab + r64[:, None]) * K + cols0[None, :], f0)
    gl0 = tl.sum(tl.where(r64[:, None] == BT - 1, G, 0.0), 0)
    egl0 = tl.exp2(gl0)
    tl.store(gl_ptr + (i_bh * NT + i_n) * K + cols0, gl0)
    tl.store(kg_ptr + (slab + r64[:, None]) * K + cols0[None, :],
             (kn0 * egl0[None, :]).to(tl.bfloat16))
    p_t0 = ((tok[None, :] * H + i_h) * K) + cols0[:, None]
    b_gT = tl.load(g_ptr + p_t0)
    GT = tl.cumsum(b_gT, 1) * RCP_LN2
    b_kT = tl.load(k_ptr + p_t0).to(tl.float32)
    knT = (b_kT * tl.exp2(-GT)).to(tl.bfloat16)
    acc_a += tl.dot(qg0, knT)
    acc_l += tl.dot(kp0, knT)

    # ---- K slice 1 ----
    b_g = tl.load(g_ptr + row_off[:, None] + cols1[None, :])
    G = tl.cumsum(b_g, 0) * RCP_LN2
    b_q = tl.load(q_ptr + row_off[:, None] + cols1[None, :]).to(tl.float32)
    b_k = tl.load(k_ptr + row_off[:, None] + cols1[None, :]).to(tl.float32)
    e_pos = tl.exp2(G)
    e_neg = tl.exp2(-G)
    qg1 = (b_q * e_pos * scale).to(tl.bfloat16)
    kp1 = (b_k * e_pos).to(tl.bfloat16)
    kn1 = (b_k * e_neg).to(tl.bfloat16)
    f1 = (b_beta[:, None] * b_k * e_pos).to(tl.bfloat16)
    tl.store(qg_ptr + (slab + r64[:, None]) * K + cols1[None, :], qg1)
    tl.store(f_ptr + (slab + r64[:, None]) * K + cols1[None, :], f1)
    gl1 = tl.sum(tl.where(r64[:, None] == BT - 1, G, 0.0), 0)
    egl1 = tl.exp2(gl1)
    tl.store(gl_ptr + (i_bh * NT + i_n) * K + cols1, gl1)
    tl.store(kg_ptr + (slab + r64[:, None]) * K + cols1[None, :],
             (kn1 * egl1[None, :]).to(tl.bfloat16))
    p_t1 = ((tok[None, :] * H + i_h) * K) + cols1[:, None]
    b_gT = tl.load(g_ptr + p_t1)
    GT = tl.cumsum(b_gT, 1) * RCP_LN2
    b_kT = tl.load(k_ptr + p_t1).to(tl.float32)
    knT = (b_kT * tl.exp2(-GT)).to(tl.bfloat16)
    acc_a += tl.dot(qg1, knT)
    acc_l += tl.dot(kp1, knT)

    # ---- mask and store Aqk; store L = strict_tril(-beta_row * acc_l) ----
    m_lo = r64[:, None] >= r64[None, :]
    tl.store(a_ptr + (slab + r64[:, None]) * BT + r64[None, :],
             tl.where(m_lo, acc_a, 0.0).to(a_ptr.dtype.element_ty))
    L = tl.where(r64[:, None] > r64[None, :], -b_beta[:, None] * acc_l, 0.0)
    tl.store(l_ptr + (slab + r64[:, None]) * BT + r64[None, :],
             L.to(l_ptr.dtype.element_ty))


# ---------------------------------------------------------------------------
# Kernel 1b: per-chunk triangular inverse + w/u (compute-bound)
# ---------------------------------------------------------------------------
@triton.jit(do_not_specialize=["T", "NT"])
def _kda_wu_kernel(
    v_ptr, beta_ptr, l_ptr, f_ptr, w_ptr, u_ptr,
    T, NT,
    H: tl.constexpr,
    K: tl.constexpr,
    V: tl.constexpr,
    BT: tl.constexpr,
):
    i_n = tl.program_id(0)
    i_bh = tl.program_id(1)
    i_b = i_bh // H
    i_h = i_bh % H

    BK: tl.constexpr = K // 2
    r64 = tl.arange(0, BT)
    tok = i_b * T + i_n * BT + r64
    slab = (i_bh * NT + i_n) * BT
    cols0 = tl.arange(0, BK)
    cols1 = BK + tl.arange(0, BK)
    b_beta = tl.load(beta_ptr + tok * H + i_h).to(tl.float32)

    L = tl.load(l_ptr + (slab + r64[:, None]) * BT + r64[None, :]).to(tl.float32)

    # X = prod_{i=0..5} (I + L^{2^i}) = sum_{j=0}^{63} L^j.  Powers are built
    # on the fly so only a few 64x64 tiles are live at any time.
    I64 = (r64[:, None] == r64[None, :]).to(tl.float32)
    X = I64 + L
    P = L
    for _ in tl.static_range(5):
        P = tl.dot(P, P, input_precision="tf32")
        X = tl.dot(X, I64 + P, input_precision="tf32")
    Xb = X.to(tl.bfloat16)

    # w = X @ f
    b_f0 = tl.load(f_ptr + (slab + r64[:, None]) * K + cols0[None, :])
    w0 = tl.dot(Xb, b_f0)
    tl.store(w_ptr + (slab + r64[:, None]) * K + cols0[None, :],
             w0.to(w_ptr.dtype.element_ty))
    b_f1 = tl.load(f_ptr + (slab + r64[:, None]) * K + cols1[None, :])
    w1 = tl.dot(Xb, b_f1)
    tl.store(w_ptr + (slab + r64[:, None]) * K + cols1[None, :],
             w1.to(w_ptr.dtype.element_ty))

    # u = X @ (beta * v)
    b_v0 = tl.load(v_ptr + ((tok * H + i_h) * V)[:, None] + cols0[None, :])
    u0 = tl.dot(Xb, (b_beta[:, None] * b_v0.to(tl.float32)).to(tl.bfloat16))
    tl.store(u_ptr + (slab + r64[:, None]) * V + cols0[None, :],
             u0.to(u_ptr.dtype.element_ty))
    b_v1 = tl.load(v_ptr + ((tok * H + i_h) * V)[:, None] + cols1[None, :])
    u1 = tl.dot(Xb, (b_beta[:, None] * b_v1.to(tl.float32)).to(tl.bfloat16))
    tl.store(u_ptr + (slab + r64[:, None]) * V + cols1[None, :],
             u1.to(u_ptr.dtype.element_ty))


# ---------------------------------------------------------------------------
# Kernel 2: sequential inter-chunk scan (one program per head x value-block)
# ---------------------------------------------------------------------------
@triton.jit(do_not_specialize=["T", "NT"])
def _kda_scan_kernel(
    w_ptr, u_ptr, a_ptr, qg_ptr, kg_ptr, gl_ptr, o_ptr,
    T, NT,
    H: tl.constexpr,
    K: tl.constexpr,
    V: tl.constexpr,
    BT: tl.constexpr,
    BV: tl.constexpr,
):
    i_bh = tl.program_id(0)
    i_v = tl.program_id(1)
    i_b = i_bh // H
    i_h = i_bh % H

    r64 = tl.arange(0, BT)
    rk = tl.arange(0, K)
    colsv = i_v * BV + tl.arange(0, BV)

    S = tl.zeros([K, BV], dtype=tl.float32)

    for n in range(NT):
        slab = (i_bh * NT + n) * BT
        b_w = tl.load(w_ptr + (slab + r64[:, None]) * K + rk[None, :])
        b_qg = tl.load(qg_ptr + (slab + r64[:, None]) * K + rk[None, :])
        b_kg = tl.load(kg_ptr + (slab + r64[:, None]) * K + rk[None, :])
        b_a = tl.load(a_ptr + (slab + r64[:, None]) * BT + r64[None, :])
        b_u = tl.load(u_ptr + (slab + r64[:, None]) * V + colsv[None, :]).to(tl.float32)
        b_gl = tl.load(gl_ptr + (i_bh * NT + n) * K + rk)

        S_bf = S.to(tl.bfloat16)
        b_v = b_u - tl.dot(b_w, S_bf)
        b_v_bf = b_v.to(tl.bfloat16)
        b_o = tl.dot(b_qg, S_bf) + tl.dot(b_a, b_v_bf)

        tok = i_b * T + n * BT + r64
        tl.store(o_ptr + ((tok * H + i_h) * V)[:, None] + colsv[None, :],
                 b_o.to(o_ptr.dtype.element_ty))

        S = S * tl.exp2(b_gl)[:, None] + tl.dot(tl.trans(b_kg), b_v_bf)


# ---------------------------------------------------------------------------
# Host wrapper
# ---------------------------------------------------------------------------
def _prep_config():
    return {"num_warps": 4, "num_stages": 1}


def _wu_config():
    return {"num_warps": 4, "num_stages": 1}


def _scan_config():
    return {"BV": 32, "num_warps": 4, "num_stages": 3}


class Model(nn.Module):
    """KDA forward (chunk form). No learned parameters; all inputs are activations."""

    def __init__(self, B: int, T: int, H: int, K: int, V: int, chunk_size: int = 64):
        super().__init__()
        self.B, self.T, self.H, self.K, self.V = B, T, H, K, V
        self.chunk_size = chunk_size
        self.scale = float(K) ** -0.5
        # No learned params; declare a dummy buffer so state_dict is well-defined.
        self.register_buffer("_dummy", torch.zeros(1), persistent=False)
        self._ws = None
        self._ws_key = None

    def _workspace(self, B, T, H, K, V, NT, device):
        key = (B, T, H, K, V)
        if self._ws_key != key:
            n_chunk_rows = B * H * NT * self.chunk_size
            cs = self.chunk_size
            self._ws = {
                "w": torch.empty(n_chunk_rows, K, device=device, dtype=torch.bfloat16),
                "u": torch.empty(n_chunk_rows, V, device=device, dtype=torch.bfloat16),
                "qg": torch.empty(n_chunk_rows, K, device=device, dtype=torch.bfloat16),
                "kg": torch.empty(n_chunk_rows, K, device=device, dtype=torch.bfloat16),
                "a": torch.empty(n_chunk_rows, cs, device=device, dtype=torch.bfloat16),
                "l": torch.empty(n_chunk_rows, cs, device=device, dtype=torch.bfloat16),
                "f": torch.empty(n_chunk_rows, K, device=device, dtype=torch.bfloat16),
                "gl": torch.empty(B * H * NT, K, device=device, dtype=torch.float32),
            }
            self._ws_key = key
        return self._ws

    def forward(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        g: torch.Tensor,
        beta: torch.Tensor,
    ) -> torch.Tensor:
        B, T, H, K = q.shape
        V = v.shape[-1]
        BT = self.chunk_size
        assert T % BT == 0
        NT = T // BT
        ws = self._workspace(B, T, H, K, V, NT, q.device)
        o = torch.empty_like(v)

        cfg = _prep_config()
        _kda_prep_kernel[(NT, B * H)](
            q, k, g, beta,
            ws["a"], ws["qg"], ws["kg"], ws["gl"], ws["l"], ws["f"],
            self.scale,
            T, NT,
            H=H, K=K, BT=BT,
            num_warps=cfg["num_warps"], num_stages=cfg["num_stages"],
        )

        wcfg = _wu_config()
        _kda_wu_kernel[(NT, B * H)](
            v, beta, ws["l"], ws["f"], ws["w"], ws["u"],
            T, NT,
            H=H, K=K, V=V, BT=BT,
            num_warps=wcfg["num_warps"], num_stages=wcfg["num_stages"],
        )

        scfg = _scan_config()
        _kda_scan_kernel[(B * H, V // scfg["BV"])](
            ws["w"], ws["u"], ws["a"], ws["qg"], ws["kg"], ws["gl"], o,
            T, NT,
            H=H, K=K, V=V, BT=BT,
            BV=scfg["BV"],
            num_warps=scfg["num_warps"], num_stages=scfg["num_stages"],
        )
        return o


# Module-level shape shims (overridden by check.py / benchmark.py per shape).
B = 2
T = 1024
H = 8
K = 128
V = 128
CHUNK_SIZE = 64


def get_inputs():
    """Return a list of activations for one forward call.

    bf16 for q/k/v/beta; fp32 for the log-decay g (per FLA convention).
    """
    torch.manual_seed(0)
    q = torch.randn(B, T, H, K, dtype=torch.bfloat16) * 0.1
    k = torch.randn(B, T, H, K, dtype=torch.bfloat16) * 0.1
    v = torch.randn(B, T, H, V, dtype=torch.bfloat16) * 0.1
    # log-decay: small negative numbers so exp(g) is in (0, 1).
    g = (torch.randn(B, T, H, K, dtype=torch.float32) * 0.1 - 0.05)
    beta = torch.sigmoid(torch.randn(B, T, H, dtype=torch.bfloat16))
    return [q, k, v, g, beta]


def get_init_inputs():
    return [B, T, H, K, V, CHUNK_SIZE]

20260805_092921_or-fable_qwen_qwen3.8-max_02_kda_cutlass