KernelBench hard · RTX PRO 6000

KDA CUTLASS Claude Fable 5

0.84%geomean peak fraction across shapes

manually audited: clean

Genuine from-scratch chunked KDA (Kimi Delta Attention) forward in eight custom Triton kernels: chunk-local decay cumsum with per-channel max anchoring, delta-rule WY triangular inverse via the nilpotent doubling identity, per-chunk affine transition operators (An, bn), a two-level parallel composition tree, a single short serial scan (NT/4 tensor-core steps), and a parallel output replay. No forbidden ops (no fla.ops.kda / chunk_kda; "fla" grep hits are the "flat" buffer name), no reference import, no grader tampering. The per-shape CUDA graph is launch-overhead elimination, not caching: inputs are copied into static buffers on every call before replay, and the empirical test proves live recompute. Rebench (clean sequential re-grade) peak_fraction 0.0084 vs original in-run 0.0082; both PASS check. The tiny fraction is structural for this problem, not weakness: 0.10-0.18 ms wall per shape at 220-300 GB/s, graded against the B200 dense-bf16 tensor peak with a chunked-linear-attention FLOP formula the op never executes densely.

harnessor-fable (containerized, live CUDA, B200)
Kernel source (redacted)
"""Kimi Delta Attention (KDA) forward, chunk form — custom Triton kernels for B200 (SM100).

Written from scratch against the chunk-parallel formulation in reference.py.

The inter-chunk recurrence is linear in the K x V state S:

    S_{n+1} = diag(exp(gc_n[-1])) @ S_n + kg_n^T @ (u_n - w_n @ S_n)
            = An @ S_n + bn,   An = diag(dlast_n) - kg_n^T @ w_n  (K x K)
                               bn = kg_n^T @ u_n                  (K x V)
    o_n     = qg_n @ S_n + Aqk_n @ (u_n - w_n @ S_n)
            = (qg_n - Aqk_n @ w_n) @ S_n + Aqk_n @ u_n = qgp_n @ S_n + ou_n

so everything except one short scan is embarrassingly parallel over chunks:

  k_pre      chunk-local cumsum of g + every decay factor, derived from one
             exp tile + one reciprocal via a per-channel max anchor.
  k_A        decay-weighted grams kp @ kn^T and qp @ kn^T (per-channel max
             anchor cancels), then the delta-rule WY transform
             (I - tril(beta k k^T))^{-1} via the nilpotent doubling identity
             (I-M)^{-1} = prod (I + M^{2^k}).
  k_wu       w = A @ (exp(gc) k), u = A @ v.
  k_trans    An (dense, diagonal folded) and bn per chunk.
  k_op       qgp = qg - Aqk @ w, ou = Aqk @ u.
  k_compose  two parallel tree levels build 4-chunk composed operators
             A4 = A3 A2 A1 A0, b4 = A3(A2(A1 b0 + b1) + b2) + b3.
  k_serial   the only sequential kernel: NT/4 steps of S = A4 @ S + b4
             (one dependent tensor-core dot per step), bf16 group snapshots.
  k_output   parallel per group: replay the <=3 intra-group hops with An/bn
             and emit o = qgp @ S_n + ou for the 4 chunks.

Model.forward wraps the pipeline in a per-shape CUDA graph: inputs are copied
into static buffers on EVERY call and the graph replays the kernel pipeline,
so each call fully recomputes from the live input values (no result caching).
The graph only removes CPU launch overhead. KDA_DISABLE_GRAPH=1 runs eagerly.
"""
from __future__ import annotations

import os

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"]


