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 tokensregimecompute

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