KernelBench hard · H100
KDA CUTLASS Qwen 3.8 Max
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.
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
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