KernelBench hard · RTX PRO 6000
KDA CUTLASS GLM-5.3
4.56%geomean peak fraction across shapes
manually audited: clean
Isolated-regrade 0.0456 on RTX PRO 6000. Custom KDA chunk forward: CUDA K1 (SM120 mma.sync via load_inline) plus a Triton two-level inter-chunk scan (k2a/k2b/k2c). Workspace is shape-keyed (B,T,H) scratch, rewritten every call; output is a fresh empty tensor. No FLA import. Lint CLEAN. template_mutated=false. Numeric stress on; check.py unmodified.
harnesszai-claudeagent session4h 48mtotal wall4h 59mcheck37sbenchmark2soutput tokens583,747cost$39.21regimecompute
Per-shape vs governing ceilingeach shape graded against whichever binds — bf16 compute or HBM bandwidth
2×1024×8×128×128×640.090 ms4.8%0.28 TB/s · 16% of 1.8 TB/s HBM · also 24 TFLOPS (5% of compute)
2×2048×8×128×128×640.146 ms5.9%0.35 TB/s · 19% of 1.8 TB/s HBM · also 29 TFLOPS (6% of compute)
1×4096×8×128×128×640.155 ms5.5%0.33 TB/s · 18% of 1.8 TB/s HBM · also 28 TFLOPS (6% of compute)
1×2048×4×128×128×640.078 ms2.8%0.16 TB/s · 9% of 1.8 TB/s HBM · also 14 TFLOPS (3% of compute)
compute-bound memory-bound · bar + right column = official fraction of the ceiling (the geomean input)
geomean(4.8% · 5.9% · 5.5% · 2.8%) = 4.6%
Kernel source (redacted)
"""Kimi Delta Attention (KDA) forward, chunk form — custom Triton kernels for SM120.
Two-kernel decomposition. Per chunk (BT=64, K=V=128; T % 64 == 0 so no masks):
K1 (grid NT x B*H, chunk-parallel)
g -> in-chunk fp32 cumsum
A0 = (k e^{g-g_l}) (k e^{g_l-g})^T (bf16 tensor-core dot)
B = tril(beta_row * A0, -1); M = (I+B)^-1 via exact nilpotent doubling
u = M (beta v), w = M (beta e^g k)
P' = diag(e^{g_l}) - (k e^{g_l-g})^T w [the chunk's state operator]
C = (k e^{g_l-g})^T u
Aqk = scale * tril(q e^{g-g_l} (k e^{g_l-g})^T, incl diag)
qe = scale q e^g - Aqk w, t = Aqk u
stores: P', C (B,H,NT,128,128 bf16), qe, t (token-major bf16)
K2 (grid V/BV x B*H, sequential over NT, fused output)
S = C_i + P'_i @ S (one dot per step; decay folded into P')
o_i = qe_i @ S_i + t_i (uses the pre-update state)
o = qg @ S_i + Aqk (u_i - w_i S_i) = (qg - Aqk w) S_i + Aqk u (exact identity).
"""
from __future__ import annotations
import torch
import torch.nn as nn
import os
import triton
import triton.language as tl
RCP_LN2 = tl.constexpr(1.4426950408889634)
@triton.jit
def _col2(x, AX: tl.constexpr):
"""[AX, 128] -> two [AX, 64] column halves."""
return tl.split(tl.permute(tl.reshape(x, (AX, 2, 64)), (0, 2, 1)))
@triton.jit
def _row2(x, AY: tl.constexpr):
"""[64, AY] -> two [32, AY] row halves."""
return tl.split(tl.permute(tl.reshape(x, (2, 32, AY)), (1, 2, 0)))
@triton.jit
def _col2b(x, AX: tl.constexpr):
"""[AX, 64] -> two [AX, 32] column halves."""
return tl.split(tl.permute(tl.reshape(x, (AX, 2, 32)), (0, 2, 1)))
# ---------------------------------------------------------------------------
# K1: chunk-local preparation (full-tile form, no layout shuffles)
# ---------------------------------------------------------------------------
@triton.jit
def _kda_k1(
q, k, v, g, beta,
P, C, Pd, qe, t,
T,
H: tl.constexpr,
NT,
scale,
):
pid_t = tl.program_id(0)
pid_bh = tl.program_id(1).to(tl.int64)
i_b = pid_bh // H
i_h = pid_bh - i_b * H
o_c = tl.arange(0, 64)
o_k = tl.arange(0, 128)
tok0 = pid_t * 64
rowb = (i_b * T * H + i_h)
b_k = tl.load(k + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
b_g = tl.load(g + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
b_beta = tl.load(beta + rowb + (tok0 + o_c) * H).to(tl.float32)
# cumsum via triangular-ones dot (tf32) instead of 6 shuffle rounds
ltri = tl.where(o_c[:, None] >= o_c[None, :], 1.0, 0.0)
b_g = tl.dot(ltri, b_g * RCP_LN2, input_precision="tf32")
b_glast = tl.sum(tl.where((o_c[:, None] == 63), b_g, 0.0), 0)
b_d = tl.exp2(b_glast)
e_up = tl.exp2(b_g - b_glast[None, :])
e_all = e_up * b_d[None, :]
b_kf = b_k.to(tl.float32)
b_kl = (b_kf * e_up).to(tl.bfloat16)
b_kr = (b_kf * tl.exp2(b_glast[None, :] - b_g)).to(tl.bfloat16)
# ---- A0, B; M = (I+B)^-1 by bf16 doubling ----
b_B = tl.dot(b_kl, tl.trans(b_kr))
b_B = tl.where(o_c[:, None] > o_c[None, :], b_B * b_beta[:, None], 0.0).to(tl.bfloat16)
eye = tl.where(o_c[:, None] == o_c[None, :], 1.0, 0.0).to(tl.bfloat16)
X = (eye.to(tl.float32) - b_B.to(tl.float32)).to(tl.bfloat16)
b_B = tl.dot(b_B, b_B).to(tl.bfloat16) # B^2
X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16)
b_B = tl.dot(b_B, b_B).to(tl.bfloat16) # B^4
X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16)
b_B = tl.dot(b_B, b_B).to(tl.bfloat16) # B^8
X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16)
b_B = tl.dot(b_B, b_B).to(tl.bfloat16) # B^16
X = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16)
b_B = tl.dot(b_B, b_B).to(tl.bfloat16) # B^32
M = tl.dot(X, (eye.to(tl.float32) + b_B.to(tl.float32)).to(tl.bfloat16)).to(tl.bfloat16)
# ---- w = M kb, u = M vb ----
b_v = tl.load(v + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
b_kb = (b_kf * (b_beta[:, None] * e_all)).to(tl.bfloat16)
b_vb = (b_v.to(tl.float32) * b_beta[:, None]).to(tl.bfloat16)
b_w = tl.dot(M, b_kb).to(tl.bfloat16)
b_u = tl.dot(M, b_vb).to(tl.bfloat16)
# ---- P'' = -kr^T w, C = kr^T u ----
pbase = (pid_bh * NT + pid_t) * (128 * 128)
b_krt = tl.trans(b_kr)
tl.store(P + pbase + o_k[:, None] * 128 + o_k[None, :], (-tl.dot(b_krt, b_w)).to(tl.bfloat16))
tl.store(C + pbase + o_k[:, None] * 128 + o_k[None, :], tl.dot(b_krt, b_u).to(tl.bfloat16))
tl.store(Pd + (pid_bh * NT + pid_t) * 128 + o_k, b_d)
# ---- Aqk, qe, t ----
b_q = tl.load(q + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
b_ql = (b_q.to(tl.float32) * (scale * e_up)).to(tl.bfloat16)
b_qg = (b_q.to(tl.float32) * (scale * e_all)).to(tl.bfloat16)
b_Aqk = tl.dot(b_ql, tl.trans(b_kr))
b_Aqk = tl.where(o_c[:, None] >= o_c[None, :], b_Aqk, 0.0)
b_Aqk = b_Aqk.to(tl.bfloat16)
b_qe = (b_qg.to(tl.float32) - tl.dot(b_Aqk, b_w)).to(tl.bfloat16)
b_t = tl.dot(b_Aqk, b_u).to(tl.bfloat16)
tl.store(qe + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :], b_qe)
tl.store(t + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :], b_t)
# ---------------------------------------------------------------------------
# K2 pass A: per-group local scan (S=0 init) + group composite A_g
# grid (NV, BH, G), BV columns each; A maintained full-width redundantly
# ---------------------------------------------------------------------------
@triton.jit
def _kda_k2a(
P, C, Pd, A, Bg,
T,
H: tl.constexpr,
NT,
L: tl.constexpr,
BV: tl.constexpr,
):
i_v = tl.program_id(0)
pid_bh = tl.program_id(1).to(tl.int64)
i_g = tl.program_id(2)
i_b = pid_bh // H
i_h = pid_bh - i_b * H
o_k = tl.arange(0, 128)
o_v = i_v * BV + tl.arange(0, BV)
rowb = (i_b * T * H + i_h)
b_S = tl.zeros([128, BV], dtype=tl.float32)
b_A = tl.where(o_k[:, None] == o_k[None, :], 1.0, 0.0) # [128,128] fp32
for j in range(L):
i_t = i_g * L + j
pbase = (pid_bh * NT + i_t) * (128 * 128)
b_P = tl.load(P + pbase + o_k[:, None] * 128 + o_k[None, :])
b_C = tl.load(C + pbase + o_k[:, None] * 128 + o_v[None, :])
b_d = tl.load(Pd + (pid_bh * NT + i_t) * 128 + o_k)
b_S = b_d[:, None] * b_S + b_C.to(tl.float32) + tl.dot(b_P, b_S.to(tl.bfloat16))
b_A = b_d[:, None] * b_A + tl.dot(b_P, b_A.to(tl.bfloat16))
abase = ((pid_bh * tl.num_programs(2)) + i_g) * (128 * 128)
tl.store(A + abase + o_k[:, None] * 128 + o_k[None, :], b_A.to(tl.bfloat16))
tl.store(Bg + abase + o_k[:, None] * 128 + o_v[None, :], b_S.to(tl.bfloat16))
# ---------------------------------------------------------------------------
# K2 pass B: tiny sequential scan over the G group composites -> boundary states
# grid (NV, BH)
# ---------------------------------------------------------------------------
@triton.jit
def _kda_k2b(
A, Bg, Sb,
H: tl.constexpr,
G: tl.constexpr,
BV: tl.constexpr,
):
i_v = tl.program_id(0)
pid_bh = tl.program_id(1).to(tl.int64)
o_k = tl.arange(0, 128)
o_v = i_v * BV + tl.arange(0, BV)
b_S = tl.zeros([128, BV], dtype=tl.float32)
for i_g in range(G):
abase = ((pid_bh * G) + i_g) * (128 * 128)
b_A = tl.load(A + abase + o_k[:, None] * 128 + o_k[None, :])
b_Bg = tl.load(Bg + abase + o_k[:, None] * 128 + o_v[None, :])
tl.store(Sb + abase + o_k[:, None] * 128 + o_v[None, :], b_S.to(tl.bfloat16))
b_S = b_Bg.to(tl.float32) + tl.dot(b_A, b_S.to(tl.bfloat16))
# ---------------------------------------------------------------------------
# K2 pass C: final scan within each group from its boundary state + outputs
# grid (NV, BH, G)
# ---------------------------------------------------------------------------
@triton.jit
def _kda_k2c(
P, C, Pd, qe, t, Sb,
o,
T,
H: tl.constexpr,
NT,
L: tl.constexpr,
BV: tl.constexpr,
):
i_v = tl.program_id(0)
pid_bh = tl.program_id(1).to(tl.int64)
i_g = tl.program_id(2)
i_b = pid_bh // H
i_h = pid_bh - i_b * H
o_c = tl.arange(0, 64)
o_k = tl.arange(0, 128)
o_v = i_v * BV + tl.arange(0, BV)
rowb = (i_b * T * H + i_h)
abase = ((pid_bh * tl.num_programs(2)) + i_g) * (128 * 128)
b_S = tl.load(Sb + abase + o_k[:, None] * 128 + o_v[None, :]).to(tl.float32)
for j in range(L):
i_t = i_g * L + j
tok0 = i_t * 64
pbase = (pid_bh * NT + i_t) * (128 * 128)
b_P = tl.load(P + pbase + o_k[:, None] * 128 + o_k[None, :])
b_C = tl.load(C + pbase + o_k[:, None] * 128 + o_v[None, :])
b_d = tl.load(Pd + (pid_bh * NT + i_t) * 128 + o_k)
b_qe = tl.load(qe + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
b_t = tl.load(t + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :])
b_o = tl.dot(b_qe, b_S.to(tl.bfloat16)) + b_t.to(tl.float32)
tl.store(o + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :],
b_o.to(tl.bfloat16))
b_S = b_d[:, None] * b_S + b_C.to(tl.float32) + tl.dot(b_P, b_S.to(tl.bfloat16))
# ---------------------------------------------------------------------------
# K2 (fallback): single sequential scan with fused output
# ---------------------------------------------------------------------------
@triton.jit
def _kda_k2(
P, C, Pd, qe, t,
o,
T,
H: tl.constexpr,
NT,
BV: tl.constexpr,
):
i_v = tl.program_id(0)
pid_bh = tl.program_id(1).to(tl.int64)
i_b = pid_bh // H
i_h = pid_bh - i_b * H
o_c = tl.arange(0, 64)
o_k = tl.arange(0, 128)
o_v = i_v * BV + tl.arange(0, BV)
rowb = (i_b * T * H + i_h)
b_S = tl.zeros([128, BV], dtype=tl.float32)
for i_t in range(NT):
tok0 = i_t * 64
pbase = (pid_bh * NT + i_t) * (128 * 128)
b_P = tl.load(P + pbase + o_k[:, None] * 128 + o_k[None, :]) # [128,128]
b_C = tl.load(C + pbase + o_k[:, None] * 128 + o_v[None, :]) # [128,BV]
b_d = tl.load(Pd + (pid_bh * NT + i_t) * 128 + o_k) # [128] fp32
b_qe = tl.load(qe + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_k[None, :])
b_t = tl.load(t + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :])
b_o = tl.dot(b_qe, b_S.to(tl.bfloat16)) + b_t.to(tl.float32) # pre-update S
tl.store(o + rowb * 128 + (tok0 + o_c)[:, None] * (H * 128) + o_v[None, :],
b_o.to(tl.bfloat16))
b_S = b_d[:, None] * b_S + b_C.to(tl.float32) + tl.dot(b_P, b_S.to(tl.bfloat16))
# ---------------------------------------------------------------------------
# Model
# ---------------------------------------------------------------------------
def _kda_forward(q, k, v, g, beta, scale, chunk_size=64):
B, T, H, K = q.shape
V = v.shape[-1]
assert K == 128 and V == 128 and chunk_size == 64 and T % 64 == 0
NT = T // 64
BH = B * H
dev = q.device
bufs = _BUF_CACHE.get((B, T, H))
if bufs is None:
# two-level scan split: aim for ~128 CTAs in k2a/k2c (2*BV-halves per bh)
tgt = max(1, 128 // (2 * BH))
G = 1
while G * 2 <= tgt and NT % (G * 2) == 0:
G *= 2
bufs = (
torch.empty(B, H, NT, K, V, dtype=torch.bfloat16, device=dev), # P
torch.empty(B, H, NT, K, V, dtype=torch.bfloat16, device=dev), # C
torch.empty(B * H * NT, K, dtype=torch.float32, device=dev), # Pd
torch.empty(B, T, H, K, dtype=torch.bfloat16, device=dev), # qe
torch.empty(B, T, H, V, dtype=torch.bfloat16, device=dev), # t
torch.empty(B * H * G, K, V, dtype=torch.bfloat16, device=dev), # A
torch.empty(B * H * G, K, V, dtype=torch.bfloat16, device=dev), # Bg
torch.empty(B * H * G, K, V, dtype=torch.bfloat16, device=dev), # Sb
G,
)
_BUF_CACHE[(B, T, H)] = bufs
P_, C_, Pd_, qe_, t_, A_, Bg_, Sb_, G = bufs
L = NT // G
o = torch.empty(B, T, H, V, dtype=torch.bfloat16, device=dev)
_run_k1(q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale, BH)
_kda_k2a[(2, BH, G)](P_, C_, Pd_, A_, Bg_, T, H, NT, L=L, BV=64,
num_warps=8, num_stages=2)
_kda_k2b[(2, BH)](A_, Bg_, Sb_, H, G=G, BV=64, num_warps=8, num_stages=2)
_kda_k2c[(2, BH, G)](P_, C_, Pd_, qe_, t_, Sb_, o, T, H, NT, L=L, BV=64,
num_warps=8, num_stages=2)
return o
# ---------------- CUDA K1 (SM120 mma.sync path) ----------------
_CUDA_K1_SRC = r"""
#include <torch/extension.h>
#include <cuda_bf16.h>
#define DEV __device__ __forceinline__
DEV int lane_id() { return threadIdx.x & 31; }
DEV uint32_t smem_u32(const void* p) { return (uint32_t)__cvta_generic_to_shared(p); }
DEV void ldm_x4(uint32_t& r0, uint32_t& r1, uint32_t& r2, uint32_t& r3, const void* p) {
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];\n"
: "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3) : "r"(smem_u32(p)));
}
DEV void ldm_x4t(uint32_t& r0, uint32_t& r1, uint32_t& r2, uint32_t& r3, const void* p) {
asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 {%0,%1,%2,%3}, [%4];\n"
: "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3) : "r"(smem_u32(p)));
}
DEV void stm_x4(const void* p, uint32_t r0, uint32_t r1, uint32_t r2, uint32_t r3) {
asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1,%2,%3,%4};\n"
:: "r"(smem_u32(p)), "r"(r0), "r"(r1), "r"(r2), "r"(r3));
}
DEV void mma_bf16(float& d0, float& d1, float& d2, float& d3,
uint32_t a0, uint32_t a1, uint32_t a2, uint32_t a3, uint32_t b0, uint32_t b1) {
asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%0,%1,%2,%3};\n"
: "+f"(d0), "+f"(d1), "+f"(d2), "+f"(d3)
: "r"(a0), "r"(a1), "r"(a2), "r"(a3), "r"(b0), "r"(b1));
}
DEV uint32_t pack_bf16(float lo, float hi) {
__nv_bfloat162 v = __floats2bfloat162_rn(lo, hi);
return *reinterpret_cast<uint32_t*>(&v);
}
DEV float bf2f(__nv_bfloat16 x) { return __bfloat162float(x); }
#ifdef K1_DBG
__device__ bool k1_dbg_root() { return blockIdx.x == 0 && blockIdx.y == 0; }
DEV void dump_f32(float* dbg, long base, const float* buf, int n) {
if (blockIdx.x != 0 || blockIdx.y != 0) return;
for (int i = threadIdx.x; i < n; i += 256) dbg[base + i] = buf[i];
}
DEV void dump_bf(float* dbg, long base, const __nv_bfloat16* buf, int n) {
if (!k1_dbg_root()) return;
for (int i = threadIdx.x; i < n; i += 256) dbg[base + i] = bf2f(buf[i]);
}
#endif
// row pitch must be a multiple of 8 elems (16B units): S % 8 == 0 makes the XOR
// swizzle conflict-free for ldmatrix/stmatrix 8-row groups (unit%S differs only via q^r).
constexpr int LD128 = 128;
constexpr int LD64 = 64;
constexpr float RCP_LN2 = 1.4426950408889634f;
struct Smem {
__nv_bfloat16 kl[64 * LD128]; // k*e_up (A0 operand) -> qe/t/P/C store staging
__nv_bfloat16 qlq[64 * LD128]; // q*scale*e_up (Aqk operand), written in prologue
__nv_bfloat16 krt[128 * LD64]; // kr^T row-major [k=128][c=64]
__nv_bfloat16 kb[64 * LD128]; // k*beta*e_all -> w
__nv_bfloat16 vb[64 * LD128]; // v*beta -> u
__nv_bfloat16 ba[64 * LD64]; // chain ping / M
__nv_bfloat16 bb[64 * LD64]; // chain pong / Aqk
float beta_s[64];
float d_s[128];
float part[512];
};
// swizzled offsets: XOR the 8-element chunk index with (row & 7)
DEV int o128(int r, int c) { return r * LD128 + ((c >> 3) ^ (r & 7)) * 8 + (c & 7); }
DEV int o64(int r, int c) { return r * LD64 + ((c >> 3) ^ (r & 7)) * 8 + (c & 7); }
// acc[RM/16][RN/8][4] += A @ B for this warp's [RM, RN] slice.
// A: row-major [M,K] smem (swizzled), warp rows arow0..+RM, row pitch aw
// B: row-major [K,N] smem (swizzled), warp cols bcol0..+RN, row pitch bw
template <int RM, int RN>
DEV void dot_mn(float acc[RM / 16][RN / 8][4], const __nv_bfloat16* A, int arow0, int aw,
const __nv_bfloat16* B, int bcol0, int bw, int K) {
int l = lane_id();
for (int k0 = 0; k0 < K; k0 += 16) {
uint32_t a[RM / 16][4];
#pragma unroll
for (int mi = 0; mi < RM / 16; mi++) {
int r = arow0 + mi * 16 + (l % 16);
int c = k0 + (l / 16) * 8;
ldm_x4(a[mi][0], a[mi][1], a[mi][2], a[mi][3],
A + (aw == LD128 ? o128(r, c) : o64(r, c)));
}
uint32_t b[RN / 4]; // each x4t yields 2 frags x 2 regs
#pragma unroll
for (int np = 0; np < RN / 16; np++) {
uint32_t r0, r1, r2, r3;
int g = l / 8;
int kk = k0 + (l % 8) + (g & 1) * 8;
int nn = bcol0 + np * 16 + (g >> 1) * 8;
ldm_x4t(r0, r1, r2, r3, B + (bw == LD128 ? o128(kk, nn) : o64(kk, nn)));
b[np * 4 + 0] = r0; b[np * 4 + 1] = r1;
b[np * 4 + 2] = r2; b[np * 4 + 3] = r3;
}
#pragma unroll
for (int mi = 0; mi < RM / 16; mi++)
#pragma unroll
for (int ni = 0; ni < RN / 8; ni++)
mma_bf16(acc[mi][ni][0], acc[mi][ni][1], acc[mi][ni][2], acc[mi][ni][3],
a[mi][0], a[mi][1], a[mi][2], a[mi][3], b[ni * 2], b[ni * 2 + 1]);
}
}
// zero acc
template <int RM, int RN>
DEV void zero_acc(float acc[RM / 16][RN / 8][4]) {
#pragma unroll
for (int i = 0; i < RM / 16; i++)
#pragma unroll
for (int j = 0; j < RN / 8; j++)
#pragma unroll
for (int p = 0; p < 4; p++) acc[i][j][p] = 0.f;
}
// stage warp's [RM, RN] acc slice to a 64- or 128-wide smem tile at (row0, col0).
// mode: 0 plain, 1 negate. LD picks o64/o128 addressing.
template <int RM, int RN, int LD>
DEV void stage_acc(float acc[RM / 16][RN / 8][4], int row0, int col0, __nv_bfloat16* base,
bool negate) {
int l = lane_id();
#pragma unroll
for (int mi = 0; mi < RM / 16; mi++)
#pragma unroll
for (int ni = 0; ni < RN / 8; ni += 2) {
int r = row0 + mi * 16;
int c = col0 + ni * 8;
int lr = r + (l % 8) + ((l / 8) & 1) * 8;
int lc = c + (l / 16) * 8;
int off = (LD == LD128) ? o128(lr, lc) : o64(lr, lc);
float s0 = negate ? -acc[mi][ni][0] : acc[mi][ni][0];
float s1 = negate ? -acc[mi][ni][1] : acc[mi][ni][1];
float s2 = negate ? -acc[mi][ni][2] : acc[mi][ni][2];
float s3 = negate ? -acc[mi][ni][3] : acc[mi][ni][3];
float u0 = negate ? -acc[mi][ni + 1][0] : acc[mi][ni + 1][0];
float u1 = negate ? -acc[mi][ni + 1][1] : acc[mi][ni + 1][1];
float u2 = negate ? -acc[mi][ni + 1][2] : acc[mi][ni + 1][2];
float u3 = negate ? -acc[mi][ni + 1][3] : acc[mi][ni + 1][3];
stm_x4(base + off, pack_bf16(s0, s1), pack_bf16(s2, s3),
pack_bf16(u0, u1), pack_bf16(u2, u3));
}
}
// copy a [64,128] bf16 smem tile (swizzled, pitch LD128) to gmem rows rowG..rowG+63,
// gmem row pitch = gstride elements, i.e. dst + (rowG+i)*gstride + j.
DEV void tile_to_gmem(const __nv_bfloat16* src, __nv_bfloat16* dst, int rowG, long gstride) {
int t = threadIdx.x;
int r = t >> 2;
int u = (t & 3) * 4; // uint4 index (8 bf16 each)
const uint4* s4 = reinterpret_cast<const uint4*>(src);
uint4* d4 = reinterpret_cast<uint4*>(dst);
long grow = (long)(rowG + r) * (gstride >> 3); // gstride in bf16 -> in uint4
#pragma unroll
for (int i = 0; i < 4; i++) {
int uu = u + i;
int soff = o128(r, uu * 8) >> 3; // /8 elems -> /8 in uint4? (8 bf16 = 1 uint4) but swizzle offset is in elems: elem offset o128(r, uu*8) is a multiple of 8 -> uint4 index = /8
d4[grow + uu] = s4[soff];
}
}
// ---------------------------------------------------------------------------
// K1: grid (NT, BH), 256 threads (8 warps: wm = w>>2, wn = w&3).
// ---------------------------------------------------------------------------
__global__ void __launch_bounds__(256, 1) kda_k1(
const __nv_bfloat16* __restrict__ q,
const __nv_bfloat16* __restrict__ k,
const __nv_bfloat16* __restrict__ v,
const float* __restrict__ g,
const __nv_bfloat16* __restrict__ beta,
__nv_bfloat16* __restrict__ P,
__nv_bfloat16* __restrict__ Cm,
float* __restrict__ Pd,
__nv_bfloat16* __restrict__ qe,
__nv_bfloat16* __restrict__ tt,
long T, int H, long NT, float scale
#ifdef K1_DBG
, float* dbg
#endif
) {
extern __shared__ __nv_bfloat16 smem_raw[];
Smem& s = *reinterpret_cast<Smem*>(smem_raw);
int pid_t = blockIdx.x;
long bh = blockIdx.y;
long i_b = bh / H, i_h = bh - i_b * H;
long tok0 = (long)pid_t * 64;
int tid = threadIdx.x, l = lane_id(), w = tid >> 5;
int wm = w >> 2, wn = w & 3;
long HK = (long)H * 128;
long rowK = (i_b * T + tok0) * HK + i_h * 128; // k/v/q/g row (t-major, K-contig)
long rowB = (i_b * T + tok0) * H + i_h; // beta row
long outQ = (i_b * T + tok0) * HK + i_h * 128; // qe/t rows
// ---------------- prologue ----------------
// thread owns 4 columns c0..c0+3 and 8 rows r0..r0+7 (rg = tid>>5, cp = tid&31).
// All gmem loads vectorized: g float4, k/v/q 8B. part[] overlays s.ba (free here).
{
int cp = tid & 31;
int c0 = cp * 4;
int rg = tid >> 5;
int r0 = rg * 8;
const float4* gp = reinterpret_cast<const float4*>(g + rowK + (long)r0 * HK + c0);
const uint2* kp = reinterpret_cast<const uint2*>(k + rowK + (long)r0 * HK + c0);
const uint2* vp = reinterpret_cast<const uint2*>(v + rowK + (long)r0 * HK + c0);
const uint2* qp = reinterpret_cast<const uint2*>(q + rowK + (long)r0 * HK + c0);
float4 gx[8]; uint2 kr4[8], vr4[8], qr4[8];
#pragma unroll
for (int i = 0; i < 8; i++) gx[i] = gp[(long)i * (HK >> 2)];
#pragma unroll
for (int i = 0; i < 8; i++) kr4[i] = kp[(long)i * (HK >> 2)];
#pragma unroll
for (int i = 0; i < 8; i++) vr4[i] = vp[(long)i * (HK >> 2)];
#pragma unroll
for (int i = 0; i < 8; i++) qr4[i] = qp[(long)i * (HK >> 2)];
float acc[8][4];
float run[4] = {0.f, 0.f, 0.f, 0.f};
#pragma unroll
for (int i = 0; i < 8; i++) {
run[0] += gx[i].x * RCP_LN2; acc[i][0] = run[0];
run[1] += gx[i].y * RCP_LN2; acc[i][1] = run[1];
run[2] += gx[i].z * RCP_LN2; acc[i][2] = run[2];
run[3] += gx[i].w * RCP_LN2; acc[i][3] = run[3];
}
float* part = reinterpret_cast<float*>(s.ba); // [8][128], ba not used yet
#pragma unroll
for (int j = 0; j < 4; j++) part[rg * 128 + c0 + j] = run[j];
if (tid < 64) s.beta_s[tid] = bf2f(beta[rowB + (long)tid * H]);
__syncthreads();
float glast[4], pbelow[4];
#pragma unroll
for (int j = 0; j < 4; j++) {
float tot = 0.f, below = 0.f;
for (int rr = 0; rr < 8; rr++) {
float p = part[rr * 128 + c0 + j];
tot += p;
if (rr < rg) below += p;
}
glast[j] = tot;
pbelow[j] = below;
}
if (rg == 0) {
float4 dd = make_float4(exp2f(glast[0]), exp2f(glast[1]),
exp2f(glast[2]), exp2f(glast[3]));
reinterpret_cast<float4*>(s.d_s)[cp] = dd;
reinterpret_cast<float4*>(Pd + (bh * NT + pid_t) * 128)[cp] = dd;
}
__nv_bfloat162* kl2 = reinterpret_cast<__nv_bfloat162*>(s.kl);
__nv_bfloat162* ql2 = reinterpret_cast<__nv_bfloat162*>(s.qlq);
__nv_bfloat162* kb2 = reinterpret_cast<__nv_bfloat162*>(s.kb);
__nv_bfloat162* vb2 = reinterpret_cast<__nv_bfloat162*>(s.vb);
#pragma unroll
for (int i = 0; i < 8; i++) {
int r = r0 + i;
__nv_bfloat162 k01 = *reinterpret_cast<__nv_bfloat162*>(&kr4[i]);
__nv_bfloat162 k23 = *reinterpret_cast<__nv_bfloat162*>(&kr4[i].y);
__nv_bfloat162 v01 = *reinterpret_cast<__nv_bfloat162*>(&vr4[i]);
__nv_bfloat162 v23 = *reinterpret_cast<__nv_bfloat162*>(&vr4[i].y);
__nv_bfloat162 q01 = *reinterpret_cast<__nv_bfloat162*>(&qr4[i]);
__nv_bfloat162 q23 = *reinterpret_cast<__nv_bfloat162*>(&qr4[i].y);
float bt = s.beta_s[r];
float e_up[4], e_dn[4], ea[4];
#pragma unroll
for (int j = 0; j < 4; j++) {
float gc = acc[i][j] + pbelow[j];
e_up[j] = exp2f(gc - glast[j]);
e_dn[j] = exp2f(glast[j] - gc);
ea[j] = exp2f(gc);
}
int o = o128(r, c0) >> 1;
kl2[o] = __floats2bfloat162_rn(bf2f(k01.x) * e_up[0], bf2f(k01.y) * e_up[1]);
kl2[o + 1] = __floats2bfloat162_rn(bf2f(k23.x) * e_up[2], bf2f(k23.y) * e_up[3]);
ql2[o] = __floats2bfloat162_rn(bf2f(q01.x) * scale * e_up[0], bf2f(q01.y) * scale * e_up[1]);
ql2[o + 1] = __floats2bfloat162_rn(bf2f(q23.x) * scale * e_up[2], bf2f(q23.y) * scale * e_up[3]);
kb2[o] = __floats2bfloat162_rn(bf2f(k01.x) * bt * ea[0], bf2f(k01.y) * bt * ea[1]);
kb2[o + 1] = __floats2bfloat162_rn(bf2f(k23.x) * bt * ea[2], bf2f(k23.y) * bt * ea[3]);
vb2[o] = __floats2bfloat162_rn(bf2f(v01.x) * bt, bf2f(v01.y) * bt);
vb2[o + 1] = __floats2bfloat162_rn(bf2f(v23.x) * bt, bf2f(v23.y) * bt);
s.krt[o64(c0, r)] = __float2bfloat16(bf2f(k01.x) * e_dn[0]);
s.krt[o64(c0 + 1, r)] = __float2bfloat16(bf2f(k01.y) * e_dn[1]);
s.krt[o64(c0 + 2, r)] = __float2bfloat16(bf2f(k23.x) * e_dn[2]);
s.krt[o64(c0 + 3, r)] = __float2bfloat16(bf2f(k23.y) * e_dn[3]);
}
__syncthreads();
}
#ifdef K1_DBG
dump_bf(dbg, 0, s.kl, 64*LD128);
dump_bf(dbg, 61440, s.qlq, 64*LD128);
dump_f32(dbg, 70000, s.beta_s, 64);
dump_bf(dbg, 8192, s.krt, 128*LD64);
dump_bf(dbg, 16384, s.kb, 64*LD128);
dump_bf(dbg, 24576, s.vb, 64*LD128);
__syncthreads();
#endif
// ---------------- A0 = kl @ krT ; B = tril(A0*beta, -1) ; X = I - B ----------------
#ifdef K1_EARLY
if (K1_EARLY == 1) return;
#endif
float a0[2][2][4];
zero_acc<32, 16>(a0);
dot_mn<32, 16>(a0, s.kl, wm * 32, LD128, s.krt, wn * 16, LD64, 128);
__syncthreads(); // kl free after all warps' A0
#pragma unroll
for (int mi = 0; mi < 2; mi++)
#pragma unroll
for (int ni = 0; ni < 2; ni++)
#pragma unroll
for (int p = 0; p < 4; p++) {
int r, c;
// frag coords within [32,16] warp slice
r = wm * 32 + mi * 16 + (l >> 2);
c = wn * 16 + ni * 8 + ((l & 3) * 2 + (p & 1));
if (p >= 2) r += 8;
float val = (r > c) ? a0[mi][ni][p] * s.beta_s[r] : 0.f;
a0[mi][ni][p] = val;
}
float xacc[2][2][4];
#pragma unroll
for (int mi = 0; mi < 2; mi++)
#pragma unroll
for (int ni = 0; ni < 2; ni++)
#pragma unroll
for (int p = 0; p < 4; p++) {
int r = wm * 32 + mi * 16 + (l >> 2) + (p >= 2 ? 8 : 0);
int c = wn * 16 + ni * 8 + ((l & 3) * 2 + (p & 1));
xacc[mi][ni][p] = (r == c) ? 1.f : -a0[mi][ni][p];
}
#ifdef K1_DBG
if (blockIdx.x == 0 && blockIdx.y == 0 && w == 0)
for (int mi = 0; mi < 2; mi++) for (int ni = 0; ni < 2; ni++) for (int p = 0; p < 4; p++)
dbg[71000 + mi * 8 + ni * 4 + p] = a0[mi][ni][p];
__syncthreads();
#endif
stage_acc<32, 16, LD64>(a0, wm * 32, wn * 16, s.ba, false);
__syncthreads();
#ifdef K1_DBG
dump_bf(dbg, 32768, s.ba, 64*LD64);
__syncthreads();
#endif
#ifdef K1_EARLY
if (K1_EARLY == 2) return;
#endif
// ---------------- chain: X <- X + X @ B^(2^i) ----------------
__nv_bfloat16* cur = s.ba;
__nv_bfloat16* oth = s.bb;
#pragma unroll
for (int it = 0; it < 5; it++) {
float sq[2][2][4];
zero_acc<32, 16>(sq);
dot_mn<32, 16>(sq, cur, wm * 32, LD64, cur, wn * 16, LD64, 64);
stage_acc<32, 16, LD64>(sq, wm * 32, wn * 16, oth, false);
__syncthreads(); // cur dead: all warps done squaring
stage_acc<32, 16, LD64>(xacc, wm * 32, wn * 16, cur, false);
__syncthreads();
dot_mn<32, 16>(xacc, cur, wm * 32, LD64, oth, wn * 16, LD64, 64); // acc-init = X
__syncthreads();
__nv_bfloat16* tmp = cur; cur = oth; oth = tmp;
}
// M = xacc
stage_acc<32, 16, LD64>(xacc, wm * 32, wn * 16, s.ba, false);
__syncwarp();
#ifdef K1_DBG
dump_bf(dbg, 36864, s.ba, 64*LD64);
__syncthreads();
#endif
// ---------------- w = M @ kb ; u = M @ vb ----------------
float wacc[2][4][4], uacc[2][4][4];
zero_acc<32, 32>(wacc);
zero_acc<32, 32>(uacc);
dot_mn<32, 32>(wacc, s.ba, wm * 32, LD64, s.kb, wn * 32, LD128, 64);
dot_mn<32, 32>(uacc, s.ba, wm * 32, LD64, s.vb, wn * 32, LD128, 64);
stage_acc<32, 32, LD128>(wacc, wm * 32, wn * 32, s.kb, false);
stage_acc<32, 32, LD128>(uacc, wm * 32, wn * 32, s.vb, false);
__syncthreads();
#ifdef K1_DBG
dump_bf(dbg, 40960, s.kb, 64*LD128);
dump_bf(dbg, 49152, s.vb, 64*LD128);
__syncthreads();
#endif
#ifdef K1_EARLY
if (K1_EARLY == 3) return;
#endif
// ---------------- Aqk = tril(ql @ krT, incl diag) ; ql precomputed in prologue ----------------
float aqk[2][2][4];
zero_acc<32, 16>(aqk);
dot_mn<32, 16>(aqk, s.qlq, wm * 32, LD128, s.krt, wn * 16, LD64, 128);
#pragma unroll
for (int mi = 0; mi < 2; mi++)
#pragma unroll
for (int ni = 0; ni < 2; ni++)
#pragma unroll
for (int p = 0; p < 4; p++) {
int r = wm * 32 + mi * 16 + (l >> 2) + (p >= 2 ? 8 : 0);
int c = wn * 16 + ni * 8 + ((l & 3) * 2 + (p & 1));
aqk[mi][ni][p] = (r >= c) ? aqk[mi][ni][p] : 0.f; // ql already has scale
}
stage_acc<32, 16, LD64>(aqk, wm * 32, wn * 16, s.bb, false);
__syncthreads();
#ifdef K1_DBG
dump_bf(dbg, 57344, s.bb, 64*LD64);
#endif
// ---------------- qe = ql*d - Aqk@w ; t = Aqk@u ----------------
float qeacc[2][4][4], tacc[2][4][4];
zero_acc<32, 32>(qeacc);
zero_acc<32, 32>(tacc);
dot_mn<32, 32>(qeacc, s.bb, wm * 32, LD64, s.kb, wn * 32, LD128, 64);
dot_mn<32, 32>(tacc, s.bb, wm * 32, LD64, s.vb, wn * 32, LD128, 64);
// pack qe with + ql[r][c]*d[c], stage into kl (element-wise read then overwrite)
#pragma unroll
for (int mi = 0; mi < 2; mi++)
#pragma unroll
for (int ni = 0; ni < 4; ni += 2) {
int r = wm * 32 + mi * 16;
int cc = wn * 32 + ni * 8;
int lr = r + (l % 8) + ((l / 8) & 1) * 8;
int lc = cc + (l / 16) * 8;
int off = o128(lr, lc);
int rr = r + (l >> 2); // fragment row (address row differs!)
int cb = cc + 2 * (l & 3); // fragment cols
float q00 = bf2f(s.qlq[o128(rr, cb)]) * s.d_s[cb];
float q01 = bf2f(s.qlq[o128(rr, cb + 1)]) * s.d_s[cb + 1];
float q02 = bf2f(s.qlq[o128(rr + 8, cb)]) * s.d_s[cb];
float q03 = bf2f(s.qlq[o128(rr + 8, cb + 1)]) * s.d_s[cb + 1];
float q10 = bf2f(s.qlq[o128(rr, cb + 8)]) * s.d_s[cb + 8];
float q11 = bf2f(s.qlq[o128(rr, cb + 9)]) * s.d_s[cb + 9];
float q12 = bf2f(s.qlq[o128(rr + 8, cb + 8)]) * s.d_s[cb + 8];
float q13 = bf2f(s.qlq[o128(rr + 8, cb + 9)]) * s.d_s[cb + 9];
stm_x4(s.kl + off,
pack_bf16(q00 - qeacc[mi][ni][0], q01 - qeacc[mi][ni][1]),
pack_bf16(q02 - qeacc[mi][ni][2], q03 - qeacc[mi][ni][3]),
pack_bf16(q10 - qeacc[mi][ni + 1][0], q11 - qeacc[mi][ni + 1][1]),
pack_bf16(q12 - qeacc[mi][ni + 1][2], q13 - qeacc[mi][ni + 1][3]));
}
__syncwarp();
__syncthreads();
tile_to_gmem(s.kl, qe + outQ, 0, HK);
__syncthreads();
stage_acc<32, 32, LD128>(tacc, wm * 32, wn * 32, s.kl, false);
__syncthreads();
tile_to_gmem(s.kl, tt + outQ, 0, HK);
__syncthreads();
#ifdef K1_EARLY
if (K1_EARLY == 4) return;
#endif
// ---------------- P = -(krt @ w) ; C = krt @ u (all 8 warps, [64,32] tiles) --------
// output rows 0-63 stage into qlq, rows 64-127 into kl (both free); kb/vb hold w/u.
long pbase = ((i_b * H + i_h) * NT + pid_t) * 128 * 128;
{
float pacc[4][4][4];
zero_acc<64, 32>(pacc);
dot_mn<64, 32>(pacc, s.krt, wm * 64, LD64, s.kb, wn * 32, LD128, 64);
__syncthreads();
stage_acc<64, 32, LD128>(pacc, 0, wn * 32, wm ? s.kl : s.qlq, true);
__syncthreads();
tile_to_gmem(s.qlq, P + pbase, 0, 128);
tile_to_gmem(s.kl, P + pbase + 64 * 128, 0, 128);
__syncthreads();
}
{
float cacc[4][4][4];
zero_acc<64, 32>(cacc);
dot_mn<64, 32>(cacc, s.krt, wm * 64, LD64, s.vb, wn * 32, LD128, 64);
__syncthreads();
stage_acc<64, 32, LD128>(cacc, 0, wn * 32, wm ? s.kl : s.qlq, false);
__syncthreads();
tile_to_gmem(s.qlq, Cm + pbase, 0, 128);
tile_to_gmem(s.kl, Cm + pbase + 64 * 128, 0, 128);
__syncthreads();
}
}
void kda_k1_launch(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor g,
torch::Tensor beta, torch::Tensor P, torch::Tensor C, torch::Tensor Pd,
torch::Tensor qe, torch::Tensor t, long T, long H, long NT, double scale
#ifdef K1_DBG
, torch::Tensor dbg
#endif
) {
dim3 grid((unsigned)NT, (unsigned)(q.size(0) * H));
size_t smem = sizeof(Smem);
static bool set = false;
if (!set) {
cudaError_t e = cudaFuncSetAttribute(kda_k1, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
TORCH_CHECK(e == cudaSuccess, "setattr failed: ", cudaGetErrorString(e), " smem=", smem);
set = true;
}
cudaError_t el = cudaGetLastError();
TORCH_CHECK(el == cudaSuccess, "pre-launch: ", cudaGetErrorString(el));
kda_k1<<<grid, 256, smem>>>(
reinterpret_cast<const __nv_bfloat16*>(q.data_ptr()),
reinterpret_cast<const __nv_bfloat16*>(k.data_ptr()),
reinterpret_cast<const __nv_bfloat16*>(v.data_ptr()),
g.data_ptr<float>(),
reinterpret_cast<const __nv_bfloat16*>(beta.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(P.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(C.data_ptr()),
Pd.data_ptr<float>(),
reinterpret_cast<__nv_bfloat16*>(qe.data_ptr()),
reinterpret_cast<__nv_bfloat16*>(t.data_ptr()),
T, (int)H, NT, (float)scale
#ifdef K1_DBG
, dbg.data_ptr<float>()
#endif
);
}
"""
def _build_cuda_k1():
"""Compile the custom CUDA K1 kernel; returns None if unavailable."""
try:
if not torch.cuda.is_available():
return None
major, _ = torch.cuda.get_device_capability(0)
if major != 12:
return None
from torch.utils.cpp_extension import load_inline
bdir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "_build")
os.makedirs(bdir, exist_ok=True)
return load_inline(
name="kda_cuda_k1_v3",
cpp_sources="void kda_k1_launch(torch::Tensor q, torch::Tensor k, torch::Tensor v, "
"torch::Tensor g, torch::Tensor beta, torch::Tensor P, torch::Tensor C, "
"torch::Tensor Pd, torch::Tensor qe, torch::Tensor t, long T, long H, "
"long NT, double scale);",
cuda_sources=_CUDA_K1_SRC,
functions=["kda_k1_launch"],
extra_cuda_cflags=["-O3", "--use_fast_math", "-arch=sm_120a", "-std=c++20"],
build_directory=bdir,
verbose=False,
)
except Exception:
return None
_CUDA_K1 = None
_CUDA_K1_TRIED = False
def _run_k1(q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale, BH=1):
"""Dispatch K1 to the CUDA kernel when available, else Triton."""
global _CUDA_K1, _CUDA_K1_TRIED
if not _CUDA_K1_TRIED:
_CUDA_K1 = _build_cuda_k1()
_CUDA_K1_TRIED = True
if _CUDA_K1 is not None:
_CUDA_K1.kda_k1_launch(q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale)
else:
_kda_k1[(NT, BH_global(q))](
q, k, v, g, beta, P_, C_, Pd_, qe_, t_, T, H, NT, scale, num_warps=8)
_BUF_CACHE: dict = {}
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_forward(q, k, v, g, beta, self.scale, 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]
20260822_053629_zai-claude_glm-5.3_02_kda_cutlass