KernelBench hard · RTX PRO 6000
KDA CUTLASS DeepSeek V4.1 Flash
3.74%geomean peak fraction across shapes
manually audited: clean
Two-kernel Triton chunk-parallel KDA forward. The intra kernel builds the decay-factored Gram matrix and inverts (I+L) exactly through a blocked product form on 16x16 diagonal blocks; the chain kernel walks the inter-chunk recurrence term-for-term against the reference. No CUTLASS needed: the prompt permits Triton and FLA is not even installed on the box.
harnessdeepseek-claudeagent session2h 15mtotal wall2h 16mcheck9sbenchmark2soutput tokens446,987cost$23.77gpu-lock wait1h 35mgpu-lock held7mregimecompute
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
2×1024×8×128×128×640.105 ms4.1%0.24 TB/s · 13% of 1.8 TB/s HBM · also 20 TFLOPS (4% of compute)
2×2048×8×128×128×640.159 ms5.4%0.32 TB/s · 18% of 1.8 TB/s HBM · also 27 TFLOPS (5% of compute)
1×4096×8×128×128×640.198 ms4.3%0.25 TB/s · 14% of 1.8 TB/s HBM · also 22 TFLOPS (4% of compute)
1×2048×4×128×128×640.104 ms2.1%0.12 TB/s · 7% of 1.8 TB/s HBM · also 10 TFLOPS (2% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(4.1% · 5.4% · 4.3% · 2.1%) = 3.7%
Kernel source (redacted)
"""Kimi Delta Attention (KDA) chunk-parallel forward -- custom Triton kernels.
Two-kernel design (both written from scratch; no library attention ops):
K1 intra-chunk pass, grid (NT, B*H), one program per (batch*head, chunk):
gc = cumsum(g) (chunk-local, per channel)
P = k * exp(-gc) Kp = k * exp(gc) Qs = q*scale*exp(gc)
G = Kp @ P^T (natural decay direction)
L = diag(beta) tril(G, -1)
Minv = (I + L)^-1 (blocked triangular inverse)
u = Minv @ (beta*v) w = Minv @ (beta*Kp)
Pt = P * exp(gc_last) d = exp(gc_last)
Aqk = tril(Qs @ P^T) (j <= i)
emits u, w, Qs, Pt (bf16), Aqk (bf16), d (fp32).
K2 inter-chunk state pass, grid (NV, B*H), one program per (batch*head, V block):
S = diag(d) S + Pt^T (u - w S)
o = Qs S + Aqk (u - w S)
The triangular inverse is factored as (I+L)^-1 = (I+N)^-1 Md with Md the
inverse of the block-diagonal (BC x BC) part and N = Md @ L_off strictly
block-lower (nilpotent). Md itself is evaluated by the exact product form
(I+A)(I+A^2)(I+A^4)(I+A^8) with A = -L_diag and L_diag^BC = 0, which never
forms large intermediate powers the way a plain Neumann series would.
Precision: the (C,K) operand tiles are rounded to bf16 for the tensor-core MMAs
(every accumulator stays fp32); the only values that need more than bf16 are the
(C,C) inverse chain, which keeps tf32. That measured ~0.5% relative deviation
from the all-tf32 version -- far inside the 0.05 tolerance -- and removed the
ptxas spills that were costing ~1.6x on this kernel.
Performance notes (RTX PRO 6000 / SM120):
* the chain kernel's cost is its serial dependency (state -> dot -> state),
so it wants the *smallest* useful V block: 16 columns with every (b,h)
chain in flight measured 3x faster than 64-column blocks.
* this harness flushes L2 (128 MB zero_) before each timed call, so the first
touch of anything is a DRAM miss and every dirty flush line we evict costs a
writeback: streamed inputs are loaded evict_first, the intermediate tiles
that the chain kernel re-reads nv times are stored evict_last.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import triton
import triton.language as tl
_TF32 = tl.constexpr("tf32")
# The harness flushes L2 (128 MB zero_) before every timed call, so the first
# touch of anything is a DRAM miss and every dirty flush line we evict costs a
# writeback. q/k/v/g are streamed once per program -> evict-first; the tiles
# the chain kernel reads back nv times -> evict-last so they stay resident.
_EV_FIRST = tl.constexpr("evict_first")
_EV_LAST = tl.constexpr("evict_last")
@triton.jit
def _kda_intra_kernel(
q_ptr, k_ptr, v_ptr, g_ptr, beta_ptr,
u_ptr, w_ptr, qs_ptr, pt_ptr, aqk_ptr, d_ptr,
scale, T, NT,
H: tl.constexpr, K: tl.constexpr, V: tl.constexpr,
C: tl.constexpr, BC: tl.constexpr,
):
pid_c = tl.program_id(0)
pid_bh = tl.program_id(1)
b = pid_bh // H
h = pid_bh - b * H
rc = tl.arange(0, C)
rk = tl.arange(0, K)
rv = tl.arange(0, V)
rows = pid_c * C + rc
qk_off = (b * T + rows)[:, None] * (H * K) + h * K + rk[None, :]
q = tl.load(q_ptr + qk_off, eviction_policy=_EV_FIRST)
k = tl.load(k_ptr + qk_off, eviction_policy=_EV_FIRST)
g = tl.load(g_ptr + qk_off, eviction_policy=_EV_FIRST)
v = tl.load(v_ptr + (b * T + rows)[:, None] * (H * V) + h * V + rv[None, :],
eviction_policy=_EV_FIRST)
beta = tl.load(beta_ptr + (b * T + rows) * H + h)
gc = tl.cumsum(g, axis=0)
glast = tl.sum(tl.where(rc[:, None] == C - 1, gc, 0.0), axis=0) # (K,)
# Keep the (C,K) operand tiles in bf16: they only ever feed tensor-core MMAs
# whose accumulators are fp32, and the fp32 copies were the register hogs
# (ptxas was spilling 24-70 words/thread, which cost more than the rounding).
kf = k.to(tl.float32)
em = tl.exp(-gc)
eg = tl.exp(gc)
P = (kf * em).to(tl.bfloat16) # k e^{-gc}
Kp = (kf * eg).to(tl.bfloat16) # k e^{gc}
Qs = (q.to(tl.float32) * (scale * eg)).to(tl.bfloat16) # q*scale e^{gc}
P32 = P.to(tl.float32)
# ---- intra-chunk Gram in the decaying direction: G[a,b] = k_a k_b e^{gc_a-gc_b}
G = tl.dot(Kp, tl.trans(P), out_dtype=tl.float32)
strict_lower = rc[:, None] > rc[None, :]
L = tl.where(strict_lower, G, 0.0) * beta[:, None] # diag(beta) tril(G,-1)
A = -L
eye = (rc[:, None] == rc[None, :]).to(tl.float32)
if C > BC:
blk = (rc[:, None] // BC) == (rc[None, :] // BC)
Ad = tl.where(blk, A, 0.0)
else:
Ad = A
# Md = (I - Ad)^{-1} = (I+Ad)(I+Ad^2)(I+Ad^4)(I+Ad^8); Ad^BC = 0 for BC=16
M = eye + Ad
A2 = tl.dot(Ad, Ad, input_precision=_TF32)
M = M + tl.dot(M, A2, input_precision=_TF32)
A4 = tl.dot(A2, A2, input_precision=_TF32)
M = M + tl.dot(M, A4, input_precision=_TF32)
A8 = tl.dot(A4, A4, input_precision=_TF32)
Md = M + tl.dot(M, A8, input_precision=_TF32)
if C > BC:
Loff = L - tl.where(blk, L, 0.0)
N = tl.dot(Md, Loff, input_precision=_TF32) # strictly block lower
if C // BC > 2:
N2 = tl.dot(N, N, input_precision=_TF32)
N3 = tl.dot(N2, N, input_precision=_TF32)
Q = eye - N + N2 - N3
else:
Q = eye - N
Minv = tl.dot(Q, Md, input_precision=_TF32)
else:
Minv = Md
Mib = Minv.to(tl.bfloat16)
betab = beta.to(tl.bfloat16)
u = tl.dot(Mib, betab[:, None] * v, out_dtype=tl.float32) # (C,V)
w = tl.dot(Mib, betab[:, None] * Kp, out_dtype=tl.float32) # (C,K)
dl = tl.exp(glast)
Pt = P32 * dl[None, :] # (C,K)
Aqk = tl.dot(Qs, tl.trans(P), out_dtype=tl.float32)
Aqk = tl.where(rc[:, None] >= rc[None, :], Aqk, 0.0) # keep j <= i (diag incl.)
idx = pid_bh * NT + pid_c
cu = u_ptr + idx * (C * V) + rc[:, None] * V + rv[None, :]
ck = rc[:, None] * K + rk[None, :]
tl.store(cu, u.to(tl.bfloat16), eviction_policy=_EV_LAST)
tl.store(w_ptr + idx * (C * K) + ck, w.to(tl.bfloat16), eviction_policy=_EV_LAST)
tl.store(qs_ptr + idx * (C * K) + ck, Qs.to(tl.bfloat16), eviction_policy=_EV_LAST)
tl.store(pt_ptr + idx * (K * C) + rk[:, None] * C + rc[None, :],
tl.trans(Pt).to(tl.bfloat16), eviction_policy=_EV_LAST)
tl.store(aqk_ptr + idx * (C * C) + rc[:, None] * C + rc[None, :],
Aqk.to(tl.bfloat16), eviction_policy=_EV_LAST)
tl.store(d_ptr + idx * K + rk, dl, eviction_policy=_EV_LAST)
@triton.jit
def _kda_chain_kernel(
u_ptr, w_ptr, qs_ptr, pt_ptr, aqk_ptr, d_ptr, o_ptr,
T, NT,
H: tl.constexpr, K: tl.constexpr, V: tl.constexpr,
C: tl.constexpr, BV: tl.constexpr,
):
pid_v = tl.program_id(0)
pid_bh = tl.program_id(1)
b = pid_bh // H
h = pid_bh - b * H
rk = tl.arange(0, K)
rc = tl.arange(0, C)
rv = pid_v * BV + tl.arange(0, BV)
S = tl.zeros((K, BV), dtype=tl.float32)
for c in range(0, NT):
idx = pid_bh * NT + c
u = tl.load(u_ptr + idx * (C * V) + rc[:, None] * V + rv[None, :],
eviction_policy=_EV_LAST).to(tl.float32)
w = tl.load(w_ptr + idx * (C * K) + rc[:, None] * K + rk[None, :],
eviction_policy=_EV_LAST)
qs = tl.load(qs_ptr + idx * (C * K) + rc[:, None] * K + rk[None, :],
eviction_policy=_EV_LAST)
pt = tl.load(pt_ptr + idx * (K * C) + rk[:, None] * C + rc[None, :],
eviction_policy=_EV_LAST)
aqk = tl.load(aqk_ptr + idx * (C * C) + rc[:, None] * C + rc[None, :],
eviction_policy=_EV_LAST)
vi = u - tl.dot(w, S.to(tl.bfloat16)) # (C,BV)
vib = vi.to(tl.bfloat16)
o = tl.dot(qs, S.to(tl.bfloat16)) + tl.dot(aqk, vib) # (C,BV)
d = tl.load(d_ptr + idx * K + rk, eviction_policy=_EV_LAST)
S = S * d[:, None] + tl.dot(pt, vib)
tl.store(o_ptr + (b * T + c * C + rc)[:, None] * (H * V) + h * V + rv[None, :],
o.to(tl.bfloat16), eviction_policy=_EV_FIRST)
# ---------------------------------------------------------------------------
# host side
# ---------------------------------------------------------------------------
_DEF_C = 32
_DEF_BC = 16
def _pick_cfg(B, T, H, K, V):
"""(chunk, BV) for the intra kernel and BV for the chain kernel.
The chain kernel is latency-bound on its serial dependency (state -> dot ->
state), so its speed is set by how many independent (batch*head, V-block)
chains can run concurrently. Measured on this part: 16-column V blocks with
every (b,h) chain in flight is 3x faster than 64-column blocks, even though
it re-reads the V-independent tiles more often (those stay L2-resident).
"""
C = 32 if T % 32 == 0 else 64
nv = 1
while nv < 8 and V // (nv * 2) >= 16:
nv *= 2
return C, V // nv
class Model(nn.Module):
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
self.register_buffer("_dummy", torch.zeros(1), persistent=False)
self._bufs = None
self._key = None
def _get_bufs(self, B, T, H, K, V, C, BV, device):
NT = T // C
bhnt = B * H * NT
key = (B, T, H, K, V, C, BV, str(device))
if self._key == key:
return self._bufs
f32 = torch.float32
bf = torch.bfloat16
bufs = {
"u": torch.empty((bhnt, C, V), dtype=bf, device=device),
"w": torch.empty((bhnt, C, K), dtype=bf, device=device),
"qs": torch.empty((bhnt, C, K), dtype=bf, device=device),
"pt": torch.empty((bhnt, K, C), dtype=bf, device=device),
"aqk": torch.empty((bhnt, C, C), dtype=bf, device=device),
"d": torch.empty((bhnt, K), dtype=f32, device=device),
}
self._bufs = bufs
self._key = key
return bufs
def forward(self, q, k, v, g, beta):
B, T, H, K = q.shape
V = v.shape[-1]
C, BV = _pick_cfg(B, T, H, K, V)
if T % C or V % BV:
C, BV = self.chunk_size, V
NT = T // C
dev = q.device
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
g = g.contiguous()
beta = beta.contiguous()
if q.dtype != torch.bfloat16:
q = q.to(torch.bfloat16)
k = k.to(torch.bfloat16)
v = v.to(torch.bfloat16)
beta = beta.to(torch.bfloat16)
if g.dtype != torch.float32:
g = g.to(torch.float32)
bufs = self._get_bufs(B, T, H, K, V, C, BV, dev)
BC = _DEF_BC if C >= _DEF_BC else C
o = torch.empty((B, T, H, V), dtype=torch.bfloat16, device=dev)
_kda_intra_kernel[(NT, B * H)](
q, k, v, g, beta,
bufs["u"], bufs["w"], bufs["qs"], bufs["pt"], bufs["aqk"], bufs["d"],
self.scale, T, NT,
H=H, K=K, V=V, C=C, BC=BC,
num_warps=4, num_stages=1,
)
nv = V // BV
_kda_chain_kernel[(nv, B * H)](
bufs["u"], bufs["w"], bufs["qs"], bufs["pt"], bufs["aqk"], bufs["d"], o,
T, NT,
H=H, K=K, V=V, C=C, BV=BV,
num_warps=4, num_stages=3,
)
return o
def get_init_inputs():
return [2, 1024, 8, 128, 128, 64]
def get_inputs():
torch.manual_seed(0)
B, T, H, K, V = 2, 1024, 8, 128, 128
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
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]
20260910_202119_deepseek-claude_deepseek-flash_02_kda_cutlass