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 tokens—regimecompute
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