KernelBench hard · B200
KDA CUTLASS Claude Fable 5
0.84%geomean peak fraction across shapes
manually audited: clean
harnessor-fableagent session2h 7mtotal wall2h 9mcheck2mbenchmark18soutput tokens—gpu-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