KernelBench hard · B200

KDA CUTLASS Claude Fable 5

0.84%geomean peak fraction across shapes

manually audited: clean

harnessor-fableagent session2h 7mtotal wall2h 9mcheck2mbenchmark18soutput tokensgpu-lock wait1h 18mgpu-lock held9mregimecompute

Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth

2×1024×8×128×128×640.113 ms0.9%0.22 TB/s · 3% of 8.0 TB/s HBM · also 19 TFLOPS (1% of compute)
2×2048×8×128×128×640.169 ms1.1%0.30 TB/s · 4% of 8.0 TB/s HBM · also 25 TFLOPS (1% of compute)
1×4096×8×128×128×640.179 ms1.1%0.28 TB/s · 4% of 8.0 TB/s HBM · also 24 TFLOPS (1% of compute)
1×2048×4×128×128×640.099 ms0.5%0.13 TB/s · 2% of 8.0 TB/s HBM · also 11 TFLOPS (0% of compute)

compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)

geomean(0.9% · 1.1% · 1.1% · 0.5%) = 0.8%

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