KernelBench hard · RTX PRO 6000

KDA CUTLASS GLM-5.3 Flash

2.26%geomean peak fraction across shapes

manually audited: clean

Five Triton kernels implement chunk-form Kimi Delta Attention: blocked 16x16 unit-lower inverse, segment affine transitions, outer scan, then a parallel output correction. Isolated regrade 0.0226 on RTX PRO 6000 (in-run 0.0219). Workspace is shape-keyed scratch; each forward allocates a fresh output and launches all five kernels. Honest LOW; launch-bound.

harnessor-fableagent session2h 7mtotal wall2h 7mcheck7sbenchmark2soutput tokensregimecompute

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

2×1024×8×128×128×640.192 ms2.2%0.13 TB/s · 7% of 1.8 TB/s HBM · also 11 TFLOPS (2% of compute)
2×2048×8×128×128×640.297 ms2.9%0.17 TB/s · 9% of 1.8 TB/s HBM · also 14 TFLOPS (3% of compute)
1×4096×8×128×128×640.327 ms2.6%0.15 TB/s · 9% of 1.8 TB/s HBM · also 13 TFLOPS (3% of compute)
1×2048×4×128×128×640.141 ms1.5%0.09 TB/s · 5% of 1.8 TB/s HBM · also 8 TFLOPS (2% of compute)

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

geomean(2.2% · 2.9% · 2.6% · 1.5%) = 2.3%

Kernel source (redacted)
"""Chunk-parallel Kimi Delta Attention forward for SM120 (RTX PRO 6000 Blackwell).

Written from the math in reference.py as three Triton kernels:

  1. _prep_wu_kernel -- per (chunk, batch*head): builds the strictly-lower
     intra-chunk matrix L[c,i] = -beta_c * <k_c e^g_c, k_i e^-g_i>, forms the
     unit lower-triangular system N = I - L, and solves N W = diag(beta) E by
     16x16-blocked forward substitution, emitting
         w = M @ (e^g k),  u = M @ v,   M = N^{-1} diag(beta).
     The inverse of each 16x16 diagonal block is computed with an exact
     row-recurrence; block combinations are tensor-core dots.
  The chunk recurrence is affine in the state, so it is solved at two levels:
  chunks are grouped into segments of W; each segment's net transition
      S_out = A @ S_in + b
  is built exactly (A as a dense K x K matrix), a short outer scan walks the
  segments, and per-chunk outputs are then finished in a fully parallel pass
  corrected by the checkpointed state at each segment start. Sequential depth
  drops from NT chunk steps to W + NT/W.

  1. _prep_wu_kernel -- per (chunk, batch*head): builds the strictly-lower
     intra-chunk matrix L[c,i] = -beta_c * <k_c e^g_c, k_i e^-g_i>, forms the
     unit lower-triangular system N = I - L, and solves N W = diag(beta) E by
     16x16-blocked forward substitution, emitting
         w = M @ (e^g k),  u = M @ v,   M = N^{-1} diag(beta).
     Also emits decay-normalized k/q (kg, qg) and per-chunk decay stats.
  2. _seg_a_kernel -- per (segment, batch*head): builds A = diag(e^{Lam}) -
     sum_n kgR_n^T (w_n D_n) and per-chunk prefix/suffix decay factors.
  3. _seg_b_kernel -- per (segment, batch*head, V-tile): runs the local chunk
     chain from a zero seed, emitting output partials o_loc and b (V-slice).
  4. _outer_scan_kernel -- per (batch*head, V-tile): scans segment transitions,
     leaving checkpoint states at every segment start.
  5. _corr_out_kernel -- per (chunk, batch*head, V-tile), fully parallel:
         o = o_loc + (q e^g D_n) @ cp - tril((scale q e^g)(k e^-g)^T) @ ((w D_n) @ cp)
"""
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 _inv_unit_lower16(z):
    """Return y with (I - z) @ y = I, for strictly-lower 16x16 fp32 z."""
    o = tl.arange(0, 16)
    y = (o[:, None] == o[None, :]).to(tl.float32)
    for i in tl.static_range(1, 16):
        sel = o == i
        zi = tl.where(sel[:, None] & (o[None, :] < i), z, 0.0)
        zv = tl.sum(zi, axis=0)
        row = sel.to(tl.float32) + tl.sum(zv[:, None] * y, axis=0)
        y = tl.where(sel[:, None], row[None, :], y)
    return y


