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