# ---------------------------------------------------------------------------
# k_pre: chunk-local cumsum + all decay factors
# ---------------------------------------------------------------------------
@triton.jit
def _kda_pre(
    q_ptr, k_ptr, g_ptr,
    kp_ptr, kn_ptr, qp_ptr, em_ptr, elm_ptr, dlast_ptr,
    H, T, NT, scale,
    BT: tl.constexpr, K: tl.constexpr,
):
    pid_n = tl.program_id(0)
    pid_bh = tl.program_id(1)
    b = pid_bh // H
    h = pid_bh % H

    offs_c = tl.arange(0, BT)
    offs_k = tl.arange(0, K)
    row = (b * T + pid_n * BT + offs_c) * H + h
    ptrs = row[:, None] * K + offs_k[None, :]

    g = tl.load(g_ptr + ptrs)                          # (BT, K) fp32
    gc = tl.cumsum(g, axis=0)                          # chunk-local cumsum
    gl = tl.sum(g, axis=0)                             # last cumsum row, exact
    m = tl.max(gc, axis=0)                             # (K,) per-channel anchor
    ep = tl.exp(gc - m[None, :])                       # <= 1
    en = 1.0 / ep                                      # exp(m - gc)
    em = tl.exp(m)
    egl = tl.exp(gl)
    elm = egl / em                                     # exp(gl - m)

    kf = tl.load(k_ptr + ptrs).to(tl.float32)
    qf = tl.load(q_ptr + ptrs).to(tl.float32) * scale

    # Downstream kernels rebuild the remaining decay products per channel:
    #   gk = kp * em, qg = qp * em, kg = kn * elm.
    tl.store(kp_ptr + ptrs, (kf * ep).to(tl.bfloat16))
    tl.store(kn_ptr + ptrs, (kf * en).to(tl.bfloat16))
    tl.store(qp_ptr + ptrs, (qf * ep).to(tl.bfloat16))
    kbase = (pid_bh * NT + pid_n) * K + offs_k
    tl.store(em_ptr + kbase, em)
    tl.store(elm_ptr + kbase, elm)
    tl.store(dlast_ptr + kbase, egl)


# ---------------------------------------------------------------------------
# k_A: triangular inverse (WY transform); axis2 program 1 computes Aqk instead
# ---------------------------------------------------------------------------
@triton.jit
def _kda_A(
    kp_ptr, kn_ptr, qp_ptr, beta_ptr, a_ptr, aqk_ptr,
    H, T, NT,
    BT: tl.constexpr, K: tl.constexpr,
):
    pid_n = tl.program_id(0)
    pid_bh = tl.program_id(1)
    which = tl.program_id(2)
    b = pid_bh // H
    h = pid_bh % H

    offs_c = tl.arange(0, BT)
    offs_k = tl.arange(0, K)
    row = (b * T + pid_n * BT + offs_c) * H + h
    ptrs = row[:, None] * K + offs_k[None, :]

    kn = tl.load(kn_ptr + ptrs)
    knt = tl.trans(kn)
    base = (pid_bh * NT + pid_n) * BT * BT
    idx = base + offs_c[:, None] * BT + offs_c[None, :]

    if which == 1:
        qp = tl.load(qp_ptr + ptrs)
        aqk = tl.dot(qp, knt)                          # (BT, BT) fp32
        aqk = tl.where(offs_c[:, None] >= offs_c[None, :], aqk, 0.0)
        tl.store(aqk_ptr + idx, aqk.to(tl.bfloat16))
    else:
        kp = tl.load(kp_ptr + ptrs)                    # (BT, K) bf16
        afull = tl.dot(kp, knt)                        # (BT, BT) fp32
        beta = tl.load(beta_ptr + row).to(tl.float32)  # (BT,)
        strict_lower = offs_c[:, None] > offs_c[None, :]
        mm = tl.where(strict_lower, -(afull * beta[:, None]), 0.0)

        # (I-M)^{-1} = (I+M)(I+M^2)(I+M^4)(I+M^8)(I+M^16)(I+M^32); M^64 = 0.
        eye = tl.where(offs_c[:, None] == offs_c[None, :], 1.0, 0.0)
        x = eye + mm
        p = tl.dot(mm, mm)
        for _ in range(4):
            x = x + tl.dot(x, p)
            p = tl.dot(p, p)
        x = x + tl.dot(x, p)

        ab = (x * beta[None, :]).to(tl.bfloat16)       # column scale by beta
        tl.store(a_ptr + idx, ab)