@triton.jit
def _prep_wu_kernel(
    q_ptr, k_ptr, v_ptr, g_ptr, beta_ptr,
    w_ptr, u_ptr, kg_ptr, qg_ptr, gl_ptr, eg_ptr, ei_ptr,
    scale, T,
    H: tl.constexpr, K: tl.constexpr, V: tl.constexpr,
    BT: tl.constexpr, BC: tl.constexpr, NSUB: tl.constexpr,
):
    i_t = tl.program_id(0)
    i_bh = tl.program_id(1)
    i_b = i_bh // H
    i_h = i_bh % H

    o_c = tl.arange(0, BC)
    o_k = tl.arange(0, K)
    o_v = tl.arange(0, V)
    t0 = i_t * BT

    # --- load the four token sub-blocks of k / g / beta -----------------------
    off0 = (i_b * T + t0 + 0 * BC + o_c) * H + i_h
    off1 = (i_b * T + t0 + 1 * BC + o_c) * H + i_h
    off2 = (i_b * T + t0 + 2 * BC + o_c) * H + i_h
    off3 = (i_b * T + t0 + 3 * BC + o_c) * H + i_h

    bk0 = tl.load(k_ptr + off0[:, None] * K + o_k[None, :]).to(tl.float32)
    bk1 = tl.load(k_ptr + off1[:, None] * K + o_k[None, :]).to(tl.float32)
    bk2 = tl.load(k_ptr + off2[:, None] * K + o_k[None, :]).to(tl.float32)
    bk3 = tl.load(k_ptr + off3[:, None] * K + o_k[None, :]).to(tl.float32)
    bg0 = tl.load(g_ptr + off0[:, None] * K + o_k[None, :])
    bg1 = tl.load(g_ptr + off1[:, None] * K + o_k[None, :])
    bg2 = tl.load(g_ptr + off2[:, None] * K + o_k[None, :])
    bg3 = tl.load(g_ptr + off3[:, None] * K + o_k[None, :])
    bt0 = tl.load(beta_ptr + off0).to(tl.float32)
    bt1 = tl.load(beta_ptr + off1).to(tl.float32)
    bt2 = tl.load(beta_ptr + off2).to(tl.float32)
    bt3 = tl.load(beta_ptr + off3).to(tl.float32)

    # Per-chunk decay normalizers shared by the state and output passes.
    # g is a per-token log-decay: form the in-chunk cumulative sum first.
    # (Computed as triangular matmuls -- tl.cumsum is miscompiled on this
    # Triton/Blackwell combination when several scans share a kernel.)
    ltri = (o_c[:, None] >= o_c[None, :]).to(tl.float32)
    acc = tl.zeros((K,), dtype=tl.float32)
    ga0 = tl.dot(ltri, bg0, input_precision="ieee") + acc[None, :]
    acc += tl.sum(bg0, 0)
    ga1 = tl.dot(ltri, bg1, input_precision="ieee") + acc[None, :]
    acc += tl.sum(bg1, 0)
    ga2 = tl.dot(ltri, bg2, input_precision="ieee") + acc[None, :]
    acc += tl.sum(bg2, 0)
    ga3 = tl.dot(ltri, bg3, input_precision="ieee") + acc[None, :]
    gl = acc + tl.sum(bg3, 0)
    pgl = gl_ptr + (i_bh * (T // BT) + i_t) * K + o_k
    tl.store(pgl, gl)

    bq0 = tl.load(q_ptr + off0[:, None] * K + o_k[None, :]).to(tl.float32)
    bq1 = tl.load(q_ptr + off1[:, None] * K + o_k[None, :]).to(tl.float32)
    bq2 = tl.load(q_ptr + off2[:, None] * K + o_k[None, :]).to(tl.float32)
    bq3 = tl.load(q_ptr + off3[:, None] * K + o_k[None, :]).to(tl.float32)

    kn0 = (bk0 * tl.exp(ga0)).to(tl.bfloat16)
    kn1 = (bk1 * tl.exp(ga1)).to(tl.bfloat16)
    kn2 = (bk2 * tl.exp(ga2)).to(tl.bfloat16)
    kn3 = (bk3 * tl.exp(ga3)).to(tl.bfloat16)
    kd0 = (bk0 * tl.exp(-ga0)).to(tl.bfloat16)
    kd1 = (bk1 * tl.exp(-ga1)).to(tl.bfloat16)
    kd2 = (bk2 * tl.exp(-ga2)).to(tl.bfloat16)
    kd3 = (bk3 * tl.exp(-ga3)).to(tl.bfloat16)

    # --- strictly-lower block rows of L (beta-scaled rows) -------------------
    l10 = -tl.dot(kn1, tl.trans(kd0)) * bt1[:, None]
    l20 = -tl.dot(kn2, tl.trans(kd0)) * bt2[:, None]
    l21 = -tl.dot(kn2, tl.trans(kd1)) * bt2[:, None]
    l30 = -tl.dot(kn3, tl.trans(kd0)) * bt3[:, None]
    l31 = -tl.dot(kn3, tl.trans(kd1)) * bt3[:, None]
    l32 = -tl.dot(kn3, tl.trans(kd2)) * bt3[:, None]

    # --- inverses of the diagonal units N_ss = I - L_ss ----------------------
    m_l = o_c[:, None] > o_c[None, :]
    y0 = _inv_unit_lower16(tl.where(m_l, -(tl.dot(kn0, tl.trans(kd0)) * bt0[:, None]), 0.0))
    y1 = _inv_unit_lower16(tl.where(m_l, -(tl.dot(kn1, tl.trans(kd1)) * bt1[:, None]), 0.0))
    y2 = _inv_unit_lower16(tl.where(m_l, -(tl.dot(kn2, tl.trans(kd2)) * bt2[:, None]), 0.0))
    y3 = _inv_unit_lower16(tl.where(m_l, -(tl.dot(kn3, tl.trans(kd3)) * bt3[:, None]), 0.0))

    # --- forward substitution over row-blocks for w = M @ (e^g k) -------------
    wa0 = kn0.to(tl.float32) * bt0[:, None]
    w0 = tl.dot(y0, wa0, input_precision="tf32")
    wa1 = kn1.to(tl.float32) * bt1[:, None] + tl.dot(l10, w0, input_precision="tf32")
    w1 = tl.dot(y1, wa1, input_precision="tf32")
    wa2 = (kn2.to(tl.float32) * bt2[:, None]
           + tl.dot(l20, w0, input_precision="tf32")
           + tl.dot(l21, w1, input_precision="tf32"))
    w2 = tl.dot(y2, wa2, input_precision="tf32")
    wa3 = (kn3.to(tl.float32) * bt3[:, None]
           + tl.dot(l30, w0, input_precision="tf32")
           + tl.dot(l31, w1, input_precision="tf32")
           + tl.dot(l32, w2, input_precision="tf32"))
    w3 = tl.dot(y3, wa3, input_precision="tf32")

    tl.store(w_ptr + off0[:, None] * K + o_k[None, :], w0.to(tl.bfloat16))
    tl.store(w_ptr + off1[:, None] * K + o_k[None, :], w1.to(tl.bfloat16))
    tl.store(w_ptr + off2[:, None] * K + o_k[None, :], w2.to(tl.bfloat16))
    tl.store(w_ptr + off3[:, None] * K + o_k[None, :], w3.to(tl.bfloat16))

    # --- forward substitution over row-blocks for u = M @ v -------------------
    bv0 = tl.load(v_ptr + off0[:, None] * V + o_v[None, :]).to(tl.float32)
    bv1 = tl.load(v_ptr + off1[:, None] * V + o_v[None, :]).to(tl.float32)
    bv2 = tl.load(v_ptr + off2[:, None] * V + o_v[None, :]).to(tl.float32)
    bv3 = tl.load(v_ptr + off3[:, None] * V + o_v[None, :]).to(tl.float32)

    ua0 = bv0 * bt0[:, None]
    u0 = tl.dot(y0, ua0, input_precision="tf32")
    ua1 = bv1 * bt1[:, None] + tl.dot(l10, u0, input_precision="tf32")
    u1 = tl.dot(y1, ua1, input_precision="tf32")
    ua2 = (bv2 * bt2[:, None]
           + tl.dot(l20, u0, input_precision="tf32")
           + tl.dot(l21, u1, input_precision="tf32"))
    u2 = tl.dot(y2, ua2, input_precision="tf32")
    ua3 = (bv3 * bt3[:, None]
           + tl.dot(l30, u0, input_precision="tf32")
           + tl.dot(l31, u1, input_precision="tf32")
           + tl.dot(l32, u2, input_precision="tf32"))
    u3 = tl.dot(y3, ua3, input_precision="tf32")

    tl.store(u_ptr + off0[:, None] * V + o_v[None, :], u0.to(tl.bfloat16))
    tl.store(u_ptr + off1[:, None] * V + o_v[None, :], u1.to(tl.bfloat16))
    tl.store(u_ptr + off2[:, None] * V + o_v[None, :], u2.to(tl.bfloat16))
    tl.store(u_ptr + off3[:, None] * V + o_v[None, :], u3.to(tl.bfloat16))

    # Decay-normalized k (k e^{g_last - g}) and decayed queries for the
    # state / output passes; both consume these instead of fp32 g.
    egl = tl.exp(gl)
    tl.store(eg_ptr + (i_bh * (T // BT) + i_t) * K + o_k, egl)
    tl.store(ei_ptr + (i_bh * (T // BT) + i_t) * K + o_k, tl.exp(-gl))
    kg0 = (bk0 * egl[None, :] * tl.exp(-ga0)).to(tl.bfloat16)
    kg1 = (bk1 * egl[None, :] * tl.exp(-ga1)).to(tl.bfloat16)
    kg2 = (bk2 * egl[None, :] * tl.exp(-ga2)).to(tl.bfloat16)
    kg3 = (bk3 * egl[None, :] * tl.exp(-ga3)).to(tl.bfloat16)
    tl.store(kg_ptr + off0[:, None] * K + o_k[None, :], kg0)
    tl.store(kg_ptr + off1[:, None] * K + o_k[None, :], kg1)
    tl.store(kg_ptr + off2[:, None] * K + o_k[None, :], kg2)
    tl.store(kg_ptr + off3[:, None] * K + o_k[None, :], kg3)

    qg0 = (bq0 * (scale * tl.exp(ga0))).to(tl.bfloat16)
    qg1 = (bq1 * (scale * tl.exp(ga1))).to(tl.bfloat16)
    qg2 = (bq2 * (scale * tl.exp(ga2))).to(tl.bfloat16)
    qg3 = (bq3 * (scale * tl.exp(ga3))).to(tl.bfloat16)
    tl.store(qg_ptr + off0[:, None] * K + o_k[None, :], qg0)
    tl.store(qg_ptr + off1[:, None] * K + o_k[None, :], qg1)
    tl.store(qg_ptr + off2[:, None] * K + o_k[None, :], qg2)
    tl.store(qg_ptr + off3[:, None] * K + o_k[None, :], qg3)


@triton.jit
def _seg_a_kernel(
    w_ptr, kg_ptr, gl_ptr, a_ptr, dn_ptr, sf_ptr,
    T,
    H: tl.constexpr, K: tl.constexpr,
    BT: tl.constexpr, W: tl.constexpr,
):
    """Per (segment, batch*head): build the segment's affine state transition
    A = diag(e^{Lam}) - sum_n kgR_n^T (w_n D_n) as an exact K x K matrix, and
    emit per-chunk decay prefixes dn[n] and suffix factors sf[n]."""
    i_s = tl.program_id(0)
    i_bh = tl.program_id(1)
    NT = T // BT
    o_k = tl.arange(0, K)
    eye = (o_k[:, None] == o_k[None, :]).to(tl.float32)

    base = i_bh * NT
    lam_end = tl.zeros((K,), dtype=tl.float32)
    for r in range(0, W):
        lam_end += tl.load(gl_ptr + (base + i_s * W + r) * K + o_k)

    lam = tl.zeros((K,), dtype=tl.float32)
    A = eye * tl.exp(lam_end)[None, :]
    for r in range(0, W):
        n = i_s * W + r
        rows = n * BT + tl.arange(0, BT)
        off = ((i_bh // H) * T + rows) * H + (i_bh % H)
        gl = tl.load(gl_ptr + (base + n) * K + o_k)
        tl.store(dn_ptr + (base + n) * K + o_k, tl.exp(lam))
        sfac = tl.exp(lam_end - lam - gl)
        tl.store(sf_ptr + (base + n) * K + o_k, sfac)
        b_w = tl.load(w_ptr + off[:, None] * K + o_k[None, :]).to(tl.float32)
        b_kg = tl.load(kg_ptr + off[:, None] * K + o_k[None, :]).to(tl.float32)
        wd = b_w * tl.exp(lam)[None, :]
        kgr = b_kg * sfac[None, :]
        A -= tl.dot(tl.trans(kgr), wd, input_precision="tf32")
        lam += gl
    tl.store(a_ptr + (i_bh * (T // (BT * W)) + i_s) * K * K
             + o_k[:, None] * K + o_k[None, :], A.to(tl.float16))


@triton.jit
def _seg_b_kernel(
    w_ptr, u_ptr, kg_ptr, qg_ptr, eg_ptr, ei_ptr, gl_ptr,
    oloc_ptr, b_ptr,
    T,
    H: tl.constexpr, K: tl.constexpr, V: tl.constexpr,
    BT: tl.constexpr, BV: tl.constexpr, W: tl.constexpr,
):
    """Per (segment, batch*head, V-tile): run the chunk chain from a zero seed.
    Emits the local output partials o_loc and the segment offset term
    b = sum_n kgR_n^T u~_n (this V-slice)."""
    i_s = tl.program_id(0)
    i_bh = tl.program_id(1)
    i_v = tl.program_id(2)
    i_b = i_bh // H
    i_h = i_bh % H
    NT = T // BT

    o_c = tl.arange(0, BT)
    o_k = tl.arange(0, K)
    o_v = i_v * BV + tl.arange(0, BV)

    base = i_bh * NT
    lam_end = tl.zeros((K,), dtype=tl.float32)
    for r in range(0, W):
        lam_end += tl.load(gl_ptr + (base + i_s * W + r) * K + o_k)

    S = tl.zeros((K, BV), dtype=tl.float32)
    bacc = tl.zeros((K, BV), dtype=tl.float32)
    lam = tl.zeros((K,), dtype=tl.float32)
    for r in range(0, W):
        n = i_s * W + r
        rows = n * BT + o_c
        off = (i_b * T + rows) * H + i_h
        b_w = tl.load(w_ptr + off[:, None] * K + o_k[None, :])
        b_u = tl.load(u_ptr + off[:, None] * V + o_v[None, :]).to(tl.float32)
        b_kg = tl.load(kg_ptr + off[:, None] * K + o_k[None, :])
        b_qg = tl.load(qg_ptr + off[:, None] * K + o_k[None, :])
        egn = tl.load(eg_ptr + (base + n) * K + o_k)
        ein = tl.load(ei_ptr + (base + n) * K + o_k)
        gl = tl.load(gl_ptr + (base + n) * K + o_k)

        # segment-local corrected values and output partials
        vp = (b_u - tl.dot(b_w, S.to(tl.bfloat16))).to(tl.bfloat16)
        kd = (b_kg.to(tl.float32) * ein[None, :]).to(tl.bfloat16)
        m_l = o_c[:, None] >= o_c[None, :]
        aq = tl.dot(b_qg, tl.trans(kd))
        aq = tl.where(m_l, aq, 0.0).to(tl.bfloat16)
        ol = tl.dot(b_qg, S.to(tl.bfloat16)) + tl.dot(aq, vp)
        tl.store(oloc_ptr + off[:, None] * V + o_v[None, :], ol.to(tl.bfloat16))

        sfac = tl.exp(lam_end - lam - gl)
        bacc += tl.dot(tl.trans((b_kg.to(tl.float32) * sfac[None, :]).to(tl.bfloat16)), vp)
        S = S * egn[:, None] + tl.dot(tl.trans(b_kg), vp)
        lam += gl
    tl.store(b_ptr + (i_bh * (NT // W) + i_s) * K * V
             + o_k[:, None] * V + o_v[None, :], bacc.to(tl.float16))


@triton.jit
def _outer_scan_kernel(
    a_ptr, b_ptr, cp_ptr,
    H: tl.constexpr, K: tl.constexpr, V: tl.constexpr,
    BV: tl.constexpr, NSEG: tl.constexpr,
):
    """Per (batch*head, V-tile): scan the segment transitions, leaving the
    state at the START of every segment in cp."""
    i_bh = tl.program_id(0)
    i_v = tl.program_id(1)
    o_k = tl.arange(0, K)
    o_v = i_v * BV + tl.arange(0, BV)

    S = tl.zeros((K, BV), dtype=tl.float32)
    for s in range(0, NSEG):
        tl.store(cp_ptr + (i_bh * NSEG + s) * K * V
                 + o_k[:, None] * V + o_v[None, :], S.to(tl.float16))
        a = tl.load(a_ptr + (i_bh * NSEG + s) * K * K
                    + o_k[:, None] * K + o_k[None, :])
        bb = tl.load(b_ptr + (i_bh * NSEG + s) * K * V
                     + o_k[:, None] * V + o_v[None, :]).to(tl.float32)
        S = tl.dot(a, S.to(tl.float16)) + bb


@triton.jit
def _corr_out_kernel(
    w_ptr, kg_ptr, qg_ptr, ei_ptr, dn_ptr, cp_ptr, oloc_ptr, o_ptr,
    T,
    H: tl.constexpr, K: tl.constexpr, V: tl.constexpr,
    BT: tl.constexpr, BV: tl.constexpr, W: tl.constexpr,
):
    """Per (chunk, batch*head, V-tile), fully parallel: add the checkpointed
    state's contribution to the local partials and store the final output."""
    i_n = tl.program_id(0)
    i_bh = tl.program_id(1)
    i_v = tl.program_id(2)
    i_b = i_bh // H
    i_h = i_bh % H
    NT = T // BT
    NSEG = NT // W

    o_c = tl.arange(0, BT)
    o_k = tl.arange(0, K)
    o_v = i_v * BV + tl.arange(0, BV)
    rows = i_n * BT + o_c
    off = (i_b * T + rows) * H + i_h

    b_w = tl.load(w_ptr + off[:, None] * K + o_k[None, :])
    b_qg = tl.load(qg_ptr + off[:, None] * K + o_k[None, :])
    b_kg = tl.load(kg_ptr + off[:, None] * K + o_k[None, :])
    ein = tl.load(ei_ptr + (i_bh * NT + i_n) * K + o_k)
    dn = tl.load(dn_ptr + (i_bh * NT + i_n) * K + o_k)

    kd = (b_kg.to(tl.float32) * ein[None, :]).to(tl.float16)
    cq = (b_qg.to(tl.float32) * dn[None, :]).to(tl.float16)
    cw = (b_w.to(tl.float32) * dn[None, :]).to(tl.float16)

    cp = tl.load(cp_ptr + (i_bh * NSEG + i_n // W) * K * V
                 + o_k[:, None] * V + o_v[None, :])
    ol = tl.load(oloc_ptr + off[:, None] * V + o_v[None, :]).to(tl.float32)

    m_l = o_c[:, None] >= o_c[None, :]
    aq = tl.dot(b_qg.to(tl.float16), tl.trans(kd))
    aq = tl.where(m_l, aq, 0.0).to(tl.float16)

    t_inter = tl.dot(cq, cp)
    t_vp = tl.dot(cw, cp).to(tl.float16)
    acc = ol + t_inter - tl.dot(aq, t_vp)

    tl.store(o_ptr + off[:, None] * V + o_v[None, :], acc.to(tl.bfloat16))


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__()
        assert chunk_size == 64, "this kernel is specialized for chunk_size=64"
        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._ws = None  # lazily-built workspace cache, keyed by (B,T,H,K,V,device)
        self.register_buffer("_dummy", torch.zeros(1), persistent=False)

    def _workspace(self, B, T, H, K, V, dev):
        key = (B, T, H, K, V, dev)
        ws = self._ws
        if ws is not None and ws[0] == key:
            return ws[1]
        NT = T // self.chunk_size
        W = self._pick_w(NT)
        NSEG = NT // W
        ws = (
            torch.empty(B * T * H * K, device=dev, dtype=torch.bfloat16),  # w
            torch.empty(B * T * H * V, device=dev, dtype=torch.bfloat16),  # u
            torch.empty(B * T * H * K, device=dev, dtype=torch.bfloat16),  # kg
            torch.empty(B * T * H * K, device=dev, dtype=torch.bfloat16),  # qg
            torch.empty(B * H * NT * K, device=dev, dtype=torch.float32),  # gl
            torch.empty(B * H * NT * K, device=dev, dtype=torch.float32),  # eg
            torch.empty(B * H * NT * K, device=dev, dtype=torch.float32),  # ei
            torch.empty(B * T * H * V, device=dev, dtype=torch.bfloat16),  # oloc
            torch.empty(B * H * NSEG * K * K, device=dev, dtype=torch.float16),  # A
            torch.empty(B * H * NSEG * K * V, device=dev, dtype=torch.float16),  # b
            torch.empty(B * H * NSEG * K * V, device=dev, dtype=torch.float16),  # cp
            torch.empty(B * H * NT * K, device=dev, dtype=torch.float32),  # dn
            torch.empty(B * H * NT * K, device=dev, dtype=torch.float32),  # sf
        )
        self._ws = (key, ws)
        return ws

    def _pick_w(self, NT):
        best, best_cost = 1, NT + 1
        for w in (2, 4, 8, 16):
            if NT % w == 0:
                cost = w + NT // w
                if cost < best_cost or (cost == best_cost and w < best):
                    best, best_cost = w, cost
        return best

    def forward(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        g: torch.Tensor,
        beta: torch.Tensor,
    ) -> torch.Tensor:
        B, T, H, K = q.shape
        V = v.shape[-1]
        BT = self.chunk_size
        NT = T // BT
        dev = q.device

        if not q.is_contiguous():
            q = q.contiguous()
        if not k.is_contiguous():
            k = k.contiguous()
        if not v.is_contiguous():
            v = v.contiguous()
        if not g.is_contiguous():
            g = g.contiguous()
        if not beta.is_contiguous():
            beta = beta.contiguous()

        w, u, kg, qg, gl, eg, ei, oloc, mat_a, vec_b, cp, dn, sf = \
            self._workspace(B, T, H, K, V, dev)
        o = torch.empty_like(v)
        W = self._pick_w(NT)
        NSEG = NT // W

        BC = 16
        _prep_wu_kernel[(NT, B * H)](
            q, k, v, g, beta, w, u, kg, qg, gl, eg, ei, self.scale, T,
            H=H, K=K, V=V, BT=BT, BC=BC, NSUB=BT // BC,
            num_warps=4, num_stages=1,
        )
        _seg_a_kernel[(NSEG, B * H)](
            w, kg, gl, mat_a, dn, sf, T,
            H=H, K=K, BT=BT, W=W,
            num_warps=4, num_stages=1,
        )
        BV = 16
        _seg_b_kernel[(NSEG, B * H, V // BV)](
            w, u, kg, qg, eg, ei, gl, oloc, vec_b, T,
            H=H, K=K, V=V, BT=BT, BV=BV, W=W,
            num_warps=4, num_stages=1,
        )
        _outer_scan_kernel[(B * H, V // BV)](
            mat_a, vec_b, cp,
            H=H, K=K, V=V, BV=BV, NSEG=NSEG,
            num_warps=4, num_stages=1,
        )
        _corr_out_kernel[(NT, B * H, V // BV)](
            w, kg, qg, ei, dn, cp, oloc, o, T,
            H=H, K=K, V=V, BT=BT, BV=BV, W=W,
            num_warps=4, num_stages=1,
        )
        return o

20260822_061013_or-fable_stealth_ox-alpha_02_kda_cutlass