KernelBench hard · RTX PRO 6000
KDA CUTLASS Grok 4.6
4.06%geomean peak fraction across shapes
manually audited: clean
Custom Triton KDA chunk-form (intra-chunk Akk, 16x16 unit-lower inverse, WY factors, tiled inter-chunk recurrence). CUDA-graph keyed on Q/K/V/g/beta data_ptr; recaptures on a new key. Same pattern as published gpt-5.5 KDA.
harnessgrokagent session1h 21mtotal wall1h 23mcheck6sbenchmark2soutput tokens—regimecompute
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
2×1024×8×128×128×640.085 ms5.1%0.30 TB/s · 17% of 1.8 TB/s HBM · also 25 TFLOPS (5% of compute)
2×2048×8×128×128×640.146 ms5.9%0.35 TB/s · 19% of 1.8 TB/s HBM · also 29 TFLOPS (6% of compute)
1×4096×8×128×128×640.196 ms4.4%0.26 TB/s · 14% of 1.8 TB/s HBM · also 22 TFLOPS (4% of compute)
1×2048×4×128×128×640.103 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(5.1% · 5.9% · 4.4% · 2.1%) = 4.1%
Kernel source (redacted)
"""Kimi Delta Attention chunk-form forward for SM120.
Custom Triton kernels:
1. intra-chunk Akk (per-channel decay)
2. 16x16-block unit-lower inverse times beta
3. WY factors w, u plus Aqk / qg / kbar
4. tiled inter-chunk state recurrence and output
CUDA graphs replay the four launches on the hot path.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import triton
import triton.language as tl
OP_TYPE = "linear_attention"
SUPPORTED_PRECISIONS = ["bf16"]
HARDWARE_REQUIRED = ["RTX_PRO_6000", "H100", "B200"]
@triton.jit
def _build_A_kernel(
k_ptr, g_ptr, beta_ptr, A_ptr,
T, H, NT,
BT: tl.constexpr,
BK: tl.constexpr,
):
i_t = tl.program_id(0)
i_bh = tl.program_id(1)
i_b = i_bh // H
i_h = i_bh % H
t0 = i_t * BT
K = 128
o_c = tl.arange(0, BT)
o_k = tl.arange(0, BK)
b_A = tl.zeros([BT, BT], dtype=tl.float32)
b_k = tl.load(k_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + o_k[None, :]).to(tl.float32)
b_g = tl.cumsum(tl.load(g_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + o_k[None, :]), 0)
b_A += tl.dot(b_k * tl.exp(b_g), tl.trans(b_k * tl.exp(-b_g)))
b_k = tl.load(k_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + (o_k + BK)[None, :]).to(tl.float32)
b_g = tl.cumsum(tl.load(g_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + (o_k + BK)[None, :]), 0)
b_A += tl.dot(b_k * tl.exp(b_g), tl.trans(b_k * tl.exp(-b_g)))
b_beta = tl.load(beta_ptr + (i_b * T + t0 + o_c) * H + i_h).to(tl.float32)
b_A = tl.where(o_c[:, None] > o_c[None, :], b_A * b_beta[:, None], 0.0)
nid = i_bh * NT + i_t
tl.store(A_ptr + nid * BT * BT + o_c[:, None] * BT + o_c[None, :], b_A)
@triton.jit
def _invert_A_kernel(
A_ptr, beta_ptr, T, H, NT,
BT: tl.constexpr,
):
"""A <- (I+A)^{-1} * diag(beta). A strictly lower, packed [BH*NT, BT, BT]."""
i_t = tl.program_id(0)
i_bh = tl.program_id(1)
i_b = i_bh // H
i_h = i_bh % H
nid = i_bh * NT + i_t
t0 = i_t * BT
A = A_ptr + nid * BT * BT
o = tl.arange(0, 16)
m_A = o[:, None] > o[None, :]
m_I = o[:, None] == o[None, :]
b11 = -tl.where(m_A, tl.load(A + o[:, None] * BT + o[None, :]).to(tl.float32), 0.0)
b22 = -tl.where(m_A, tl.load(A + (16 + o)[:, None] * BT + (16 + o)[None, :]).to(tl.float32), 0.0)
b33 = -tl.where(m_A, tl.load(A + (32 + o)[:, None] * BT + (32 + o)[None, :]).to(tl.float32), 0.0)
b44 = -tl.where(m_A, tl.load(A + (48 + o)[:, None] * BT + (48 + o)[None, :]).to(tl.float32), 0.0)
for i in range(2, 16):
a = -tl.load(A + i * BT + o)
a = tl.where(o < i, a, 0.0)
a += tl.sum(a[:, None] * b11, 0)
b11 = tl.where((o == i)[:, None], a[None, :], b11)
for i in range(2, 16):
a = -tl.load(A + (16 + i) * BT + (16 + o))
a = tl.where(o < i, a, 0.0)
a += tl.sum(a[:, None] * b22, 0)
b22 = tl.where((o == i)[:, None], a[None, :], b22)
for i in range(2, 16):
a = -tl.load(A + (32 + i) * BT + (32 + o))
a = tl.where(o < i, a, 0.0)
a += tl.sum(a[:, None] * b33, 0)
b33 = tl.where((o == i)[:, None], a[None, :], b33)
for i in range(2, 16):
a = -tl.load(A + (48 + i) * BT + (48 + o))
a = tl.where(o < i, a, 0.0)
a += tl.sum(a[:, None] * b44, 0)
b44 = tl.where((o == i)[:, None], a[None, :], b44)
b11 += m_I
b22 += m_I
b33 += m_I
b44 += m_I
a21 = tl.load(A + (16 + o)[:, None] * BT + o[None, :]).to(tl.float32)
a31 = tl.load(A + (32 + o)[:, None] * BT + o[None, :]).to(tl.float32)
a32 = tl.load(A + (32 + o)[:, None] * BT + (16 + o)[None, :]).to(tl.float32)
a41 = tl.load(A + (48 + o)[:, None] * BT + o[None, :]).to(tl.float32)
a42 = tl.load(A + (48 + o)[:, None] * BT + (16 + o)[None, :]).to(tl.float32)
a43 = tl.load(A + (48 + o)[:, None] * BT + (32 + o)[None, :]).to(tl.float32)
ai21 = -tl.dot(tl.dot(b22, a21), b11)
ai32 = -tl.dot(tl.dot(b33, a32), b22)
ai43 = -tl.dot(tl.dot(b44, a43), b33)
ai31 = -tl.dot(b33, tl.dot(a31, b11) + tl.dot(a32, ai21))
ai42 = -tl.dot(b44, tl.dot(a42, b22) + tl.dot(a43, ai32))
ai41 = -tl.dot(b44, tl.dot(a41, b11) + tl.dot(a42, ai21) + tl.dot(a43, ai31))
b0 = tl.load(beta_ptr + (i_b * T + t0 + o) * H + i_h).to(tl.float32)
b1 = tl.load(beta_ptr + (i_b * T + t0 + 16 + o) * H + i_h).to(tl.float32)
b2 = tl.load(beta_ptr + (i_b * T + t0 + 32 + o) * H + i_h).to(tl.float32)
b3 = tl.load(beta_ptr + (i_b * T + t0 + 48 + o) * H + i_h).to(tl.float32)
tl.store(A + o[:, None] * BT + o[None, :], b11 * b0[None, :])
tl.store(A + (16 + o)[:, None] * BT + (16 + o)[None, :], b22 * b1[None, :])
tl.store(A + (32 + o)[:, None] * BT + (32 + o)[None, :], b33 * b2[None, :])
tl.store(A + (48 + o)[:, None] * BT + (48 + o)[None, :], b44 * b3[None, :])
tl.store(A + (16 + o)[:, None] * BT + o[None, :], ai21 * b0[None, :])
tl.store(A + (32 + o)[:, None] * BT + (16 + o)[None, :], ai32 * b1[None, :])
tl.store(A + (48 + o)[:, None] * BT + (32 + o)[None, :], ai43 * b2[None, :])
tl.store(A + (32 + o)[:, None] * BT + o[None, :], ai31 * b0[None, :])
tl.store(A + (48 + o)[:, None] * BT + (16 + o)[None, :], ai42 * b1[None, :])
tl.store(A + (48 + o)[:, None] * BT + o[None, :], ai41 * b0[None, :])
@triton.jit
def _build_wy_kernel(
q_ptr, k_ptr, v_ptr, g_ptr, A_ptr,
w_ptr, u_ptr, aqk_ptr, qg_ptr, kbar_ptr, glast_ptr,
scale, T, H, NT,
BT: tl.constexpr,
BK: tl.constexpr,
):
i_t = tl.program_id(0)
i_bh = tl.program_id(1)
i_b = i_bh // H
i_h = i_bh % H
t0 = i_t * BT
K = 128
V = 128
o_c = tl.arange(0, BT)
o_k = tl.arange(0, BK)
nid = i_bh * NT + i_t
pack = nid * BT
b_A = tl.load(A_ptr + nid * BT * BT + o_c[:, None] * BT + o_c[None, :])
b_Aqk = tl.zeros([BT, BT], dtype=tl.float32)
b_k = tl.load(k_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + o_k[None, :]).to(tl.float32)
b_g = tl.cumsum(tl.load(g_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + o_k[None, :]), 0)
b_eg = tl.exp(b_g)
b_gl = tl.sum(tl.where(o_c[:, None] == (BT - 1), b_g, 0.0), axis=0)
tl.store(glast_ptr + nid * K + o_k, b_gl)
b_q = tl.load(q_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + o_k[None, :]).to(tl.float32)
b_qg = b_q * scale * b_eg
tl.store(w_ptr + (pack + o_c)[:, None] * K + o_k[None, :], tl.dot(b_A, b_k * b_eg).to(tl.bfloat16))
tl.store(qg_ptr + (pack + o_c)[:, None] * K + o_k[None, :], b_qg.to(tl.bfloat16))
tl.store(kbar_ptr + (pack + o_c)[:, None] * K + o_k[None, :], (b_k * tl.exp(b_gl[None, :] - b_g)).to(tl.bfloat16))
b_Aqk += tl.dot(b_qg, tl.trans(b_k * tl.exp(-b_g)))
b_k = tl.load(k_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + (o_k + BK)[None, :]).to(tl.float32)
b_g = tl.cumsum(tl.load(g_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + (o_k + BK)[None, :]), 0)
b_eg = tl.exp(b_g)
b_gl = tl.sum(tl.where(o_c[:, None] == (BT - 1), b_g, 0.0), axis=0)
tl.store(glast_ptr + nid * K + (o_k + BK), b_gl)
b_q = tl.load(q_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * K + (o_k + BK)[None, :]).to(tl.float32)
b_qg = b_q * scale * b_eg
tl.store(w_ptr + (pack + o_c)[:, None] * K + (o_k + BK)[None, :], tl.dot(b_A, b_k * b_eg).to(tl.bfloat16))
tl.store(qg_ptr + (pack + o_c)[:, None] * K + (o_k + BK)[None, :], b_qg.to(tl.bfloat16))
tl.store(kbar_ptr + (pack + o_c)[:, None] * K + (o_k + BK)[None, :], (b_k * tl.exp(b_gl[None, :] - b_g)).to(tl.bfloat16))
b_Aqk += tl.dot(b_qg, tl.trans(b_k * tl.exp(-b_g)))
b_Aqk = tl.where(o_c[:, None] >= o_c[None, :], b_Aqk, 0.0)
tl.store(aqk_ptr + (pack + o_c)[:, None] * BT + o_c[None, :], b_Aqk.to(tl.bfloat16))
b_v = tl.load(v_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * V + o_k[None, :]).to(tl.float32)
tl.store(u_ptr + (pack + o_c)[:, None] * V + o_k[None, :], tl.dot(b_A, b_v).to(tl.bfloat16))
b_v = tl.load(v_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * V + (o_k + BK)[None, :]).to(tl.float32)
tl.store(u_ptr + (pack + o_c)[:, None] * V + (o_k + BK)[None, :], tl.dot(b_A, b_v).to(tl.bfloat16))
@triton.jit(do_not_specialize=["NT"])
def _inter_kernel(
w_ptr, u_ptr, aqk_ptr, qg_ptr, kbar_ptr, glast_ptr, o_ptr,
T, H, NT,
BT: tl.constexpr,
BK: tl.constexpr,
BV: tl.constexpr,
):
i_bh = tl.program_id(0)
iv = tl.program_id(1)
i_b = i_bh // H
i_h = i_bh % H
o_c = tl.arange(0, BT)
o_k = tl.arange(0, BK)
o_v = iv * BV + tl.arange(0, BV)
K = 128
V = 128
s0 = tl.zeros([BK, BV], dtype=tl.float32)
s1 = tl.zeros([BK, BV], dtype=tl.float32)
s2 = tl.zeros([BK, BV], dtype=tl.float32)
s3 = tl.zeros([BK, BV], dtype=tl.float32)
for it in range(0, NT):
idx = (i_bh * NT + it) * BT
t0 = it * BT
b_vi = tl.load(u_ptr + (idx + o_c)[:, None] * V + o_v[None, :]).to(tl.float32)
b_w = tl.load(w_ptr + (idx + o_c)[:, None] * K + o_k[None, :])
b_vi -= tl.dot(b_w, s0.to(tl.bfloat16))
b_w = tl.load(w_ptr + (idx + o_c)[:, None] * K + (o_k + BK)[None, :])
b_vi -= tl.dot(b_w, s1.to(tl.bfloat16))
b_w = tl.load(w_ptr + (idx + o_c)[:, None] * K + (o_k + 2 * BK)[None, :])
b_vi -= tl.dot(b_w, s2.to(tl.bfloat16))
b_w = tl.load(w_ptr + (idx + o_c)[:, None] * K + (o_k + 3 * BK)[None, :])
b_vi -= tl.dot(b_w, s3.to(tl.bfloat16))
b_o = tl.zeros([BT, BV], dtype=tl.float32)
b_q = tl.load(qg_ptr + (idx + o_c)[:, None] * K + o_k[None, :])
b_o += tl.dot(b_q, s0.to(tl.bfloat16))
b_q = tl.load(qg_ptr + (idx + o_c)[:, None] * K + (o_k + BK)[None, :])
b_o += tl.dot(b_q, s1.to(tl.bfloat16))
b_q = tl.load(qg_ptr + (idx + o_c)[:, None] * K + (o_k + 2 * BK)[None, :])
b_o += tl.dot(b_q, s2.to(tl.bfloat16))
b_q = tl.load(qg_ptr + (idx + o_c)[:, None] * K + (o_k + 3 * BK)[None, :])
b_o += tl.dot(b_q, s3.to(tl.bfloat16))
b_A = tl.load(aqk_ptr + (idx + o_c)[:, None] * BT + o_c[None, :])
b_o += tl.dot(b_A, b_vi.to(tl.bfloat16))
tl.store(o_ptr + ((i_b * T + t0 + o_c) * H + i_h)[:, None] * V + o_v[None, :], b_o.to(tl.bfloat16))
b_g = tl.load(glast_ptr + (i_bh * NT + it) * K + o_k)
s0 = s0 * tl.exp(b_g)[:, None]
b_kb = tl.load(kbar_ptr + (idx + o_c)[:, None] * K + o_k[None, :])
s0 += tl.dot(tl.trans(b_kb), b_vi.to(tl.bfloat16))
b_g = tl.load(glast_ptr + (i_bh * NT + it) * K + (o_k + BK))
s1 = s1 * tl.exp(b_g)[:, None]
b_kb = tl.load(kbar_ptr + (idx + o_c)[:, None] * K + (o_k + BK)[None, :])
s1 += tl.dot(tl.trans(b_kb), b_vi.to(tl.bfloat16))
b_g = tl.load(glast_ptr + (i_bh * NT + it) * K + (o_k + 2 * BK))
s2 = s2 * tl.exp(b_g)[:, None]
b_kb = tl.load(kbar_ptr + (idx + o_c)[:, None] * K + (o_k + 2 * BK)[None, :])
s2 += tl.dot(tl.trans(b_kb), b_vi.to(tl.bfloat16))
b_g = tl.load(glast_ptr + (i_bh * NT + it) * K + (o_k + 3 * BK))
s3 = s3 * tl.exp(b_g)[:, None]
b_kb = tl.load(kbar_ptr + (idx + o_c)[:, None] * K + (o_k + 3 * BK)[None, :])
s3 += tl.dot(tl.trans(b_kb), b_vi.to(tl.bfloat16))
_WS: dict = {}
def _workspace(B, T, H, Kdim, Vdim, BT, device):
key = (B, T, H, Kdim, Vdim, BT, device)
ws = _WS.get(key)
if ws is not None:
return ws
NT = T // BT
N = B * H * NT
ws = {
"A": torch.empty(N, BT, BT, device=device, dtype=torch.float32),
"w": torch.empty(N * BT, Kdim, device=device, dtype=torch.bfloat16),
"u": torch.empty(N * BT, Vdim, device=device, dtype=torch.bfloat16),
"aqk": torch.empty(N * BT, BT, device=device, dtype=torch.bfloat16),
"qg": torch.empty(N * BT, Kdim, device=device, dtype=torch.bfloat16),
"kbar": torch.empty(N * BT, Kdim, device=device, dtype=torch.bfloat16),
"glast": torch.empty(N, Kdim, device=device, dtype=torch.float32),
"o": torch.empty(B, T, H, Vdim, device=device, dtype=torch.bfloat16),
}
_WS[key] = ws
return ws
def _kda_forward(q, k, v, g, beta, scale, chunk_size):
B, T, H, Kdim = q.shape
Vdim = v.shape[-1]
BT = chunk_size
NT = T // BT
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
g = g.contiguous()
beta = beta.contiguous()
ws = _workspace(B, T, H, Kdim, Vdim, BT, q.device)
grid = (NT, B * H)
_build_A_kernel[grid](k, g, beta, ws["A"], T, H, NT, BT=BT, BK=64, num_warps=4, num_stages=2)
_invert_A_kernel[grid](ws["A"], beta, T, H, NT, BT=BT, num_warps=4, num_stages=1)
_build_wy_kernel[grid](
q, k, v, g, ws["A"],
ws["w"], ws["u"], ws["aqk"], ws["qg"], ws["kbar"], ws["glast"],
float(scale), T, H, NT, BT=BT, BK=64, num_warps=4, num_stages=2,
)
_inter_kernel[(B * H, 4)](
ws["w"], ws["u"], ws["aqk"], ws["qg"], ws["kbar"], ws["glast"], ws["o"],
T, H, NT, BT=BT, BK=32, BV=32, num_warps=4, num_stages=2,
)
return ws["o"]
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
self.register_buffer("_dummy", torch.zeros(1), persistent=False)
self._graph = None
self._graph_ptrs = None
def forward(self, q, k, v, g, beta):
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
g = g.contiguous()
beta = beta.contiguous()
ptrs = (q.data_ptr(), k.data_ptr(), v.data_ptr(), g.data_ptr(), beta.data_ptr())
if self._graph is None or ptrs != self._graph_ptrs:
out = _kda_forward(q, k, v, g, beta, scale=self.scale, chunk_size=self.chunk_size)
try:
_kda_forward(q, k, v, g, beta, scale=self.scale, chunk_size=self.chunk_size)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
_kda_forward(q, k, v, g, beta, scale=self.scale, chunk_size=self.chunk_size)
self._graph = graph
self._graph_ptrs = ptrs
except Exception:
self._graph = None
self._graph_ptrs = None
return out
self._graph.replay()
ws = _workspace(q.shape[0], q.shape[1], q.shape[2], q.shape[3], v.shape[-1], self.chunk_size, q.device)
return ws["o"]
B = 2
T = 1024
H = 8
K = 128
V = 128
CHUNK_SIZE = 64
def get_inputs():
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
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]
20260813_072822_grok_grok-4.6_02_kda_cutlass