# ---------------------------------------------------------------------------
# k_wu: axis2 program 0: w = A @ gk;  program 1: u = A @ v
# ---------------------------------------------------------------------------
@triton.jit
def _kda_wu(
    a_ptr, kp_ptr, em_ptr, v_ptr, w_ptr, u_ptr,
    H, T, NT,
    BT: tl.constexpr, K: tl.constexpr, V: tl.constexpr,
):
    pid_n = tl.program_id(0)
    pid_bh = tl.program_id(1)
    which = tl.program_id(2)
    b = pid_bh // H
    h = pid_bh % H

    offs_c = tl.arange(0, BT)
    offs_k = tl.arange(0, K)
    row = (b * T + pid_n * BT + offs_c) * H + h

    a_base = (pid_bh * NT + pid_n) * BT * BT
    a = tl.load(a_ptr + a_base + offs_c[:, None] * BT + offs_c[None, :])
    ptr = row[:, None] * K + offs_k[None, :]
    if which == 0:
        kp = tl.load(kp_ptr + ptr).to(tl.float32)
        em = tl.load(em_ptr + (pid_bh * NT + pid_n) * K + offs_k)
        gk = (kp * em[None, :]).to(tl.bfloat16)        # k * exp(gc)
        w = tl.dot(a, gk)
        tl.store(w_ptr + ptr, w.to(tl.bfloat16))
    else:
        v = tl.load(v_ptr + ptr)
        u = tl.dot(a, v)
        tl.store(u_ptr + ptr, u.to(tl.bfloat16))


