KernelBench hard · RTX PRO 6000
KDA CUTLASS Claude Fable 5
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.
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