KernelBench hard · H100
KDA CUTLASS GPT-5.6 Sol
1.32%geomean peak fraction across shapes
manually audited: clean
harnesscodexagent session27mtotal wall28mcheck42sbenchmark7soutput tokens54,714gpu-lock wait3sgpu-lock held7mregimecompute
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
2×1024×8×128×128×640.179 ms1.6%0.14 TB/s · 7% of 2.0 TB/s HBM · also 12 TFLOPS (2% of compute)
2×2048×8×128×128×640.318 ms1.8%0.16 TB/s · 8% of 2.0 TB/s HBM · also 13 TFLOPS (2% of compute)
1×4096×8×128×128×640.432 ms1.3%0.12 TB/s · 6% of 2.0 TB/s HBM · also 10 TFLOPS (1% of compute)
1×2048×4×128×128×640.176 ms0.8%0.07 TB/s · 4% of 2.0 TB/s HBM · also 6 TFLOPS (1% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(1.6% · 1.8% · 1.3% · 0.8%) = 1.3%
Kernel source (redacted)
"""SM90-oriented Triton forward kernels for Kimi Delta Attention.
Each 64-token chunk is prepared independently with tensor-core products and
a four-block triangular inverse. A final persistent kernel owns one
(batch, head, value-tile), fuses the inter-chunk recurrence with the output
projection, and keeps its 128 x BV state on chip for the complete sequence.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import triton
import triton.language as tl
PAIR_WARPS = 4
PAIR_STAGES = 3
SOLVE_WARPS = 2
SOLVE_STAGES = 1
PROJECT_WARPS = 4
PROJECT_STAGES = 3
STATE_BV = 16
STATE_WARPS = 4
OUTPUT_BV = 128
OUTPUT_WARPS = 4
OUTPUT_STAGES = 3
@triton.jit
def _delta_forward_kernel(
q_ptr,
k_ptr,
v_ptr,
g_ptr,
beta_ptr,
out_ptr,
T,
H: tl.constexpr,
D: tl.constexpr,
BV: tl.constexpr,
SCALE: tl.constexpr,
):
pid_v = tl.program_id(0)
pid_bh = tl.program_id(1)
b = pid_bh // H
h = pid_bh - b * H
ik = tl.arange(0, D)
iv = pid_v * BV + tl.arange(0, BV)
mv = iv < D
# This tensor is distributed over the participating warps and stays in
# registers for the dynamic sequence loop.
state = tl.zeros((D, BV), tl.float32)
t = 0
while t < T:
qk_base = ((b * T + t) * H + h) * D
v_base = qk_base
bh_base = (b * T + t) * H + h
q = tl.load(q_ptr + qk_base + ik).to(tl.float32)
k = tl.load(k_ptr + qk_base + ik).to(tl.float32)
decay = tl.exp(tl.load(g_ptr + qk_base + ik))
value = tl.load(v_ptr + v_base + iv, mask=mv, other=0.0).to(tl.float32)
beta = tl.load(beta_ptr + bh_base).to(tl.float32)
state *= decay[:, None]
prediction = tl.sum(state * k[:, None], axis=0)
residual = beta * (value - prediction)
state += k[:, None] * residual[None, :]
output = tl.sum(state * q[:, None], axis=0) * SCALE
tl.store(out_ptr + v_base + iv, output, mask=mv)
t += 1
@triton.jit
def _local_prefix_kernel(
src,
dst,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
BS: tl.constexpr,
):
pid_s = tl.program_id(0)
pid_t = tl.program_id(1)
pid_bh = tl.program_id(2)
b = pid_bh // H
h = pid_bh - b * H
base = (b * T * H + h) * D
p_src = tl.make_block_ptr(
src + base, (T, D), (H * D, 1),
(pid_t * BT, pid_s * BS), (BT, BS), (1, 0),
)
p_dst = tl.make_block_ptr(
dst + base, (T, D), (H * D, 1),
(pid_t * BT, pid_s * BS), (BT, BS), (1, 0),
)
x = tl.load(p_src)
tl.store(p_dst, tl.cumsum(x, axis=0))
@triton.jit
def _build_chunks_kernel(
q_ptr,
k_ptr,
v_ptr,
gc_ptr,
beta_ptr,
score_ptr,
w_ptr,
u_ptr,
kg_ptr,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
SCALE: tl.constexpr,
):
pid_t = tl.program_id(0)
pid_bh = tl.program_id(1)
b = pid_bh // H
h = pid_bh - b * H
tc = pid_t * BT
qk_base = (b * T * H + h) * D
s_base = (b * T * H + h) * BT
beta_base = b * T * H + h
p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_k = tl.make_block_ptr(k_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_v = tl.make_block_ptr(v_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_beta = tl.make_block_ptr(beta_ptr + beta_base, (T,), (H,), (tc,), (BT,), (0,))
q = tl.load(p_q)
k = tl.load(p_k)
g = tl.load(p_g)
beta = tl.load(p_beta).to(tl.float32)
# Centering the exponentials does not change either product and keeps both
# factors well-scaled even for an unusually lopsided chunk.
anchor = tl.sum(tl.where(tl.arange(0, BT)[:, None] == BT // 2, g, 0.0), axis=0)
ep = tl.exp(g - anchor[None, :])
em = tl.exp(anchor[None, :] - g)
qg = q.to(tl.float32) * ep
kg_left = k.to(tl.float32) * ep
kg_right = k.to(tl.float32) * em
scores = tl.dot(qg, tl.trans(kg_right), input_precision="tf32") * SCALE
interaction = tl.dot(kg_left, tl.trans(kg_right), input_precision="tf32")
ii = tl.arange(0, BT)
lower = ii[:, None] > ii[None, :]
causal = ii[:, None] >= ii[None, :]
interaction = tl.where(lower, interaction * beta[:, None], 0.0)
scores = tl.where(causal, scores, 0.0)
# Invert the unit lower-triangular system by forward substitution. Each
# completed row is immediately visible to subsequent rows in `inv`.
inv = -interaction
for row_idx in range(2, BT):
row = -tl.sum(
tl.where(ii[:, None] == row_idx, interaction, 0.0), axis=0
)
row += tl.sum(row[:, None] * inv, axis=0)
inv = tl.where((ii == row_idx)[:, None], row[None, :], inv)
inv += (ii[:, None] == ii[None, :])
transform = inv * beta[None, :]
p_score = tl.make_block_ptr(
score_ptr + s_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0)
)
tl.store(p_score, scores.to(tl.bfloat16))
# These three intermediates are deliberately bf16: their consumers are
# tensor-core products and the fp32 state holds the inter-chunk accuracy.
v = tl.load(p_v)
u = tl.dot(transform.to(tl.bfloat16), v)
p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
tl.store(p_u, u.to(tl.bfloat16))
k_exp = (k.to(tl.float32) * tl.exp(g)).to(tl.bfloat16)
w = tl.dot(transform.to(tl.bfloat16), k_exp)
p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
tl.store(p_w, w.to(tl.bfloat16))
g_last = tl.sum(tl.where(ii[:, None] == BT - 1, g, 0.0), axis=0)
k_tail = k.to(tl.float32) * tl.exp(g_last[None, :] - g)
p_kg = tl.make_block_ptr(kg_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
tl.store(p_kg, k_tail.to(tl.bfloat16))
@triton.jit
def _pair_matrices_kernel(
q_ptr,
k_ptr,
g_ptr,
gc_ptr,
beta_ptr,
score_ptr,
lower_ptr,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
SCALE: tl.constexpr,
):
pid_t = tl.program_id(0)
pid_bh = tl.program_id(1)
b = pid_bh // H
h = pid_bh - b * H
tc = pid_t * BT
qk_base = (b * T * H + h) * D
m_base = (b * T * H + h) * BT
beta_base = b * T * H + h
p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_k = tl.make_block_ptr(k_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_g = tl.make_block_ptr(g_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_gc = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_b = tl.make_block_ptr(beta_ptr + beta_base, (T,), (H,), (tc,), (BT,), (0,))
q = tl.load(p_q)
k = tl.load(p_k)
g = tl.cumsum(tl.load(p_g), axis=0)
tl.store(p_gc, g)
beta = tl.load(p_b).to(tl.float32)
rows = tl.arange(0, BT)
anchor = tl.sum(tl.where(rows[:, None] == BT // 2, g, 0.0), axis=0)
ep = tl.exp(g - anchor[None, :])
em = tl.exp(anchor[None, :] - g)
right = k.to(tl.float32) * em
score = tl.dot(q.to(tl.float32) * ep, tl.trans(right), input_precision="tf32") * SCALE
lower = tl.dot(k.to(tl.float32) * ep, tl.trans(right), input_precision="tf32")
score = tl.where(rows[:, None] >= rows[None, :], score, 0.0)
lower = tl.where(rows[:, None] > rows[None, :], lower * beta[:, None], 0.0)
p_score = tl.make_block_ptr(score_ptr + m_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0))
p_lower = tl.make_block_ptr(lower_ptr + m_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0))
tl.store(p_score, score.to(tl.bfloat16))
tl.store(p_lower, lower)
@triton.jit
def _load_matrix16(ptr, base, T, H: tl.constexpr, BT: tl.constexpr,
BC: tl.constexpr, tc, RO: tl.constexpr, CO: tl.constexpr):
p = tl.make_block_ptr(
ptr + base, (T, BT), (H * BT, 1),
(tc + RO * BC, CO * BC), (BC, BC), (1, 0),
)
return tl.load(p)
@triton.jit
def _load_beta16(ptr, base, T, H: tl.constexpr, BC: tl.constexpr,
tc, BLOCK: tl.constexpr):
p = tl.make_block_ptr(ptr + base, (T,), (H,), (tc + BLOCK * BC,), (BC,), (0,))
return tl.load(p).to(tl.float32)
@triton.jit
def _store_matrix16(ptr, x, base, T, H: tl.constexpr, BT: tl.constexpr,
BC: tl.constexpr, tc, RO: tl.constexpr, CO: tl.constexpr):
p = tl.make_block_ptr(
ptr + base, (T, BT), (H * BT, 1),
(tc + RO * BC, CO * BC), (BC, BC), (1, 0),
)
tl.store(p, x.to(tl.bfloat16))
@triton.jit
def _solve64_kernel(
lower_ptr,
beta_ptr,
transform_ptr,
T,
H: tl.constexpr,
BT: tl.constexpr,
BC: tl.constexpr,
):
pid_t = tl.program_id(0)
pid_bh = tl.program_id(1)
b = pid_bh // H
h = pid_bh - b * H
tc = pid_t * BT
m_base = (b * T * H + h) * BT
beta_base = b * T * H + h
l00 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 0, 0)
l10 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 1, 0)
l11 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 1, 1)
l20 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 2, 0)
l21 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 2, 1)
l22 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 2, 2)
l30 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 0)
l31 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 1)
l32 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 2)
l33 = _load_matrix16(lower_ptr, m_base, T, H, BT, BC, tc, 3, 3)
ii = tl.arange(0, BC)
strict = ii[:, None] > ii[None, :]
ident = ii[:, None] == ii[None, :]
a00 = -tl.where(strict, l00, 0.0)
a11 = -tl.where(strict, l11, 0.0)
a22 = -tl.where(strict, l22, 0.0)
a33 = -tl.where(strict, l33, 0.0)
for row_idx in range(2, BC):
r0 = -tl.sum(tl.where(ii[:, None] == row_idx, l00, 0.0), axis=0)
r1 = -tl.sum(tl.where(ii[:, None] == row_idx, l11, 0.0), axis=0)
r2 = -tl.sum(tl.where(ii[:, None] == row_idx, l22, 0.0), axis=0)
r3 = -tl.sum(tl.where(ii[:, None] == row_idx, l33, 0.0), axis=0)
r0 += tl.sum(r0[:, None] * a00, axis=0)
r1 += tl.sum(r1[:, None] * a11, axis=0)
r2 += tl.sum(r2[:, None] * a22, axis=0)
r3 += tl.sum(r3[:, None] * a33, axis=0)
select = (ii == row_idx)[:, None]
a00 = tl.where(select, r0[None, :], a00)
a11 = tl.where(select, r1[None, :], a11)
a22 = tl.where(select, r2[None, :], a22)
a33 = tl.where(select, r3[None, :], a33)
a00 += ident
a11 += ident
a22 += ident
a33 += ident
a10 = -tl.dot(tl.dot(a11, l10, input_precision="tf32"), a00, input_precision="tf32")
a21 = -tl.dot(tl.dot(a22, l21, input_precision="tf32"), a11, input_precision="tf32")
a20 = -tl.dot(
a22,
tl.dot(l20, a00, input_precision="tf32") + tl.dot(l21, a10, input_precision="tf32"),
input_precision="tf32",
)
a32 = -tl.dot(tl.dot(a33, l32, input_precision="tf32"), a22, input_precision="tf32")
a31 = -tl.dot(
a33,
tl.dot(l31, a11, input_precision="tf32") + tl.dot(l32, a21, input_precision="tf32"),
input_precision="tf32",
)
a30 = -tl.dot(
a33,
tl.dot(l30, a00, input_precision="tf32")
+ tl.dot(l31, a10, input_precision="tf32")
+ tl.dot(l32, a20, input_precision="tf32"),
input_precision="tf32",
)
b0 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 0)
b1 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 1)
b2 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 2)
b3 = _load_beta16(beta_ptr, beta_base, T, H, BC, tc, 3)
_store_matrix16(transform_ptr, a00 * b0[None, :], m_base, T, H, BT, BC, tc, 0, 0)
_store_matrix16(transform_ptr, a10 * b0[None, :], m_base, T, H, BT, BC, tc, 1, 0)
_store_matrix16(transform_ptr, a11 * b1[None, :], m_base, T, H, BT, BC, tc, 1, 1)
_store_matrix16(transform_ptr, a20 * b0[None, :], m_base, T, H, BT, BC, tc, 2, 0)
_store_matrix16(transform_ptr, a21 * b1[None, :], m_base, T, H, BT, BC, tc, 2, 1)
_store_matrix16(transform_ptr, a22 * b2[None, :], m_base, T, H, BT, BC, tc, 2, 2)
_store_matrix16(transform_ptr, a30 * b0[None, :], m_base, T, H, BT, BC, tc, 3, 0)
_store_matrix16(transform_ptr, a31 * b1[None, :], m_base, T, H, BT, BC, tc, 3, 1)
_store_matrix16(transform_ptr, a32 * b2[None, :], m_base, T, H, BT, BC, tc, 3, 2)
_store_matrix16(transform_ptr, a33 * b3[None, :], m_base, T, H, BT, BC, tc, 3, 3)
@triton.jit
def _project_chunks_kernel(
k_ptr,
v_ptr,
gc_ptr,
transform_ptr,
w_ptr,
u_ptr,
kg_ptr,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
):
pid_t = tl.program_id(0)
pid_bh = tl.program_id(1)
b = pid_bh // H
h = pid_bh - b * H
tc = pid_t * BT
qk_base = (b * T * H + h) * D
m_base = (b * T * H + h) * BT
p_a = tl.make_block_ptr(transform_ptr + m_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0))
a = tl.load(p_a)
ii = tl.arange(0, BT)
a = tl.where(ii[:, None] >= ii[None, :], a, 0.0)
for block in range(2):
d0 = block * 64
p_k = tl.make_block_ptr(k_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
p_v = tl.make_block_ptr(v_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
k = tl.load(p_k)
v = tl.load(p_v)
g = tl.load(p_g)
u = tl.dot(a, v)
w = tl.dot(a, (k.to(tl.float32) * tl.exp(g)).to(tl.bfloat16))
last = tl.sum(tl.where(ii[:, None] == BT - 1, g, 0.0), axis=0)
kg = k.to(tl.float32) * tl.exp(last[None, :] - g)
p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
p_kg = tl.make_block_ptr(kg_ptr + qk_base, (T, D), (H * D, 1), (tc, d0), (BT, 64), (1, 0))
tl.store(p_u, u.to(tl.bfloat16))
tl.store(p_w, w.to(tl.bfloat16))
tl.store(p_kg, kg.to(tl.bfloat16))
@triton.jit
def _chunks_recurrence_kernel(
q_ptr,
gc_ptr,
score_ptr,
w_ptr,
u_ptr,
kg_ptr,
out_ptr,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
BV: tl.constexpr,
SCALE: tl.constexpr,
):
pid_v = tl.program_id(0)
pid_bh = tl.program_id(1)
b = pid_bh // H
h = pid_bh - b * H
v0 = pid_v * BV
qk_base = (b * T * H + h) * D
s_base = (b * T * H + h) * BT
state0 = tl.zeros((64, BV), tl.float32)
state1 = tl.zeros((64, BV), tl.float32)
chunk = 0
while chunk < T // BT:
tc = chunk * BT
p_w0 = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, 64), (1, 0))
p_w1 = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 64), (BT, 64), (1, 0))
w0 = tl.load(p_w0)
w1 = tl.load(p_w1)
correction = tl.dot(w0, state0.to(tl.bfloat16))
correction += tl.dot(w1, state1.to(tl.bfloat16))
p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
value = tl.load(p_u).to(tl.float32) - correction
value_bf = value.to(tl.bfloat16)
p_q0 = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, 64), (1, 0))
p_q1 = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 64), (BT, 64), (1, 0))
p_g0 = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, 64), (1, 0))
p_g1 = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 64), (BT, 64), (1, 0))
g0 = tl.load(p_g0)
g1 = tl.load(p_g1)
q0 = (tl.load(p_q0).to(tl.float32) * tl.exp(g0)).to(tl.bfloat16)
q1 = (tl.load(p_q1).to(tl.float32) * tl.exp(g1)).to(tl.bfloat16)
output = tl.dot(q0, state0.to(tl.bfloat16))
output += tl.dot(q1, state1.to(tl.bfloat16))
output *= SCALE
p_score = tl.make_block_ptr(
score_ptr + s_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0)
)
score = tl.load(p_score)
output += tl.dot(score, value_bf)
p_out = tl.make_block_ptr(out_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
tl.store(p_out, output.to(tl.bfloat16))
ik = tl.arange(0, 64)
last0 = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + ik)
last1 = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + 64 + ik)
state0 *= tl.exp(last0)[:, None]
state1 *= tl.exp(last1)[:, None]
p_kg0 = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (0, tc), (64, BT), (0, 1))
p_kg1 = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (64, tc), (64, BT), (0, 1))
state0 += tl.dot(tl.load(p_kg0), value_bf)
state1 += tl.dot(tl.load(p_kg1), value_bf)
chunk += 1
@triton.jit
def _chunks_recurrence128_kernel(
q_ptr,
gc_ptr,
score_ptr,
w_ptr,
u_ptr,
kg_ptr,
out_ptr,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
BV: tl.constexpr,
SCALE: tl.constexpr,
):
pid_v = tl.program_id(0)
pid_bh = tl.program_id(1)
b = pid_bh // H
h = pid_bh - b * H
v0 = pid_v * BV
qk_base = (b * T * H + h) * D
s_base = (b * T * H + h) * BT
state = tl.zeros((D, BV), tl.float32)
chunk = 0
while chunk < T // BT:
tc = chunk * BT
p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
value = tl.load(p_u).to(tl.float32) - tl.dot(tl.load(p_w), state.to(tl.bfloat16))
value_bf = value.to(tl.bfloat16)
p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
g = tl.load(p_g)
qg = (tl.load(p_q).to(tl.float32) * tl.exp(g)).to(tl.bfloat16)
output = tl.dot(qg, state.to(tl.bfloat16)) * SCALE
p_score = tl.make_block_ptr(score_ptr + s_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0))
output += tl.dot(tl.load(p_score), value_bf)
p_out = tl.make_block_ptr(out_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
tl.store(p_out, output.to(tl.bfloat16))
ik = tl.arange(0, D)
last = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + ik)
state *= tl.exp(last)[:, None]
p_kg = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (0, tc), (D, BT), (0, 1))
state += tl.dot(tl.load(p_kg), value_bf)
chunk += 1
@triton.jit
def _state_chunks_kernel(
gc_ptr,
w_ptr,
u_ptr,
kg_ptr,
states_ptr,
values_ptr,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
BV: tl.constexpr,
):
pid_v = tl.program_id(0)
pid_bh = tl.program_id(1)
b = pid_bh // H
h = pid_bh - b * H
v0 = pid_v * BV
nt = T // BT
qk_base = (b * T * H + h) * D
state_base = (b * nt * H + h) * D * D
state = tl.zeros((D, BV), tl.float32)
chunk = 0
while chunk < nt:
tc = chunk * BT
p_state = tl.make_block_ptr(
states_ptr + state_base + chunk * H * D * D,
(D, D), (D, 1), (0, v0), (D, BV), (1, 0),
)
tl.store(p_state, state.to(tl.bfloat16))
p_w = tl.make_block_ptr(w_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_u = tl.make_block_ptr(u_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
value = tl.load(p_u).to(tl.float32) - tl.dot(tl.load(p_w), state.to(tl.bfloat16))
value_bf = value.to(tl.bfloat16)
p_value = tl.make_block_ptr(values_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
tl.store(p_value, value_bf)
ik = tl.arange(0, D)
last = tl.load(gc_ptr + qk_base + (tc + BT - 1) * H * D + ik)
state *= tl.exp(last)[:, None]
p_kg = tl.make_block_ptr(kg_ptr + qk_base, (D, T), (1, H * D), (0, tc), (D, BT), (0, 1))
state += tl.dot(tl.load(p_kg), value_bf)
chunk += 1
@triton.jit
def _output_chunks_kernel(
q_ptr,
gc_ptr,
score_ptr,
states_ptr,
values_ptr,
out_ptr,
T,
H: tl.constexpr,
D: tl.constexpr,
BT: tl.constexpr,
BV: tl.constexpr,
SCALE: tl.constexpr,
):
pid_v = tl.program_id(0)
chunk = tl.program_id(1)
pid_bh = tl.program_id(2)
b = pid_bh // H
h = pid_bh - b * H
v0 = pid_v * BV
nt = T // BT
tc = chunk * BT
qk_base = (b * T * H + h) * D
score_base = (b * T * H + h) * BT
state_base = ((b * nt + chunk) * H + h) * D * D
p_q = tl.make_block_ptr(q_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_g = tl.make_block_ptr(gc_ptr + qk_base, (T, D), (H * D, 1), (tc, 0), (BT, D), (1, 0))
p_state = tl.make_block_ptr(states_ptr + state_base, (D, D), (D, 1), (0, v0), (D, BV), (1, 0))
qg = (tl.load(p_q).to(tl.float32) * tl.exp(tl.load(p_g))).to(tl.bfloat16)
output = tl.dot(qg, tl.load(p_state)) * SCALE
p_score = tl.make_block_ptr(score_ptr + score_base, (T, BT), (H * BT, 1), (tc, 0), (BT, BT), (1, 0))
p_value = tl.make_block_ptr(values_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
output += tl.dot(tl.load(p_score), tl.load(p_value))
p_out = tl.make_block_ptr(out_ptr + qk_base, (T, D), (H * D, 1), (tc, v0), (BT, BV), (1, 0))
tl.store(p_out, output.to(tl.bfloat16))
def _forward(q, k, v, g, beta, scale: float):
B, T, H, D = q.shape
bt = 64
nt = T // bt
gc = torch.empty_like(g)
score = torch.empty((B, T, H, bt), device=q.device, dtype=q.dtype)
w = torch.empty_like(k)
u = torch.empty_like(v)
kg = torch.empty_like(k)
out = torch.empty_like(v)
lower = torch.empty((B, T, H, bt), device=q.device, dtype=torch.float32)
transform = torch.empty_like(score)
_pair_matrices_kernel[(nt, B * H)](
q, k, g, gc, beta, score, lower,
T, H=H, D=D, BT=bt, SCALE=scale,
num_warps=PAIR_WARPS,
num_stages=PAIR_STAGES,
)
_solve64_kernel[(nt, B * H)](
lower, beta, transform, T, H=H, BT=bt, BC=16,
num_warps=SOLVE_WARPS,
num_stages=SOLVE_STAGES,
)
_project_chunks_kernel[(nt, B * H)](
k, v, gc, transform, w, u, kg,
T, H=H, D=D, BT=bt,
num_warps=PROJECT_WARPS,
num_stages=PROJECT_STAGES,
)
states = torch.empty((B, nt, H, D, D), device=q.device, dtype=q.dtype)
values = torch.empty_like(v)
long_context = T >= 4096
state_bv = 32 if long_context else STATE_BV
state_warps = 8 if long_context else STATE_WARPS
state_stages = 4 if long_context else 5
_state_chunks_kernel[(D // state_bv, B * H)](
gc, w, u, kg, states, values,
T, H=H, D=D, BT=bt, BV=state_bv,
num_warps=state_warps,
num_stages=state_stages,
)
output_bv = OUTPUT_BV
_output_chunks_kernel[(D // output_bv, nt, B * H)](
q, gc, score, states, values, out,
T, H=H, D=D, BT=bt, BV=output_bv, SCALE=scale,
num_warps=OUTPUT_WARPS,
num_stages=OUTPUT_STAGES,
)
return out
class Model(nn.Module):
def __init__(self, B: int, T: int, H: int, K: int, V: int, chunk_size: int = 64):
super().__init__()
if K != 128 or V != 128 or chunk_size != 64:
raise ValueError("This kernel is specialized for K=V=128 and chunks of 64")
self.scale = float(K) ** -0.5
self.register_buffer("_dummy", torch.zeros(1), persistent=False)
def forward(self, q, k, v, g, beta):
return _forward(q, k, v, g, beta, self.scale)
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]
20260721_142442_codex_gpt-5.6-sol_02_kda_cutlass