# ---------------------------------------------------------------------------
# k_trans: An = diag(dlast) - kg^T @ w (dense), bn = kg^T @ u
#   grid axis2: (K row block) x (an | bn)
# ---------------------------------------------------------------------------
@triton.jit
def _kda_trans(
    w_ptr, u_ptr, kn_ptr, elm_ptr, dlast_ptr, an_ptr, bn_ptr,
    H, T, NT,
    BT: tl.constexpr, K: tl.constexpr, V: tl.constexpr, BKR: tl.constexpr,
):
    pid_n = tl.program_id(0)
    pid_bh = tl.program_id(1)
    pid_z = tl.program_id(2)
    pid_kb = pid_z % (K // BKR)
    which = pid_z // (K // BKR)
    b = pid_bh // H
    h = pid_bh % H

    offs_c = tl.arange(0, BT)
    offs_k = tl.arange(0, K)
    offs_kr = pid_kb * BKR + tl.arange(0, BKR)
    row = (b * T + pid_n * BT + offs_c) * H + h

    knr = tl.load(kn_ptr + row[:, None] * K + offs_kr[None, :]).to(tl.float32)
    elmr = tl.load(elm_ptr + (pid_bh * NT + pid_n) * K + offs_kr)
    kgt = (knr * elmr[None, :]).to(tl.bfloat16)        # k * exp(gl - gc), (BT, BKR)
    base = (pid_bh * NT + pid_n) * K + pid_kb * BKR
    if which == 0:
        w = tl.load(w_ptr + row[:, None] * K + offs_k[None, :])   # (BT, K)
        akw = tl.dot(tl.trans(kgt), w)                            # (BKR, K)
        dlr = tl.load(dlast_ptr + (pid_bh * NT + pid_n) * K + offs_kr)
        an = tl.where(offs_kr[:, None] == offs_k[None, :], dlr[:, None] - akw, -akw)
        tl.store(an_ptr + base * K + tl.arange(0, BKR)[:, None] * K + offs_k[None, :],
                 an.to(tl.bfloat16))
    else:
        u = tl.load(u_ptr + row[:, None] * V + offs_k[None, :])   # (BT, V)
        bn = tl.dot(tl.trans(kgt), u)                             # (BKR, V)
        tl.store(bn_ptr + base * V + tl.arange(0, BKR)[:, None] * V + offs_k[None, :],
                 bn.to(tl.bfloat16))


# ---------------------------------------------------------------------------
# k_op: axis2 program 0: qgp = qg - Aqk @ w;  program 1: ou = Aqk @ u
# ---------------------------------------------------------------------------
@triton.jit
def _kda_op(
    w_ptr, u_ptr, qp_ptr, em_ptr, aqk_ptr, qgp_ptr, ou_ptr,
    H, T, NT,
    BT: tl.constexpr, K: tl.constexpr, V: tl.constexpr,
):
    pid_n = tl.program_id(0)
    pid_bh = tl.program_id(1)
    which = tl.program_id(2)
    b = pid_bh // H
    h = pid_bh % H

    offs_c = tl.arange(0, BT)
    offs_k = tl.arange(0, K)
    row = (b * T + pid_n * BT + offs_c) * H + h
    ptr = row[:, None] * K + offs_k[None, :]

    base = (pid_bh * NT + pid_n) * BT * BT
    aqk = tl.load(aqk_ptr + base + offs_c[:, None] * BT + offs_c[None, :])
    if which == 0:
        w = tl.load(w_ptr + ptr)
        qp = tl.load(qp_ptr + ptr).to(tl.float32)
        em = tl.load(em_ptr + (pid_bh * NT + pid_n) * K + offs_k)
        qg = qp * em[None, :]                          # q * scale * exp(gc)
        qgp = qg - tl.dot(aqk, w)
        tl.store(qgp_ptr + ptr, qgp.to(tl.bfloat16))
    else:
        u = tl.load(u_ptr + ptr)
        ou = tl.dot(aqk, u)
        tl.store(ou_ptr + ptr, ou.to(tl.bfloat16))


# ---------------------------------------------------------------------------
# k_compose: one tree level; group pairs (lo, hi) -> hi ∘ lo
#   Aout = Ahi @ Alo, bout = Ahi @ blo + bhi.  Column-split for registers.
# ---------------------------------------------------------------------------
@triton.jit
def _kda_compose(
    alo_ptr, blo_ptr, ahi_ptr, bhi_ptr, aout_ptr, bout_ptr,
    NPAIR_PER_BH,
    K: tl.constexpr, V: tl.constexpr, BC: tl.constexpr,
):
    pid_p = tl.program_id(0)     # pair index within bh
    pid_bh = tl.program_id(1)
    pid_c = tl.program_id(2)     # column block
    offs_k = tl.arange(0, K)
    offs_c = pid_c * BC + tl.arange(0, BC)

    lo = (pid_bh * 2 * NPAIR_PER_BH + 2 * pid_p)
    hi = lo + 1
    out = pid_bh * NPAIR_PER_BH + pid_p

    ahi = tl.load(ahi_ptr + hi * K * K + offs_k[:, None] * K + offs_k[None, :])
    alo_c = tl.load(alo_ptr + lo * K * K + offs_k[:, None] * K + offs_c[None, :])
    aout = tl.dot(ahi, alo_c)
    tl.store(aout_ptr + out * K * K + offs_k[:, None] * K + offs_c[None, :],
             aout.to(tl.bfloat16))

    blo_c = tl.load(blo_ptr + lo * K * V + offs_k[:, None] * V + offs_c[None, :])
    bhi_c = tl.load(bhi_ptr + hi * K * V + offs_k[:, None] * V + offs_c[None, :]).to(tl.float32)
    bout = tl.dot(ahi, blo_c) + bhi_c
    tl.store(bout_ptr + out * K * V + offs_k[:, None] * V + offs_c[None, :],
             bout.to(tl.bfloat16))


# ---------------------------------------------------------------------------
# k_serial: S <- A4 @ S + b4 over NT/4 groups; bf16 snapshots per group
# ---------------------------------------------------------------------------
@triton.jit
def _kda_serial(
    a4_ptr, b4_ptr, snap_ptr,
    NG,
    K: tl.constexpr, V: tl.constexpr, BV: tl.constexpr, NSTAGE: tl.constexpr,
):
    pid_bh = tl.program_id(0)
    pid_v = tl.program_id(1)
    offs_k = tl.arange(0, K)
    offs_v = pid_v * BV + tl.arange(0, BV)

    s = tl.zeros((K, BV), dtype=tl.float32)
    for n in tl.range(NG, num_stages=NSTAGE):
        base = pid_bh * NG + n
        sb = s.to(tl.bfloat16)
        tl.store(snap_ptr + base * K * V + offs_k[:, None] * V + offs_v[None, :], sb)
        a4 = tl.load(a4_ptr + base * K * K + offs_k[:, None] * K + offs_k[None, :])
        b4 = tl.load(b4_ptr + base * K * V + offs_k[:, None] * V + offs_v[None, :]).to(tl.float32)
        s = b4 + tl.dot(a4, sb)


# ---------------------------------------------------------------------------
# k_output: per 4-chunk group, replay intra-group hops and emit o
# ---------------------------------------------------------------------------
@triton.jit
def _kda_output(
    qgp_ptr, ou_ptr, an_ptr, bn_ptr, snap_ptr, o_ptr,
    H, T, NT, NG,
    BT: tl.constexpr, K: tl.constexpr, V: tl.constexpr, BV: tl.constexpr,
    L: tl.constexpr,
):
    pid_g = tl.program_id(0)     # group index within bh
    pid_bhv = tl.program_id(1)
    NV: tl.constexpr = V // BV
    pid_bh = pid_bhv // NV
    pid_v = pid_bhv % NV
    b = pid_bh // H
    h = pid_bh % H

    offs_c = tl.arange(0, BT)
    offs_k = tl.arange(0, K)
    offs_v = pid_v * BV + tl.arange(0, BV)

    sbase = pid_bh * NG + pid_g
    sb = tl.load(snap_ptr + sbase * K * V + offs_k[:, None] * V + offs_v[None, :])

    for j in tl.static_range(L):
        n = pid_g * L + j
        row = (b * T + n * BT + offs_c) * H + h
        qgp = tl.load(qgp_ptr + row[:, None] * K + offs_k[None, :])
        ou = tl.load(ou_ptr + row[:, None] * V + offs_v[None, :]).to(tl.float32)
        o = tl.dot(qgp, sb) + ou
        tl.store(o_ptr + row[:, None] * V + offs_v[None, :], o.to(tl.bfloat16))
        if j < L - 1:
            nbase = pid_bh * NT + n
            an = tl.load(an_ptr + nbase * K * K + offs_k[:, None] * K + offs_k[None, :])
            bn = tl.load(bn_ptr + nbase * K * V + offs_k[:, None] * V + offs_v[None, :]).to(tl.float32)
            sb = (bn + tl.dot(an, sb)).to(tl.bfloat16)


# ---------------------------------------------------------------------------
# Python driver
# ---------------------------------------------------------------------------
_BT = 64
_BV_SER = int(os.environ.get("KDA_BV_SER", "16"))
_NSTAGE = 3
_BV_OUT = 64


class _ShapeState:
    """Per-shape static workspace + optional CUDA graph."""

    def __init__(self, B, T, H, K, V, scale, device):
        self.B, self.T, self.H, self.K, self.V = B, T, H, K, V
        self.scale = scale
        NT = T // _BT
        # Composed-group size: the serial scan runs NT/L steps; the output pass
        # replays L-1 intra-group hops in parallel. Deeper composition only
        # pays off when the scan is long.
        L = int(os.environ.get("KDA_L", "0")) or (4 if NT >= 64 else 2)
        NG = NT // L
        self.NT, self.NG, self.L = NT, NG, L
        BH = B * H
        dt = torch.bfloat16
        dev = device
        e = torch.empty
        # q/k/v/beta live in one flat bf16 buffer so one cat() refreshes them.
        nq = B * T * H * K
        nv = B * T * H * V
        nb = B * T * H
        self.flat = e(nq + nq + nv + nb, dtype=dt, device=dev)
        self.q = self.flat[:nq].view(B, T, H, K)
        self.k = self.flat[nq:2 * nq].view(B, T, H, K)
        self.v = self.flat[2 * nq:2 * nq + nv].view(B, T, H, V)
        self.beta = self.flat[2 * nq + nv:].view(B, T, H)
        self.g = e(B, T, H, K, dtype=torch.float32, device=dev)
        self.kp = e(B, T, H, K, dtype=dt, device=dev)
        self.kn = e(B, T, H, K, dtype=dt, device=dev)
        self.qp = e(B, T, H, K, dtype=dt, device=dev)
        self.em = e(BH, NT, K, dtype=torch.float32, device=dev)
        self.elm = e(BH, NT, K, dtype=torch.float32, device=dev)
        self.dlast = e(BH, NT, K, dtype=torch.float32, device=dev)
        self.a = e(BH, NT, _BT, _BT, dtype=dt, device=dev)
        self.aqk = e(BH, NT, _BT, _BT, dtype=dt, device=dev)
        self.w = e(B, T, H, K, dtype=dt, device=dev)
        self.u = e(B, T, H, V, dtype=dt, device=dev)
        self.an = e(BH, NT, K, K, dtype=dt, device=dev)
        self.bn = e(BH, NT, K, V, dtype=dt, device=dev)
        self.a2 = e(BH, NT // 2, K, K, dtype=dt, device=dev)
        self.b2 = e(BH, NT // 2, K, V, dtype=dt, device=dev)
        if L == 4:
            self.a4 = e(BH, NG, K, K, dtype=dt, device=dev)
            self.b4 = e(BH, NG, K, V, dtype=dt, device=dev)
        else:
            self.a4, self.b4 = self.a2, self.b2
        self.qgp = e(B, T, H, K, dtype=dt, device=dev)
        self.ou = e(B, T, H, V, dtype=dt, device=dev)
        self.snap = e(BH, NG, K, V, dtype=dt, device=dev)
        self.o = e(B, T, H, V, dtype=dt, device=dev)
        self.graph = None

    def launch(self):
        B, T, H, K, V, NT, NG = self.B, self.T, self.H, self.K, self.V, self.NT, self.NG
        BH = B * H
        _kda_pre[(NT, BH)](
            self.q, self.k, self.g,
            self.kp, self.kn, self.qp, self.em, self.elm, self.dlast,
            H, T, NT, self.scale, BT=_BT, K=K, num_warps=4)
        _kda_A[(NT, BH, 2)](
            self.kp, self.kn, self.qp, self.beta, self.a, self.aqk,
            H, T, NT, BT=_BT, K=K, num_warps=4)
        _kda_wu[(NT, BH, 2)](
            self.a, self.kp, self.em, self.v, self.w, self.u,
            H, T, NT, BT=_BT, K=K, V=V, num_warps=4)
        _kda_op[(NT, BH, 2)](
            self.w, self.u, self.qp, self.em, self.aqk, self.qgp, self.ou,
            H, T, NT, BT=_BT, K=K, V=V, num_warps=4)
        _kda_trans[(NT, BH, 4)](
            self.w, self.u, self.kn, self.elm, self.dlast, self.an, self.bn,
            H, T, NT, BT=_BT, K=K, V=V, BKR=K // 2, num_warps=4)
        _kda_compose[(NT // 2, BH, 2)](
            self.an, self.bn, self.an, self.bn, self.a2, self.b2,
            NT // 2, K=K, V=V, BC=K // 2, num_warps=8)
        if self.L == 4:
            _kda_compose[(NG, BH, 2)](
                self.a2, self.b2, self.a2, self.b2, self.a4, self.b4,
                NG, K=K, V=V, BC=K // 2, num_warps=8)
        _kda_serial[(BH, V // _BV_SER)](
            self.a4, self.b4, self.snap,
            NG, K=K, V=V, BV=_BV_SER, NSTAGE=_NSTAGE, num_warps=4)
        _kda_output[(NG, BH * (V // _BV_OUT))](
            self.qgp, self.ou, self.an, self.bn, self.snap, self.o,
            H, T, NT, NG, BT=_BT, K=K, V=V, BV=_BV_OUT, L=self.L, num_warps=4)

    def run(self, q, k, v, g, beta, use_graph):
        # Fresh input values are copied into the static buffers on EVERY call;
        # the kernels then recompute the output from those values.
        self.g.copy_(g)
        torch.cat(
            [q.reshape(-1), k.reshape(-1), v.reshape(-1), beta.reshape(-1)],
            out=self.flat,
        )
        if not use_graph:
            self.launch()
            return self.o
        if self.graph is None:
            torch.cuda.synchronize()
            for _ in range(3):
                self.launch()          # Triton JIT + allocator warmup
            torch.cuda.synchronize()
            gr = torch.cuda.CUDAGraph()
            with torch.cuda.graph(gr):
                self.launch()
            self.graph = gr
        self.graph.replay()
        return self.o


_STATES: dict = {}


def kda_chunk_forward(q, k, v, g, beta, scale, chunk_size=64):
    B, T, H, K = q.shape
    V = v.shape[-1]
    assert chunk_size == _BT and T % (_BT * 2) == 0 and K == V
    key = (B, T, H, K, V, q.device.index)
    st = _STATES.get(key)
    if st is None:
        st = _ShapeState(B, T, H, K, V, scale, q.device)
        _STATES[key] = st
    use_graph = os.environ.get("KDA_DISABLE_GRAPH", "0") != "1"
    return st.run(q, k, v, g, beta, use_graph)


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)

    def forward(self, q, k, v, g, beta):
        return kda_chunk_forward(q, k, v, g, beta, scale=self.scale, chunk_size=self.chunk_size)


# 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():
    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]

20260719_031422_or-fable_anthropic_claude-fable-5_02_kda_cutlass