[KDA] Add CuteDSL Prefill Kernel on SM100 (#27488)

Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
Yuan Luo
2026-06-10 21:25:19 +08:00
committed by GitHub
co-authored by luoyuan.luo
parent 4faaa9ba92
commit 518e35fae7
9 changed files with 3045 additions and 5 deletions
@@ -0,0 +1,357 @@
"""
Benchmark & Correctness: Triton KDA vs CuTeDSL KDA (prefill, SM100 Blackwell).
Compares:
- Triton: sglang's chunk_kda (FLA chunkwise gated delta rule, per-channel gate)
- CuteDSL: kda_blackwell pipeline (fused Triton prologue -> kkt_inv_uw -> h -> o)
KDA differs from GDN by a PER-CHANNEL decay gate (g is [T, H, K], not scalar).
The cutedsl pipeline externalizes the per-channel decay into five pre-scaled
key/query tensors computed by a fused Triton prologue; the chunk metadata is
computed once and shared across layers in a real forward, so the benchmarked
cutedsl path precomputes it outside the timed region (the realistic ceiling).
Correctness is checked against the token-by-token fused_recurrent_kda ground
truth. Reports performance (ms, approx TFLOPS, TB/s, speedup).
Usage:
python bench_kda_prefill_cutedsl.py # default sweep
python bench_kda_prefill_cutedsl.py --mode bench # benchmark only
python bench_kda_prefill_cutedsl.py --mode correctness # correctness only
"""
import argparse
import os
import sys
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "python"))
import torch
import torch.nn.functional as F
from sglang.srt.layers.attention.fla.kda import chunk_kda, fused_recurrent_kda
from sglang.srt.layers.attention.linear.kernels.kda_blackwell import prepare_metadata
from sglang.srt.layers.attention.linear.kernels.kda_blackwell.kernel_h import (
kda_h_cutedsl,
)
from sglang.srt.layers.attention.linear.kernels.kda_blackwell.kernel_kkt_inv_uw import (
kkt_inv_uw_cutedsl,
)
from sglang.srt.layers.attention.linear.kernels.kda_blackwell.kernel_o import (
kda_o_cutedsl,
)
from sglang.srt.layers.attention.linear.kernels.kda_blackwell.prologue import (
kda_prologue,
)
BT = 64 # chunk size
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _l2norm(x: torch.Tensor) -> torch.Tensor:
return F.normalize(x.float(), p=2, dim=-1)
def kda_flops(total_seq_len, num_heads, head_k, head_v):
"""Per-token-per-head: k@v outer (2*K*V) + q@state (2*K*V), plus the intra-chunk
KKT (2*K*K averaged over the chunk). Approximate (ignores the inverse)."""
return total_seq_len * num_heads * (4 * head_k * head_v + 2 * head_k * head_k)
def kda_bytes(total_seq_len, num_heads, head_k, head_v, num_seqs, dtype):
elem = dtype.itemsize
q_b = total_seq_len * num_heads * head_k * elem
k_b = total_seq_len * num_heads * head_k * elem
v_b = total_seq_len * num_heads * head_v * elem
o_b = total_seq_len * num_heads * head_v * elem
g_b = total_seq_len * num_heads * head_k * 4 # per-channel gate, fp32
beta_b = total_seq_len * num_heads * 4
state_b = 2 * num_seqs * num_heads * head_k * head_v * 4 # fp32 r/w
return q_b + k_b + v_b + o_b + g_b + beta_b + state_b
# ---------------------------------------------------------------------------
# Input factory (single sequence per benchmark point, B=1)
# ---------------------------------------------------------------------------
def make_inputs(T, H, K, V, device, dtype, seed=42):
torch.manual_seed(seed)
q = _l2norm(torch.randn(1, T, H, K, device=device)).to(dtype)
k = _l2norm(torch.randn(1, T, H, K, device=device)).to(dtype)
v = torch.randn(1, T, H, V, device=device).to(dtype)
# Mild per-channel gate (real Kimi-Linear regime; keeps exp() in fp32 range).
A_log = torch.randn(H, device=device) * 0.5 - 1.5
dt_bias = torch.randn(H, K, device=device) * 0.1
g_raw = torch.randn(1, T, H, K, device=device)
g_act = (
-A_log.exp().view(1, 1, H, 1) * F.softplus(g_raw + dt_bias.view(1, 1, H, K))
).float()
beta = torch.sigmoid(torch.randn(1, T, H, device=device)).float()
return dict(q=q, k=k, v=v, g_act=g_act, beta=beta, T=T, H=H, K=K, V=V)
# ---------------------------------------------------------------------------
# Runners
# ---------------------------------------------------------------------------
def run_recurrent(inp, scale):
"""Token-by-token ground truth. Returns (o [T,H,V], state [1,H,V,K])."""
cu = torch.tensor([0, inp["T"]], dtype=torch.int64, device=inp["q"].device)
h0 = torch.zeros(1, inp["H"], inp["V"], inp["K"], device=inp["q"].device)
o, state = fused_recurrent_kda(
q=inp["q"],
k=inp["k"],
v=inp["v"],
g=inp["g_act"],
beta=inp["beta"],
scale=scale,
initial_state=h0,
inplace_final_state=False,
use_qk_l2norm_in_kernel=False,
cu_seqlens=cu,
)
return o[0], state
def cutedsl_buffers(inp, num_sms, device):
"""Precompute metadata + preallocate (shared across layers in a real forward)."""
T, H, K, V = inp["T"], inp["H"], inp["K"], inp["V"]
cu = torch.tensor([0, T], dtype=torch.int32, device=device)
ci, co, tc, total = prepare_metadata(cu)
pad_t = total * BT
return dict(
cu=cu,
ci=ci,
co=co,
tc=tc,
total=total,
num_sms=num_sms,
h0=torch.zeros(1, H, V, K, device=device, dtype=torch.float32),
U=torch.empty(pad_t, H, V, device=device, dtype=torch.bfloat16),
W=torch.empty(pad_t, H, K, device=device, dtype=torch.bfloat16),
V_new=torch.empty(pad_t, H, V, device=device, dtype=torch.bfloat16),
h_chunks=torch.empty(total, H, V, K, device=device, dtype=torch.bfloat16),
ht=torch.empty(1, H, V, K, device=device, dtype=torch.float32),
o=torch.empty(T, H, V, device=device, dtype=torch.bfloat16),
)
def run_cutedsl_pipeline(inp, buf, scale):
"""The fused-prologue + 3 cutedsl kernels (metadata precomputed in buf)."""
q3, k3, v3 = inp["q"][0], inp["k"][0], inp["v"][0]
g3, beta3 = inp["g_act"][0], inp["beta"][0].contiguous()
KL, KR, KG, qg, qg2, g_cu = kda_prologue(
q3, k3, g3, scale, buf["cu"], buf["ci"], buf["total"]
)
kkt_inv_uw_cutedsl(
KL,
KR,
KG,
v3,
buf["U"],
buf["W"],
beta3,
buf["cu"],
buf["ci"],
buf["tc"],
num_sms=buf["num_sms"],
)
kda_h_cutedsl(
KR,
buf["U"],
buf["W"],
buf["V_new"],
g_cu,
buf["h_chunks"],
buf["h0"],
buf["ht"],
buf["cu"],
buf["co"],
)
kda_o_cutedsl(
qg,
qg2,
KR,
buf["V_new"],
buf["h_chunks"],
buf["o"],
buf["cu"],
buf["ci"],
buf["tc"],
num_sms=buf["num_sms"],
)
return buf["o"], buf["ht"]
# ---------------------------------------------------------------------------
# Correctness
# ---------------------------------------------------------------------------
def check_shape(T, H, K, V, device, dtype, num_sms):
tag = f"T={T:>5} H={H:>2} K={K:>3} V={V:>3}"
if K != 128 or V != 128:
print(f" [SKIP] {tag} (cutedsl requires K=V=128)")
return True
scale = K**-0.5
inp = make_inputs(T, H, K, V, device, dtype)
o_ref, state_ref = run_recurrent(inp, scale)
try:
buf = cutedsl_buffers(inp, num_sms, device)
o, ht = run_cutedsl_pipeline(inp, buf, scale)
torch.cuda.synchronize()
except Exception as e: # noqa: BLE001
print(f" [SKIP] {tag} (cutedsl error: {e})")
return True
finite = bool(torch.isfinite(o).all() and torch.isfinite(ht).all())
o_err = (o.float() - o_ref.float()).abs().max().item()
s_err = (ht.float() - state_ref.float()).abs().max().item()
ok = finite and o_err < 1e-2 and s_err < 5e-2
status = "PASS" if ok else "FAIL"
print(
f" [{status}] {tag} | o_err {o_err:.2e} state_err {s_err:.2e} finite={finite}"
)
return ok
# ---------------------------------------------------------------------------
# Benchmark
# ---------------------------------------------------------------------------
def bench_shape(T, H, K, V, device, dtype, num_sms):
import triton.testing
if K != 128 or V != 128:
print(f" [SKIP] T={T} H={H} K={K} V={V} (cutedsl K=V=128 only)")
return
scale = K**-0.5
inp = make_inputs(T, H, K, V, device, dtype)
q, k, v = inp["q"], inp["k"], inp["v"]
g_act, beta = inp["g_act"], inp["beta"]
h0f = torch.zeros(1, H, K, V, device=device, dtype=torch.float32)
idx = torch.zeros(1, dtype=torch.int32, device=device)
def fn_triton():
chunk_kda(
q=q,
k=k,
v=v,
g=g_act,
beta=beta,
scale=scale,
initial_state=h0f,
initial_state_indices=idx,
use_qk_l2norm_in_kernel=False,
cu_seqlens=None,
A_log=None,
dt_bias=None,
lower_bound=None,
)
buf = cutedsl_buffers(inp, num_sms, device)
def fn_cutedsl():
run_cutedsl_pipeline(inp, buf, scale)
quantiles = [0.5, 0.2, 0.8]
fn_triton()
fn_cutedsl()
torch.cuda.synchronize()
ms_triton, _, _ = triton.testing.do_bench_cudagraph(fn_triton, quantiles=quantiles)
ms_cutedsl, _, _ = triton.testing.do_bench_cudagraph(
fn_cutedsl, quantiles=quantiles
)
flops = kda_flops(T, H, K, V)
mem_bytes = kda_bytes(T, H, K, V, 1, dtype)
speedup = ms_triton / ms_cutedsl if ms_cutedsl > 0 else float("inf")
print(
f" {H:>3} {T:>7} | "
f"{ms_triton:>8.3f} {flops / ms_triton / 1e9:>7.2f} {mem_bytes / ms_triton / 1e9:>7.2f} | "
f"{ms_cutedsl:>8.3f} {flops / ms_cutedsl / 1e9:>7.2f} {mem_bytes / ms_cutedsl / 1e9:>7.2f} | "
f"{speedup:>7.2f}x"
)
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def run_correctness(device, dtype, H, num_sms):
print("=" * 72)
print("Correctness: cutedsl pipeline vs fused_recurrent_kda (ground truth)")
print("=" * 72)
all_pass = True
for T in (128, 192, 256, 512, 1024):
if not check_shape(T, H, 128, 128, device, dtype, num_sms):
all_pass = False
print("\nALL PASSED." if all_pass else "\nSOME FAILED.")
return all_pass
def run_benchmark(device, dtype, args, num_sms):
print()
print("=" * 92)
print("Benchmark: Triton chunk_kda vs CuTeDSL pipeline (do_bench_cudagraph)")
print("=" * 92)
print(f" Device SMs={num_sms}, K=V=128, dtype={dtype}, metadata precomputed")
print(
f" {'H':>3} {'T':>7} | "
f"{'tri(ms)':>8} {'TFLOP':>7} {'TB/s':>7} | "
f"{'cute(ms)':>8} {'TFLOP':>7} {'TB/s':>7} | {'speedup':>8}"
)
print(" " + "-" * 84)
for H in args.num_heads:
for T in args.seq_lens:
bench_shape(T, H, 128, 128, device, dtype, num_sms)
def main():
parser = argparse.ArgumentParser(
description="Benchmark & Correctness: Triton KDA vs CuTeDSL KDA (SM100)"
)
parser.add_argument(
"--mode", choices=["all", "correctness", "bench"], default="all"
)
parser.add_argument("--dtype", choices=["float16", "bfloat16"], default="bfloat16")
parser.add_argument("--num-heads", type=int, nargs="+", default=[32])
parser.add_argument(
"--seq-lens", type=int, nargs="+", default=[512, 1024, 2048, 4096, 8192]
)
args = parser.parse_args()
device = "cuda"
dtype = getattr(torch, args.dtype)
cap = torch.cuda.get_device_capability()
print(f"Device: {torch.cuda.get_device_name()} (SM {cap[0]}{cap[1]})")
if cap[0] < 10:
print("ERROR: CuTeDSL KDA prefill requires SM100+ (Blackwell). Exiting.")
return 1
num_sms = torch.cuda.get_device_properties(0).multi_processor_count
if args.mode in ("all", "correctness"):
all_pass = run_correctness(device, dtype, args.num_heads[0], num_sms)
if not all_pass and args.mode == "all":
print("\nSkipping benchmark due to correctness failures.")
return 1
if args.mode in ("all", "bench"):
run_benchmark(device, dtype, args, num_sms)
return 0
if __name__ == "__main__":
sys.exit(main())
@@ -60,10 +60,28 @@ class KDAKernelDispatcher:
if prefill_backend.is_triton():
self.extend_kernel = triton_kernel
elif prefill_backend.is_cutedsl():
if not is_cuda():
raise ValueError("KDA CuTe DSL backend requires CUDA")
from sglang.srt.layers.attention.linear.kernels.kda_cutedsl import (
CuteDSLKDAKernel,
)
cutedsl_kernel = CuteDSLKDAKernel()
if getattr(cutedsl_kernel, "supports_prefill", False):
# SM100 chunk prefill pipeline.
self.extend_kernel = cutedsl_kernel
else:
# CuTe DSL prefill kernels need SM100 (Blackwell); on older GPUs
# fall back to the Triton chunk kernel.
self.extend_kernel = triton_kernel
rank0_log(
"KDA cutedsl prefill needs SM100; falling back to Triton extend."
)
else:
raise ValueError(
f"Unsupported KDA prefill backend: {prefill_backend}. "
"KDA currently only supports 'triton'."
"KDA supports 'triton' or 'cutedsl' (cutedsl prefill needs SM100)."
)
self.supports_packed_decode = getattr(
@@ -0,0 +1,221 @@
# SPDX-License-Identifier: Apache-2.0
# KDA (Kimi Delta Attention) SM100/Blackwell CuteDSL prefill pipeline.
#
# Mirrors gdn_blackwell but for KDA's PER-CHANNEL decay gate. A fused Triton
# prologue computes the per-chunk cumsum g_cu and five pre-scaled key/query
# tensors; three cutedsl kernels then run the chunked gated delta rule:
# prologue -> kkt_inv_uw (U,W) -> h (V_new, per-chunk state, final state) -> o
import torch
from .kernel_h import kda_h_cutedsl
from .kernel_kkt_inv_uw import kkt_inv_uw_cutedsl
from .kernel_o import kda_o_cutedsl
from .prologue import kda_prologue
__all__ = ["chunk_kda_cutedsl", "prepare_metadata"]
def prepare_metadata(cu_seqlens: torch.Tensor, chunk_size: int = 64):
"""Build (chunk_indices [NT,2], chunk_offsets [N+1], total_chunks [1]).
chunk_indices[g] = (seq_id, local_chunk_id) for global chunk g.
chunk_offsets[s] = number of chunks before sequence s.
"""
dev = cu_seqlens.device
cs = cu_seqlens.to(torch.int64)
seqlens = cs[1:] - cs[:-1]
nchunks = (seqlens + chunk_size - 1) // chunk_size # [N]
n = seqlens.numel()
chunk_offsets = torch.zeros(n + 1, dtype=torch.int32, device=dev)
chunk_offsets[1:] = nchunks.cumsum(0).to(torch.int32)
total = int(chunk_offsets[-1].item())
seq_id = torch.repeat_interleave(torch.arange(n, device=dev), nchunks)
local = torch.arange(total, device=dev) - chunk_offsets[seq_id].to(torch.int64)
chunk_indices = torch.stack(
[seq_id.to(torch.int32), local.to(torch.int32)], dim=1
).contiguous()
total_chunks = torch.tensor([total], dtype=torch.int32, device=dev)
return chunk_indices, chunk_offsets, total_chunks, total
# Per-(Hv,K,V,device) grow-only scratch workspace. The cutedsl KKT/h/o kernels
# are fast; the per-call PyTorch overhead (re-allocating + re-zeroing the eye and
# the two pack buffers ~200MB/call, metadata recompute, a `.item()` sync) was what
# dragged the full function below Triton. Reusing scratch across calls removes it.
# Safe because KDA layers run sequentially on one CUDA stream (the next call's
# kernels are ordered after this call's), and only the returned o/ht are fresh.
_KDA_WS: dict = {}
def _kda_workspace(q, T, Hv, K, V, cu_seqlens):
import torch as _t
dev = q.device
# Key by the current CUDA stream too: the scratch is process-global and
# mutable, so two KDA forwards running concurrently on different streams
# (e.g. two-batch overlap) must not share buffers. Within one forward all
# KDA layers run on the same stream -> same key -> the reuse benefit holds.
stream = _t.cuda.current_stream(device=dev).cuda_stream
key = (Hv, K, V, dev, q.dtype, stream)
ws = _KDA_WS.get(key)
# metadata: recompute only when cu_seqlens changes (object identity -> no
# sync; within one forward all KDA layers share the same cu_seqlens object).
if ws is None or ws["cu"] is not cu_seqlens:
ci, co, tcs, total = prepare_metadata(cu_seqlens)
else:
ci, co, tcs, total = ws["ci"], ws["co"], ws["tcs"], ws["total"]
pad_t = total * 64
if ws is None or ws["Tcap"] < T or ws["padcap"] < pad_t or ws["totalcap"] < total:
Tcap = T if ws is None else max(T, ws["Tcap"])
padcap = pad_t if ws is None else max(pad_t, ws["padcap"])
totalcap = total if ws is None else max(total, ws["totalcap"])
ws = {
"kL": q.new_zeros(Tcap, Hv, K, dtype=_t.bfloat16),
"qg2": q.new_zeros(Tcap, Hv, K, dtype=_t.bfloat16),
"eye": q.new_zeros(Tcap, Hv, K, dtype=_t.bfloat16),
"U": q.new_empty(padcap, Hv, V, dtype=_t.bfloat16),
"W": q.new_empty(padcap, Hv, K, dtype=_t.bfloat16),
"Vn": q.new_empty(padcap, Hv, V, dtype=_t.bfloat16),
"hc": q.new_empty(totalcap, Hv, V, K, dtype=_t.bfloat16),
"Tcap": Tcap,
"padcap": padcap,
"totalcap": totalcap,
"cu": None,
"eye_hw": 0,
}
_KDA_WS[key] = ws
ws["ci"], ws["co"], ws["tcs"], ws["total"] = ci, co, tcs, total
# eye is the one-hot(chunk-position) identity injection: recompute only on a
# cu_seqlens change. Clear the prior high-water region then scatter the new 1s.
if ws["cu"] is not cu_seqlens:
eye = ws["eye"]
hw = max(ws["eye_hw"], T)
eye[:hw].zero_()
# Match cu_seqlens' dtype (typically int32) so searchsorted/indexing avoid
# the int64 casts, while staying correct if cu_seqlens is passed as int64.
tok = _t.arange(T, device=dev, dtype=cu_seqlens.dtype)
seq_of = _t.searchsorted(cu_seqlens, tok, right=True) - 1
pos = (tok - cu_seqlens[seq_of]) % 64
eye[tok, :, pos] = 1.0
ws["eye_hw"] = T
ws["cu"] = cu_seqlens
return ws, ci, co, tcs, total, pad_t
def chunk_kda_cutedsl(
q: torch.Tensor, # [T, Hv, K] bf16, L2-normed
k: torch.Tensor, # [T, Hv, K] bf16, L2-normed
v: torch.Tensor, # [T, Hv, V] bf16
g: torch.Tensor, # [T, Hv, K] log-decay. RAW if A_log given, else pre-activated
beta: torch.Tensor, # [T, Hv] fp32, post-sigmoid
h0: torch.Tensor, # [N, Hv, V, K] (initial recurrent state, [V,K] layout)
cu_seqlens: torch.Tensor,
scale: float | None = None,
num_sms: int | None = None,
A_log: torch.Tensor | None = None, # [Hv]; if set, activate g internally
dt_bias: torch.Tensor | None = None, # [Hv, K] or [Hv*K]
lower_bound: float | None = None,
):
"""Run the KDA chunk gated-delta-rule prefill. Returns (o [T,Hv,V], ht [N,Hv,V,K])."""
import torch.nn.functional as F
T, Hv, K = q.shape
V = v.shape[-1]
if scale is None:
scale = K**-0.5
if num_sms is None:
num_sms = torch.cuda.get_device_properties(q.device).multi_processor_count
# Gate activation (standard KDA gate). Fused into the prologue is a B2 TODO;
# for now a small PyTorch pass, matching chunk_kda's kda_gate_chunk_cumsum.
if A_log is not None:
if lower_bound is not None:
raise NotImplementedError(
"KDA cutedsl: safe_gate (lower_bound) not yet supported"
)
x = g.float()
if dt_bias is not None:
x = x + dt_bias.float().view(1, Hv, K)
g_act = -torch.exp(A_log.float()).view(1, Hv, 1) * F.softplus(x)
else:
g_act = g.float()
# Reusable scratch (eye/pack/U/W/V_new/h_chunks) + cached metadata; only the
# returned o/ht are freshly allocated. This removes the ~0.2-0.6ms/call host
# overhead (re-alloc + re-zero of ~200MB + metadata sync) that otherwise drags
# the (fast) cutedsl kernels below Triton.
ws, chunk_indices, chunk_offsets, total_chunks, total, pad_t = _kda_workspace(
q, T, Hv, K, V, cu_seqlens
)
# KL/qg2 from the prologue fold the decay with a chunk-global g_last reference
# (exp(g_cu - g_last)), which overflows fp32 for real per-channel gates. They
# are recomputed below; the prologue still gives the bounded KR/KG/qg/g_cu.
_, KR, KG, qg, _, g_cu = kda_prologue(
q, k, g_act, float(scale), cu_seqlens, chunk_indices, total
)
# Sub-chunk-normalized intra-chunk gated KKT / QK from the FLA kernel (stable),
# injected through the cutedsl KKT/Aqk MMAs as an identity-right-operand pass:
# with kL'=M (M in the first 64 K-slots) and kR'=onehot(chunk-pos), the MMA
# kL'@kR'.T == M, so kkt_inv_uw/kernel_o see the correct matrix without overflow.
from sglang.srt.layers.attention.fla.kda import chunk_kda_scaled_dot_kkt_fwd
ones_beta = q.new_ones(1, T, Hv, dtype=torch.float32)
M_kk, M_qk = chunk_kda_scaled_dot_kkt_fwd(
q.unsqueeze(0).contiguous(),
k.unsqueeze(0).contiguous(),
gk=g_cu.unsqueeze(0),
beta=ones_beta,
scale=float(scale),
cu_seqlens=cu_seqlens,
chunk_size=64,
)
# Pack M into the first 64 K-slots of the reused buffers; cols [64:128] stay 0
# (never written since the one-time zeroed alloc), so the MxI injection is exact.
kL_inj = ws["kL"][:T]
qg2_inj = ws["qg2"][:T]
kL_inj[:, :, :64] = M_kk[0].to(torch.bfloat16)
qg2_inj[:, :, :64] = M_qk[0].to(torch.bfloat16)
eye = ws["eye"][:T]
U = ws["U"][:pad_t]
W = ws["W"][:pad_t]
kkt_inv_uw_cutedsl(
kL_inj,
eye,
KG,
v,
U,
W,
beta,
cu_seqlens,
chunk_indices,
total_chunks,
num_sms=num_sms,
)
V_new = ws["Vn"][:pad_t]
h_chunks = ws["hc"][:total]
ht = torch.empty_like(h0)
kda_h_cutedsl(KR, U, W, V_new, g_cu, h_chunks, h0, ht, cu_seqlens, chunk_offsets)
o = q.new_empty(T, Hv, V, dtype=torch.bfloat16)
kda_o_cutedsl(
qg,
qg2_inj,
eye,
V_new,
h_chunks,
o,
cu_seqlens,
chunk_indices,
total_chunks,
num_sms=num_sms,
)
return o, ht
@@ -0,0 +1,690 @@
# SPDX-License-Identifier: Apache-2.0
# KDA (Kimi Delta Attention) SM100 chunk recurrent-state kernel.
#
# Idea is adopted from GDN blackwell kernel. KDA differs from GDN only in the
# decay gate, which is PER-CHANNEL (one decay per key-dim k) instead of a single
# scalar per head. The hard cross-token part of the per-channel decay is folded
# OUTSIDE this kernel into the pre-scaled key tensor `kg`:
#
# kg[c, k] = k[c, k] * exp(g_cu_last[k] - g_cu[c, k]) (bounded, <= |k|)
#
# so the only in-kernel gate logic that remains is:
# 1. state decay is PER-COLUMN: H[v, k] *= exp(g_cu_last[k]) (not a scalar)
# 2. the H_new MMA consumes `kg` (pre-scaled) instead of raw K, and v_new stays
# RAW (GDN instead scales v_new by the scalar exp(g_last - g_t) and uses raw K).
#
# Math per chunk (state S stored transposed as H = [V, K]):
# V_new = U - W @ S (gate-free; W already gated in kkt stage)
# H_scaled[v, k] = H[v, k] * exp(g_cu_last[k])
# H_new = H_scaled + V_new.T @ kg
from functools import cache
import cutlass
import torch
from cuda.bindings.driver import CUstream
from cutlass import BFloat16, Float32, Int32, Int64, Uint32, cute
from cutlass.cute.nvgpu import cpasync, warp
from quack.compile_utils import make_fake_tensor
from sglang.srt.layers.attention.cute_utils import (
EVICT_FIRST,
_tcgen05,
cvt,
fence_before_tma_store,
simple_tma_copy,
)
class Sm100KdaChunkHKernel:
"""KDA per-chunk recurrent-state update (see module docstring)."""
def __init__(
self,
H: int,
Hv: int,
K_dim: int,
V_dim: int,
h_dtype: cutlass.Numeric = Float32,
BT: int = 64,
num_stages: int = 2,
) -> None:
assert Hv % H == 0
assert K_dim == V_dim == 128
assert BT == 64
self.H = H
self.Hv = Hv
self.K_dim = K_dim
self.V_dim = V_dim
self.h_dtype = h_dtype
self.BT = BT
self.num_stages = num_stages
self.num_warps = 10
@cute.jit
def _make_bf16_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def _make_h_tma_args(self, tensor: cute.Tensor, op: cpasync.TmaCopyOp):
num_elems = 128 // (tensor.element_type.width // 8)
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(1, 1, self.V_dim, (num_elems, self.K_dim // num_elems)),
stride=(0, 0, num_elems, (1, self.V_dim * num_elems)),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, None, num_elems)),
slayout,
cta_tiler=(1, 1, self.V_dim, self.K_dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
K: cute.Tensor, # KDA: this is `kg`, the per-channel pre-scaled key [T, Hv, K]
V: cute.Tensor, # = U from kkt stage
W: cute.Tensor,
V_new: cute.Tensor,
g_cu: cute.Tensor, # KDA: [T, Hv, K] per-channel cumsum
h: cute.Tensor,
h0: cute.Tensor,
ht: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_offsets: cute.Tensor,
stream: CUstream,
):
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
K_args = self._make_bf16_tma_args(K, self.K_dim, tma_g2s, self.num_stages)
V_args = self._make_bf16_tma_args(V, self.V_dim, tma_g2s, self.num_stages)
W_args = self._make_bf16_tma_args(W, self.K_dim, tma_g2s, self.num_stages)
V_new_args = self._make_bf16_tma_args(V_new, self.V_dim, tma_s2g, 1)
H0_args = self._make_h_tma_args(h0, tma_g2s)
HT_args = self._make_h_tma_args(ht, tma_s2g)
H_args = self._make_h_tma_args(h, tma_s2g)
grid = (self.Hv, h0.shape[0], 1)
block = (self.num_warps * 32, 1, 1)
self.kernel(
K_args,
V_args,
W_args,
V_new_args,
H0_args,
HT_args,
H_args,
g_cu,
cu_seqlens,
chunk_offsets,
).launch(grid=grid, block=block, stream=stream)
@cute.kernel
def kernel(
self,
K_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
V_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
W_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
V_new_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
H0_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
HT_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
H_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
g_cu: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_offsets: cute.Tensor,
):
tid, _, _ = cute.arch.thread_idx()
head_id, seq_id, _ = cute.arch.block_idx()
warp_id = cute.arch.make_warp_uniform(tid // 32)
lane_id = tid % 32
BT = self.BT
V_dim = self.V_dim
K_dim = self.K_dim
num_stages = self.num_stages
is_f32 = self.h_dtype == Float32
K_tma_atom, tmaK, sK_layout = K_args
V_tma_atom, tmaV, sV_layout = V_args
W_tma_atom, tmaW, sW_layout = W_args
V_new_tma_atom, tmaV_new, sV_new_layout = V_new_args
H0_tma_atom, tmaH0, sH0_layout = H0_args
HT_tma_atom, tmaHT, _ = HT_args
H_tma_atom, tmaH, sH_layout = H_args
def allocate_tensor(smem, dtype, layout):
return smem.allocate_tensor(
dtype, layout.outer, byte_alignment=128, swizzle=layout.inner
)
smem = cutlass.utils.SmemAllocator()
sW = allocate_tensor(smem, BFloat16, sW_layout)[None, 0, None, None]
sV = allocate_tensor(smem, BFloat16, sV_layout)[None, 0, None, None]
sK = allocate_tensor(smem, BFloat16, sK_layout)[None, 0, None, None]
sH0 = allocate_tensor(smem, self.h_dtype, sH0_layout)[0, 0, None, None]
sH = allocate_tensor(smem, BFloat16, sH_layout)[0, 0, None, None]
sV_new = allocate_tensor(smem, BFloat16, sV_new_layout)[None, 0, None, 0]
# KDA: per-channel end-of-chunk decay exp(g_cu_last[k]); shared by all V-rows.
s_gl_exp = smem.allocate_array(Float32, K_dim)
tma_mbar = smem.allocate_array(Int64, num_stages)
wh_in_mbar = smem.allocate_array(Int64, num_stages)
wh_done_mbar = smem.allocate_array(Int64, num_stages)
vk_in_mbar = smem.allocate_array(Int64, num_stages)
vk_done_mbar = smem.allocate_array(Int64, num_stages)
h0_mbar = smem.allocate_array(Int64, 1)
taddr = smem.allocate(Int32, 4)
wh_tmem = 0
vk_tmem = wh_tmem + BT
h_tmem_base = vk_tmem + K_dim
v_tmem_base = h_tmem_base + K_dim // 2
if warp_id == 0:
with cute.arch.elect_one():
for i in cutlass.range_constexpr(num_stages):
cute.arch.mbarrier_init(tma_mbar + i, 1)
cute.arch.mbarrier_init(wh_in_mbar + i, 256)
cute.arch.mbarrier_init(wh_done_mbar + i, 1)
cute.arch.mbarrier_init(vk_in_mbar + i, 256)
cute.arch.mbarrier_init(vk_done_mbar + i, 1)
cute.arch.mbarrier_init(h0_mbar, 1)
cute.arch.mbarrier_init_fence()
elif warp_id == 1:
cpasync.prefetch_descriptor(H0_tma_atom)
cpasync.prefetch_descriptor(W_tma_atom)
cpasync.prefetch_descriptor(V_tma_atom)
cpasync.prefetch_descriptor(K_tma_atom)
cpasync.prefetch_descriptor(HT_tma_atom)
cpasync.prefetch_descriptor(H_tma_atom)
cpasync.prefetch_descriptor(V_new_tma_atom)
cute.arch.sync_threads()
bos = cu_seqlens[seq_id]
eos = cu_seqlens[seq_id + 1]
seqlen = eos - bos
num_chunks = cute.ceil_div(seqlen, BT)
if warp_id == 9:
# TMA warp
stage_id = 0
parity = 1
chunk_offset = chunk_offsets[seq_id]
# load H0
with cute.arch.elect_one():
H0_size = V_dim * K_dim * self.h_dtype.width // 8
cute.arch.mbarrier_arrive_and_expect_tx(h0_mbar, H0_size)
simple_tma_copy(
H0_tma_atom, tmaH0[seq_id, head_id, None, None], sH0, h0_mbar
)
gW_tiles = cute.logical_divide(tmaW[None, head_id, None], (BT, None))
gV_tiles = cute.logical_divide(tmaV[None, head_id, None], (BT, None))
# KDA: kg is per v-head [T, Hv, K], index by head_id (G=1 => same as k_head_id)
gK_tiles = cute.logical_divide(
cute.domain_offset((bos, 0), tmaK[None, head_id, None]),
(BT, None),
)
for chunk_id in range(num_chunks):
mbar = tma_mbar + stage_id
gW = gW_tiles[(None, chunk_offset + chunk_id), None]
gV = gV_tiles[(None, chunk_offset + chunk_id), None]
gK = gK_tiles[(None, chunk_id), None]
cute.arch.mbarrier_wait(vk_done_mbar + stage_id, parity)
with cute.arch.elect_one():
STAGE_SIZE = BT * (K_dim + V_dim + K_dim) * 2
cute.arch.mbarrier_arrive_and_expect_tx(mbar, STAGE_SIZE)
simple_tma_copy(
W_tma_atom, gW, sW[None, None, stage_id], mbar, EVICT_FIRST
)
simple_tma_copy(
V_tma_atom, gV, sV[None, None, stage_id], mbar, EVICT_FIRST
)
simple_tma_copy(K_tma_atom, gK, sK[None, None, stage_id], mbar)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
elif warp_id == 8:
# MMA warp -- IDENTICAL to GDN: sK now holds kg, so V_new.T@kg falls out.
_tcgen05.alloc(taddr)
stage_id = 0
parity = 0
wh_idesc = _tcgen05.make_bf16_idesc(V_dim, BT, negate_A=True)
vk_idesc = _tcgen05.make_bf16_idesc(V_dim, K_dim, transpose_B=True)
sdesc_template = _tcgen05.make_sdesc_128B_swizzle(BT * 128)
if cutlass.const_expr(not is_f32):
Haddr0 = sH0[None, None].iterator.toint()
Waddr0 = sW[None, None, stage_id].iterator.toint()
hdesc0_base = sdesc_template | (Haddr0 >> 4)
wdesc0_base = sdesc_template | (Waddr0 >> 4)
cute.arch.mbarrier_wait(tma_mbar + stage_id, parity)
cute.arch.mbarrier_wait(wh_in_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(K_dim // 64):
for j in cutlass.range_constexpr(64 // 16):
hdesc0 = hdesc0_base | ((i * V_dim * 128 + j * 32) >> 4)
wdesc0 = wdesc0_base | ((i * BT * 128 + j * 32) >> 4)
_tcgen05.mma_f16(wh_tmem, hdesc0, wdesc0, wh_idesc, True)
_tcgen05.commit(wh_done_mbar + stage_id)
Kaddr0 = sK[None, None, stage_id].iterator.toint()
kdesc0_base = sdesc_template | (Kaddr0 >> 4)
cute.arch.mbarrier_wait(vk_in_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for k in cutlass.range_constexpr(BT // 16):
vtmem0 = v_tmem_base + k * 8
kdesc0 = kdesc0_base | ((k * 16 * 128) >> 4)
_tcgen05.mma_ts_f16(vk_tmem, vtmem0, kdesc0, vk_idesc, True)
_tcgen05.commit(vk_done_mbar + stage_id)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
num_iters = num_chunks - int(not is_f32)
for _ in range(num_iters):
Waddr = sW[None, None, stage_id].iterator.toint()
wdesc_base = sdesc_template | (Waddr >> 4)
cute.arch.mbarrier_wait(tma_mbar + stage_id, parity)
cute.arch.mbarrier_wait(wh_in_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(K_dim // 64):
for j in cutlass.range_constexpr(64 // 16):
htmem = h_tmem_base + i * 32 + j * 8
wdesc = wdesc_base | ((i * BT * 128 + j * 32) >> 4)
_tcgen05.mma_ts_f16(wh_tmem, htmem, wdesc, wh_idesc, True)
_tcgen05.commit(wh_done_mbar + stage_id)
Kaddr = sK[None, None, stage_id].iterator.toint()
kdesc_base = sdesc_template | (Kaddr >> 4)
cute.arch.mbarrier_wait(vk_in_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for k in cutlass.range_constexpr(BT // 16):
vtmem = v_tmem_base + k * 8
kdesc = kdesc_base | ((k * 16 * 128) >> 4)
_tcgen05.mma_ts_f16(vk_tmem, vtmem, kdesc, vk_idesc, True)
_tcgen05.commit(vk_done_mbar + stage_id)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
elif warp_id >= 4:
# H warps
tid_ = tid % 128
warp_id_ = warp_id % 4
chunk_offset = chunk_offsets[seq_id]
stage_id = 0
vk_stage_id = 0
vk_parity = 0
op = cute.nvgpu.CopyUniversalOp()
cp_16B = cute.make_copy_atom(op, Float32, num_bits_per_copy=128)
##### chunk_id = 0 #####
if True:
chunk_id = 0
end_t = min(bos + (chunk_id + 1) * BT, eos)
last_idx = end_t - 1
# KDA: load per-channel end-of-chunk decay into smem (all 128 k-cols)
s_gl_exp[tid_] = cute.math.exp(
g_cu[last_idx, head_id, tid_], fastmath=True
)
if warp_id_ == 0:
cute.arch.mbarrier_wait(h0_mbar, 0)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
if cutlass.const_expr(is_f32):
for i in cutlass.range_constexpr(K_dim // 32):
h_f32 = cute.make_rmem_tensor(32, Float32)
cute.copy(cp_16B, sH0[tid_, (None, i)], h_f32)
h_bf16 = cute.make_rmem_tensor(32, BFloat16)
h_bf16.store(h_f32.load().to(BFloat16))
_tcgen05.st(
warp_id_ * 32, h_tmem_base + i * 16, "32x32b", 16, h_bf16
)
dst = cute.local_tile(sH[tid_, None], (32,), (i,))
cute.copy(cp_16B, h_bf16, dst)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(wh_in_mbar + stage_id)
# scale H for 2nd MMA -- KDA: per-column decay s_gl_exp[k]
for i in cutlass.range_constexpr(K_dim // 32):
h_f32 = cute.make_rmem_tensor(32, Float32)
if cutlass.const_expr(is_f32):
cute.copy(cp_16B, sH0[tid_, (None, i)], h_f32)
else:
h_bf16 = cute.make_rmem_tensor(32, BFloat16)
sH_src = cute.local_tile(sH0[tid_, None], (32,), (i,))
cute.copy(cp_16B, sH_src, h_bf16)
h_f32.store(
cvt.bf16x2_to_fp32x2(
cute.recast_tensor(h_bf16, Uint32)
).load()
)
for j in cutlass.range_constexpr(32):
h_f32[j] *= s_gl_exp[i * 32 + j]
_tcgen05.st(warp_id_ * 32, vk_tmem + i * 32, "32x32b", 32, h_f32)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(vk_in_mbar + stage_id)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
fence_before_tma_store()
if warp_id_ == 3:
h_src = sH if cutlass.const_expr(is_f32) else sH0
h_dst = tmaH[chunk_offset + chunk_id, head_id, None, None]
simple_tma_copy(H_tma_atom, h_src, h_dst)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
if cutlass.const_expr(not is_f32):
cute.arch.cp_async_bulk_wait_group(0, read=True)
stage_id = (stage_id + 1) % num_stages
##### subsequent chunks #####
for chunk_id in range(1, num_chunks):
end_t = min(bos + (chunk_id + 1) * BT, eos)
last_idx = end_t - 1
# KDA: refresh per-channel end-of-chunk decay for this chunk
s_gl_exp[tid_] = cute.math.exp(
g_cu[last_idx, head_id, tid_], fastmath=True
)
if warp_id_ == 0:
cute.arch.mbarrier_wait(vk_done_mbar + vk_stage_id, vk_parity)
vk_stage_id = (vk_stage_id + 1) % num_stages
if vk_stage_id == 0:
vk_parity ^= 1
elif warp_id_ == 3:
with cute.arch.elect_one():
cute.arch.cp_async_bulk_wait_group(0, read=True)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
for i in cutlass.range_constexpr(K_dim // 32):
h_f32 = _tcgen05.ld(warp_id_ * 32, vk_tmem + i * 32, "32x32b", 32)
h_bf16 = cute.make_rmem_tensor(32, BFloat16)
h_bf16.store(h_f32.to(BFloat16))
_tcgen05.st(
warp_id_ * 32, h_tmem_base + i * 16, "32x32b", 16, h_bf16
)
dst = cute.local_tile(sH[tid_, None], (32,), (i,))
cute.copy(cp_16B, h_bf16, dst)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(wh_in_mbar + stage_id)
# scale H for 2nd MMA -- KDA: per-column decay
for i in cutlass.range_constexpr(K_dim // 32):
h_f32 = cute.make_rmem_tensor(32, Float32)
h_f32.store(
_tcgen05.ld(warp_id_ * 32, vk_tmem + i * 32, "32x32b", 32)
)
for j in cutlass.range_constexpr(32):
h_f32[j] *= s_gl_exp[i * 32 + j]
_tcgen05.st(warp_id_ * 32, vk_tmem + i * 32, "32x32b", 32, h_f32)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(vk_in_mbar + stage_id)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
fence_before_tma_store()
if warp_id_ == 3:
h_dst = tmaH[chunk_offset + chunk_id, head_id, None, None]
simple_tma_copy(H_tma_atom, sH, h_dst)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
stage_id = (stage_id + 1) % num_stages
# handle final state. reuse H0 smem.
if warp_id_ == 0:
cute.arch.mbarrier_wait(vk_done_mbar + vk_stage_id, vk_parity)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
for i in cutlass.range_constexpr(K_dim // 32):
h_f32 = cute.make_rmem_tensor(32, Float32)
h_f32.store(_tcgen05.ld(warp_id_ * 32, vk_tmem + i * 32, "32x32b", 32))
if cutlass.const_expr(is_f32):
cute.copy(cp_16B, h_f32, sH0[tid_, (None, i)])
else:
h_bf16 = cute.make_rmem_tensor(32, BFloat16)
h_bf16.store(h_f32.load().to(BFloat16))
sH0_dst = cute.local_tile(sH0[tid_, None], (32,), (i,))
cute.copy(cp_16B, h_bf16, sH0_dst)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
if warp_id_ == 0:
ht_dst = tmaHT[seq_id, head_id, None, None]
simple_tma_copy(HT_tma_atom, sH0, ht_dst)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
if warp_id_ == 1:
_tcgen05.dealloc()
else:
# V warps -- KDA: v_new is NOT gate-scaled; store RAW to both tmem & gmem.
stage_id = 0
parity = 0
chunk_offset = chunk_offsets[seq_id]
ldsm_trans_op = warp.LdMatrix8x8x16bOp(num_matrices=4, transpose=True)
stsm_trans_op = warp.StMatrix8x8x16bOp(num_matrices=4, transpose=True)
ldsm_trans_atom = cute.make_copy_atom(ldsm_trans_op, BFloat16)
stsm_trans_atom = cute.make_copy_atom(stsm_trans_op, BFloat16)
gV_new_tiles = cute.logical_divide(
tmaV_new[None, head_id, None], (BT, None)
)
sV_view = cute.logical_divide(sV, (None, 8, None))
sV_new_view = cute.logical_divide(sV_new, (None, 8))
s_col = warp_id * 4 + (lane_id // 8)
sV_view = sV_view[None, (None, s_col), None]
sV_new_view = sV_new_view[None, (None, s_col)]
for chunk_id in range(num_chunks):
if warp_id == 0:
cute.arch.mbarrier_wait(tma_mbar + stage_id, parity)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
# unpack U (sV) BF16->FP32, store to tmem to init the 1st MMA acc
for i in cutlass.range_constexpr(BT // 8):
s_row = i * 8 + (lane_id % 8)
v_bf16 = cute.make_rmem_tensor(8, BFloat16)
cute.copy(ldsm_trans_atom, sV_view[s_row, None, stage_id], v_bf16)
v_fp32 = cvt.bf16x2_to_fp32x2(cute.recast_tensor(v_bf16, Uint32))
v_fp32 = cute.logical_divide(v_fp32, 4)
tcol = wh_tmem + i * 8
_tcgen05.st(warp_id * 32 + 0, tcol, "16x256b", 1, v_fp32[None, 0])
_tcgen05.st(warp_id * 32 + 16, tcol, "16x256b", 1, v_fp32[None, 1])
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(wh_in_mbar + stage_id)
# wait for 1st MMA (V_new.T) to finish
if warp_id == 2:
cute.arch.mbarrier_wait(wh_done_mbar + stage_id, parity)
elif warp_id == 3:
with cute.arch.elect_one():
cute.arch.cp_async_bulk_wait_group(0, read=True)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
for i in cutlass.range_constexpr(BT // 8):
v_new = cute.make_rmem_tensor((4, 2), Float32)
tcol = wh_tmem + i * 8
v_new[None, 0].store(
_tcgen05.ld(warp_id * 32 + 0, tcol, "16x256b", 1)
)
v_new[None, 1].store(
_tcgen05.ld(warp_id * 32 + 16, tcol, "16x256b", 1)
)
v_new_bf16 = cute.make_rmem_tensor(8, BFloat16)
v_new_bf16.store(v_new.load().to(BFloat16))
# KDA: NO per-token scaling. v_new (raw) goes to BOTH gmem and tmem.
s_row = i * 8 + (lane_id % 8)
cute.copy(stsm_trans_atom, v_new_bf16, sV_new_view[s_row, None])
v_new_bf16_42 = v_new.load().to(BFloat16).reshape((4, 2))
tcol = v_tmem_base + i * 4
_tcgen05.st(
warp_id * 32 + 0, tcol, "16x128b", 1, v_new_bf16_42[None, 0]
)
_tcgen05.st(
warp_id * 32 + 16, tcol, "16x128b", 1, v_new_bf16_42[None, 1]
)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(vk_in_mbar + stage_id)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
fence_before_tma_store()
if warp_id == 3:
gV = gV_new_tiles[(None, chunk_offset + chunk_id), None]
simple_tma_copy(V_new_tma_atom, sV_new, gV)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
@cache
@staticmethod
def compile(
H: int,
Hv: int,
K_dim: int,
V_dim: int,
h_dtype: cutlass.Numeric = Float32,
BT: int = 64,
num_stages: int = 2,
):
total_t = cute.sym_int()
pad_t = cute.sym_int()
total_chunks_n = cute.sym_int()
num_sequences = cute.sym_int()
cu_entries = cute.sym_int()
K = make_fake_tensor(BFloat16, (total_t, Hv, K_dim), divisibility=16)
V = make_fake_tensor(BFloat16, (pad_t, Hv, V_dim), divisibility=16)
W = make_fake_tensor(BFloat16, (pad_t, Hv, K_dim), divisibility=16)
V_new = make_fake_tensor(BFloat16, (pad_t, Hv, V_dim), divisibility=16)
g_cu = make_fake_tensor(Float32, (total_t, Hv, K_dim), divisibility=4)
h = make_fake_tensor(
BFloat16, (total_chunks_n, Hv, V_dim, K_dim), divisibility=16
)
h0 = make_fake_tensor(
h_dtype, (num_sequences, Hv, V_dim, K_dim), divisibility=16
)
ht = make_fake_tensor(
h_dtype, (num_sequences, Hv, V_dim, K_dim), divisibility=16
)
cu_seqlens = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
chunk_offsets = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
kernel = Sm100KdaChunkHKernel(H, Hv, K_dim, V_dim, h_dtype, BT, num_stages)
stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True)
return cute.compile(
kernel,
K,
V,
W,
V_new,
g_cu,
h,
h0,
ht,
cu_seqlens,
chunk_offsets,
stream,
options="--enable-tvm-ffi",
)
def kda_h_cutedsl(
kg: torch.Tensor,
V: torch.Tensor,
W: torch.Tensor,
V_new: torch.Tensor,
g_cu: torch.Tensor,
h: torch.Tensor,
h0: torch.Tensor,
ht: torch.Tensor,
cu_seqlens: torch.Tensor,
chunk_offsets: torch.Tensor,
BT: int = 64,
num_stages: int = 2,
) -> None:
"""KDA chunk-state kernel. `kg` = per-channel pre-scaled key [T, Hv, K]."""
_, Hv, K_dim = kg.shape
_, _, V_dim = V.shape
h_dtype = {torch.bfloat16: BFloat16, torch.float32: Float32}[h0.dtype]
Sm100KdaChunkHKernel.compile(Hv, Hv, K_dim, V_dim, h_dtype, BT, num_stages)(
kg, V, W, V_new, g_cu, h, h0, ht, cu_seqlens, chunk_offsets
)
@@ -0,0 +1,741 @@
# SPDX-License-Identifier: Apache-2.0
# KDA (Kimi Delta Attention) SM100 KKT-inverse + U/W kernel.
#
# Adapted from gdn_blackwell/kernel_kkt_inv_uw.py. KDA's decay is PER-CHANNEL, so
# (as with kernel_h/o) the gate is folded OUTSIDE this kernel into pre-scaled keys:
#
# kL [c,d] = k[c,d] * exp(g_cu[c,d] - g_cu_last[d]) (KKT left operand)
# kR [j,d] = k[j,d] * exp(g_cu_last[d] - g_cu[j,d]) (KKT right operand, bounded)
# kg [j,d] = k[j,d] * exp(g_cu[j,d]) (W operand, bounded)
#
# Then KKT[c,j] = sum_d kL[c,d]*kR[j,d] = sum_d k[c,d]*k[j,d]*exp(g_cu[c,d]-g_cu[j,d])
# carries the per-channel decay, so:
# A = strictLower(beta * KKT) (NO post-MMA Gamma; decay already inside)
# Ai = inverse(I + A) (Newton-Schulz, gate-independent -> verbatim)
# U = (Ai * beta) @ V
# W = (Ai * beta) @ kg (NO Abg; the exp(g_cu) lives in kg)
#
# Net: this kernel has NO cumsum and NO g_cu — only beta survives, exactly like GDN.
from functools import cache
import cutlass
import torch
from cuda.bindings.driver import CUstream
from cutlass import BFloat16, Float32, Int32, Int64, Uint32, cute
from cutlass.cute.nvgpu import cpasync, warp
from quack.compile_utils import make_fake_tensor
from sglang.srt.layers.attention.cute_utils import (
EVICT_FIRST,
_tcgen05,
cvt,
fence_before_tma_store,
mma_bf16,
simple_tma_copy,
)
class Sm100KdaChunkUWKernel:
"""KDA per-chunk KKT-inverse + U/W (see module docstring)."""
def __init__(
self,
H: int,
Hv: int,
K_dim: int,
V_dim: int,
num_stages: int = 2,
) -> None:
assert Hv % H == 0
assert K_dim == V_dim == 128
self.H = H
self.Hv = Hv
self.K_dim = K_dim
self.V_dim = V_dim
self.num_stages = num_stages
self.BT = 64
self.num_warps = 2 + 4 + 4
@cute.jit
def _make_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
num_stages: int,
op: cpasync.TmaCopyOp,
):
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), num_stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
KL: cute.Tensor, # k*exp(g_cu - g_cu_last) [T, Hv, K]
KR: cute.Tensor, # k*exp(g_cu_last - g_cu) [T, Hv, K]
KG: cute.Tensor, # k*exp(g_cu) [T, Hv, K]
V: cute.Tensor,
U: cute.Tensor,
W: cute.Tensor,
beta: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_indices: cute.Tensor,
total_chunks: cute.Tensor,
num_sms: Int32,
stream: CUstream,
):
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
KL_args = self._make_tma_args(KL, self.K_dim, self.num_stages, tma_g2s)
KR_args = self._make_tma_args(KR, self.K_dim, self.num_stages, tma_g2s)
KG_args = self._make_tma_args(KG, self.K_dim, self.num_stages, tma_g2s)
V_args = self._make_tma_args(V, self.V_dim, self.num_stages, tma_g2s)
U_args = self._make_tma_args(U, self.V_dim, 1, tma_s2g)
W_args = self._make_tma_args(W, self.K_dim, 1, tma_s2g)
grid = (num_sms // self.Hv, self.Hv, 1)
block = (self.num_warps * 32, 1, 1)
self.kernel(
KL_args,
KR_args,
KG_args,
V_args,
U_args,
W_args,
beta,
cu_seqlens,
chunk_indices,
total_chunks,
).launch(grid=grid, block=block, stream=stream)
@cute.kernel
def kernel(
self,
KL_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
KR_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
KG_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
V_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
U_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
W_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
beta: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_indices: cute.Tensor,
total_chunks: cute.Tensor,
):
tid, _, _ = cute.arch.thread_idx()
bid, head_id, _ = cute.arch.block_idx()
grid_x, _, _ = cute.arch.grid_dim()
warp_id = cute.arch.make_warp_uniform(tid // 32)
lane_id = tid % 32
BT = self.BT
K_dim = self.K_dim
V_dim = self.V_dim
num_stages = self.num_stages
KL_tma_atom, tmaKL, sKL_layout = KL_args
KR_tma_atom, tmaKR, sKR_layout = KR_args
KG_tma_atom, tmaKG, sKG_layout = KG_args
V_tma_atom, tmaV, sV_layout = V_args
U_tma_atom, tmaU, sU_layout = U_args
W_tma_atom, tmaW, sW_layout = W_args
def allocate_tensor(smem, dtype, layout):
return smem.allocate_tensor(
dtype, layout.outer, byte_alignment=128, swizzle=layout.inner
)
smem = cutlass.utils.SmemAllocator()
sKL = allocate_tensor(smem, BFloat16, sKL_layout)[None, 0, None, None]
sKR = allocate_tensor(smem, BFloat16, sKR_layout)[None, 0, None, None]
sKG = allocate_tensor(smem, BFloat16, sKG_layout)[None, 0, None, None]
sV = allocate_tensor(smem, BFloat16, sV_layout)[None, 0, None, None]
sU = allocate_tensor(smem, BFloat16, sU_layout)[None, 0, None, 0]
sW = allocate_tensor(smem, BFloat16, sW_layout)[None, 0, None, 0]
swizzle_128B = cute.make_swizzle(3, 4, 3)
sA_layout = cute.make_layout((BT, (64, 1)), stride=(64, (1, BT * 64)))
sA_layout = cute.make_composed_layout(swizzle_128B, 0, sA_layout)
sA = allocate_tensor(smem, BFloat16, sA_layout)
sAi = allocate_tensor(smem, BFloat16, sA_layout)
s_beta = smem.allocate_array(Float32, BT)
tma_mbar = smem.allocate_array(Int64, num_stages)
mma_kkt_mbar = smem.allocate_array(Int64, num_stages)
inv_mbar = smem.allocate_array(Int64, num_stages)
mma_u_mbar = smem.allocate_array(Int64, num_stages)
mma_w_mbar = smem.allocate_array(Int64, num_stages)
epi_mbar = smem.allocate_array(Int64, num_stages)
taddr = smem.allocate(Int32, 4)
kkt_tmem = 0
U_tmem_base = kkt_tmem + BT
Ab_tmem_base = U_tmem_base + V_dim * num_stages
assert Ab_tmem_base + (BT // 2) * num_stages <= 512
ldsm_op = warp.LdMatrix8x8x16bOp(num_matrices=4)
stsm_op = warp.StMatrix8x8x16bOp(num_matrices=4)
ldsm_trans_op = warp.LdMatrix8x8x16bOp(num_matrices=4, transpose=True)
ldsm_atom = cute.make_copy_atom(ldsm_op, BFloat16)
stsm_atom = cute.make_copy_atom(stsm_op, BFloat16)
ldsm_trans_atom = cute.make_copy_atom(ldsm_trans_op, BFloat16)
if warp_id == 0:
with cute.arch.elect_one():
for i in cutlass.range_constexpr(num_stages):
cute.arch.mbarrier_init(tma_mbar + i, 1)
cute.arch.mbarrier_init(mma_kkt_mbar + i, 1)
cute.arch.mbarrier_init(inv_mbar + i, 128)
cute.arch.mbarrier_init(mma_u_mbar + i, 1)
cute.arch.mbarrier_init(mma_w_mbar + i, 1)
cute.arch.mbarrier_init(epi_mbar + i, 128)
cute.arch.mbarrier_init_fence()
elif warp_id == 1:
cpasync.prefetch_descriptor(KL_tma_atom)
cpasync.prefetch_descriptor(KR_tma_atom)
cpasync.prefetch_descriptor(KG_tma_atom)
cpasync.prefetch_descriptor(V_tma_atom)
cpasync.prefetch_descriptor(U_tma_atom)
cpasync.prefetch_descriptor(W_tma_atom)
cute.arch.sync_threads()
num_global_chunks = total_chunks[0]
if warp_id == 9:
# TMA warp
stage_id = 0
parity = 1
for global_chunk_id in range(bid, num_global_chunks, grid_x):
seq_id = chunk_indices[global_chunk_id, 0]
chunk_id = chunk_indices[global_chunk_id, 1]
bos = cu_seqlens[seq_id]
mbar = tma_mbar + stage_id
# KDA: all keys are per v-head [T, Hv, K], index by head_id.
gKL = cute.local_tile(
cute.domain_offset((bos, 0), tmaKL[None, head_id, None]),
tiler=(BT, K_dim),
coord=(chunk_id, 0),
)
gKR = cute.local_tile(
cute.domain_offset((bos, 0), tmaKR[None, head_id, None]),
tiler=(BT, K_dim),
coord=(chunk_id, 0),
)
gKG = cute.local_tile(
cute.domain_offset((bos, 0), tmaKG[None, head_id, None]),
tiler=(BT, K_dim),
coord=(chunk_id, 0),
)
gV = cute.local_tile(
cute.domain_offset((bos, 0), tmaV[None, head_id, None]),
tiler=(BT, V_dim),
coord=(chunk_id, 0),
)
cute.arch.mbarrier_wait(mma_u_mbar + stage_id, parity)
with cute.arch.elect_one():
STAGE_SIZE = BT * (K_dim + K_dim + K_dim + V_dim) * 2
cute.arch.mbarrier_arrive_and_expect_tx(mbar, STAGE_SIZE)
simple_tma_copy(KL_tma_atom, gKL, sKL[None, None, stage_id], mbar)
simple_tma_copy(KR_tma_atom, gKR, sKR[None, None, stage_id], mbar)
simple_tma_copy(
KG_tma_atom, gKG, sKG[None, None, stage_id], mbar, EVICT_FIRST
)
simple_tma_copy(
V_tma_atom, gV, sV[None, None, stage_id], mbar, EVICT_FIRST
)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
elif warp_id == 8:
# MMA warp
_tcgen05.alloc(taddr)
stage_id = 0
parity = 0
kkt_idesc = _tcgen05.make_bf16_idesc(BT, BT)
u_idesc = _tcgen05.make_bf16_idesc(BT, V_dim, transpose_B=True)
w_idesc = _tcgen05.make_bf16_idesc(BT, K_dim, transpose_B=True)
sdesc_template = _tcgen05.make_sdesc_128B_swizzle(BT * 128)
for global_chunk_id in range(bid, num_global_chunks, grid_x):
U_tmem = U_tmem_base + V_dim * stage_id
W_tmem = U_tmem | (16 << 16)
Ab_tmem = Ab_tmem_base + (BT // 2) * stage_id
Abg_tmem = Ab_tmem | (16 << 16)
##### KKT MMA: KKT = kL @ kR.T #####
klraddr = sKL[None, None, stage_id].iterator.toint()
krraddr = sKR[None, None, stage_id].iterator.toint()
kldesc_base = sdesc_template | (klraddr >> 4)
krdesc_base = sdesc_template | (krraddr >> 4)
cute.arch.mbarrier_wait(tma_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(K_dim // 64):
for j in cutlass.range_constexpr(64 // 16):
off = (i * BT * 128 + j * 32) >> 4
_tcgen05.mma_f16(
kkt_tmem,
kldesc_base | off,
krdesc_base | off,
kkt_idesc,
(i > 0) or (j > 0),
)
_tcgen05.commit(mma_kkt_mbar + stage_id)
##### U/W MMA: U = Ab @ V, W = Ab @ kg #####
vaddr = sV[None, None, stage_id].iterator.toint()
kgaddr = sKG[None, None, stage_id].iterator.toint()
vdesc = sdesc_template | (vaddr >> 4)
kgdesc = sdesc_template | (kgaddr >> 4)
cute.arch.mbarrier_wait(epi_mbar + stage_id, parity ^ 1)
cute.arch.mbarrier_wait(inv_mbar + stage_id, parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(BT // 16):
_tcgen05.mma_ts_f16(
W_tmem, Abg_tmem + i * 8, kgdesc, w_idesc, i > 0
)
kgdesc += (16 * 128) >> 4
_tcgen05.commit(mma_w_mbar + stage_id)
for i in cutlass.range_constexpr(BT // 16):
_tcgen05.mma_ts_f16(
U_tmem, Ab_tmem + i * 8, vdesc, u_idesc, i > 0
)
vdesc += (16 * 128) >> 4
_tcgen05.commit(mma_u_mbar + stage_id)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
cute.arch.mbarrier_wait(epi_mbar + stage_id, parity ^ 1)
_tcgen05.dealloc()
elif warp_id >= 4:
# inv warps
tid_ = tid % 128
warp_id_ = warp_id % 4
stage_id = 0
parity = 0
sA_ldsm = cute.logical_divide(sA, (16, cute.make_layout((8, 2))))
sAi_ldsm = cute.logical_divide(sAi, (16, cute.make_layout((8, 2))))
sA_ldsm = sA_ldsm[(lane_id % 16, None), ((None, lane_id // 16), None)]
sAi_ldsm = sAi_ldsm[(lane_id % 16, None), ((None, lane_id // 16), None)]
for i in cutlass.range_constexpr((BT // 4 * 3) * BT // 128):
idx = i * 128 + tid_
sAi[idx // BT, idx % BT] = BFloat16(0.0)
row_indices = cute.make_rmem_tensor((1, 2, 1), Int32)
row_indices[0, 0, 0] = warp_id_ * 16 + (lane_id // 4)
row_indices[0, 1, 0] = warp_id_ * 16 + (lane_id // 4) + 8
row_indices = row_indices.load()
col_indices = cute.make_rmem_tensor((2, 1, 2), Int32)
col_indices[0, 0, 0] = (lane_id % 4) * 2 + 0
col_indices[1, 0, 0] = (lane_id % 4) * 2 + 1
col_indices[0, 0, 1] = (lane_id % 4) * 2 + 8
col_indices[1, 0, 1] = (lane_id % 4) * 2 + 9
col_indices = col_indices.load()
for global_chunk_id in range(bid, num_global_chunks, grid_x):
seq_id = chunk_indices[global_chunk_id, 0]
chunk_id = chunk_indices[global_chunk_id, 1]
bos = cu_seqlens[seq_id]
eos = cu_seqlens[seq_id + 1]
off_t = bos + chunk_id * BT
t = off_t + tid_
##### Phase 1: load beta (KDA: no cumsum) #####
if tid_ < BT:
in_bounds = t < eos
beta_val = beta[t, head_id] if in_bounds else Float32(0.0)
s_beta[tid_] = beta_val
##### Phase 2: A = strictLower(beta * kkt) #####
if warp_id_ == 0:
cute.arch.mbarrier_wait(mma_kkt_mbar + stage_id, parity)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
row_coord = (lane_id // 4, None, warp_id_)
s_beta_view = cute.make_tensor(s_beta, (8, 2, 4))
beta_row = s_beta_view[row_coord].load().reshape((1, 2, 1))
kkt = _tcgen05.ld(kkt_tmem, 0, "16x256b", BT // 8)
kkt = kkt.reshape((2, 2, 2, BT // 16))
for i in cutlass.range_constexpr(BT // 16):
# KDA: decay is already inside KKT; only beta + mask here.
A = kkt[None, None, None, i] * beta_row
A_masked = cute.where(row_indices > col_indices + i * 16, A, 0.0)
packed = cute.make_rmem_tensor(4, Uint32)
packed[0] = cvt.fp32x2_to_bf16x2(
A_masked[0, 0, 0], A_masked[1, 0, 0]
)
packed[1] = cvt.fp32x2_to_bf16x2(
A_masked[0, 1, 0], A_masked[1, 1, 0]
)
packed[2] = cvt.fp32x2_to_bf16x2(
A_masked[0, 0, 1], A_masked[1, 0, 1]
)
packed[3] = cvt.fp32x2_to_bf16x2(
A_masked[0, 1, 1], A_masked[1, 1, 1]
)
cute.copy(
stsm_atom,
cute.recast_tensor(packed, BFloat16),
sA_ldsm[warp_id_, None, i],
)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
##### Phase 3: matrix inverse (VERBATIM from GDN) #####
zeros_f32 = cute.make_rmem_tensor(4, Float32)
zeros_f32.fill(0.0)
def set_diagonal(A: cute.Tensor, lane_id: Int32):
"Set the diagonal to 1s"
if lane_id % 9 == 0:
A[0] = (A[0] & Uint32(0xFFFF0000)) | Uint32(0x00003F80)
A[3] = (A[3] & Uint32(0xFFFF0000)) | Uint32(0x00003F80)
elif lane_id % 9 == 4:
A[0] = (A[0] & Uint32(0x0000FFFF)) | Uint32(0x3F800000)
A[3] = (A[3] & Uint32(0x0000FFFF)) | Uint32(0x3F800000)
Ai_bf16 = cute.make_rmem_tensor(8, BFloat16)
mma_B_bf16 = cute.make_rmem_tensor(8, BFloat16)
M_bf16 = cute.make_rmem_tensor(8, BFloat16)
acc = cute.make_rmem_tensor((4, 2), Float32)
Ai = cute.recast_tensor(Ai_bf16, Uint32)
mma_B = cute.logical_divide(cute.recast_tensor(mma_B_bf16, Uint32), 2)
M = cute.logical_divide(cute.recast_tensor(M_bf16, Uint32), 2)
cute.copy(ldsm_atom, sA_ldsm[warp_id_, None, warp_id_], Ai_bf16)
for i in cutlass.range_constexpr(4):
Ai[i] ^= Uint32(0x80008000)
set_diagonal(Ai, lane_id)
Ai_f32 = cute.logical_divide(cvt.bf16x2_to_fp32x2(Ai), 4)
cute.copy(ldsm_trans_atom, sA_ldsm[warp_id_, None, warp_id_], M_bf16)
set_diagonal(M, lane_id)
for i in cutlass.range_constexpr(4):
M[i] ^= Uint32(0x80008000)
for _ in cutlass.range_constexpr(3):
cute.copy(stsm_atom, Ai_bf16, sA_ldsm[warp_id_, None, warp_id_])
cute.arch.sync_warp()
acc[None, 0] = mma_bf16(Ai, M[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, M[None, 1], zeros_f32)
Ai_bf16.store(acc.load().to(BFloat16))
for j in cutlass.range_constexpr(8):
Ai_f32[j] *= 2.0
cute.copy(
ldsm_trans_atom,
sA_ldsm[warp_id_, None, warp_id_],
mma_B_bf16,
)
Ai_f32[None, 0] = mma_bf16(Ai, mma_B[None, 0], Ai_f32[None, 0])
Ai_f32[None, 1] = mma_bf16(Ai, mma_B[None, 1], Ai_f32[None, 1])
Ai_bf16.store(Ai_f32.load().to(BFloat16))
cute.copy(stsm_atom, Ai_bf16, sAi_ldsm[warp_id_, None, warp_id_])
cute.arch.barrier(barrier_id=1, number_of_threads=128)
if warp_id_ > 0:
neg_Ai = cute.make_rmem_tensor(4, Uint32)
for i in cutlass.range_constexpr(4):
neg_Ai[i] = Ai[i] ^ Uint32(0x80008000)
cute.copy(
ldsm_trans_atom,
sA_ldsm[warp_id_, None, warp_id_ - 1],
mma_B_bf16,
)
acc[None, 0] = mma_bf16(neg_Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(neg_Ai, mma_B[None, 1], zeros_f32)
Ai_bf16.store(acc.load().to(BFloat16))
cute.copy(
ldsm_trans_atom,
sAi_ldsm[warp_id_ - 1, None, warp_id_ - 1],
mma_B_bf16,
)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], zeros_f32)
Ai_bf16.store(acc.load().to(BFloat16))
cute.copy(
stsm_atom,
Ai_bf16,
sAi_ldsm[warp_id_, None, warp_id_ - 1],
)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
if warp_id_ < 2:
cute.copy(
ldsm_atom,
sA_ldsm[warp_id_ + 2, None, warp_id_],
Ai_bf16,
)
cute.copy(
ldsm_trans_atom,
sAi_ldsm[warp_id_, None, warp_id_],
mma_B_bf16,
)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], zeros_f32)
cute.copy(
ldsm_atom,
sA_ldsm[warp_id_ + 2, None, warp_id_ + 1],
Ai_bf16,
)
cute.copy(
ldsm_trans_atom,
sAi_ldsm[warp_id_ + 1, None, warp_id_],
mma_B_bf16,
)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], acc[None, 0])
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], acc[None, 1])
tmp = cute.make_rmem_tensor(8, BFloat16)
tmp.store(acc.load().to(BFloat16))
cute.copy(stsm_atom, tmp, sAi_ldsm[warp_id_ + 2, None, warp_id_])
cute.arch.sync_warp()
cute.copy(
ldsm_atom, sAi_ldsm[warp_id_ + 2, None, warp_id_ + 2], Ai_bf16
)
for i in cutlass.range_constexpr(4):
Ai[i] ^= Uint32(0x80008000)
cute.copy(
ldsm_trans_atom,
sAi_ldsm[warp_id_ + 2, None, warp_id_],
mma_B_bf16,
)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], zeros_f32)
tmp.store(acc.load().to(BFloat16))
cute.copy(stsm_atom, tmp, sAi_ldsm[warp_id_ + 2, None, warp_id_])
cute.arch.barrier(barrier_id=1, number_of_threads=128)
if warp_id_ == 0:
cute.copy(ldsm_atom, sA_ldsm[3, None, 0], Ai_bf16)
cute.copy(ldsm_trans_atom, sAi_ldsm[0, None, 0], mma_B_bf16)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], zeros_f32)
for i in cutlass.range_constexpr(1, 3):
cute.copy(ldsm_atom, sA_ldsm[3, None, i], Ai_bf16)
cute.copy(ldsm_trans_atom, sAi_ldsm[i, None, 0], mma_B_bf16)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], acc[None, 0])
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], acc[None, 1])
tmp = cute.make_rmem_tensor(8, BFloat16)
tmp.store(acc.load().to(BFloat16))
cute.copy(stsm_atom, tmp, sAi_ldsm[3, None, 0])
cute.arch.sync_warp()
cute.copy(ldsm_atom, sAi_ldsm[3, None, 3], Ai_bf16)
for i in cutlass.range_constexpr(4):
Ai[i] ^= Uint32(0x80008000)
cute.copy(ldsm_trans_atom, sAi_ldsm[3, None, 0], mma_B_bf16)
acc[None, 0] = mma_bf16(Ai, mma_B[None, 0], zeros_f32)
acc[None, 1] = mma_bf16(Ai, mma_B[None, 1], zeros_f32)
tmp.store(acc.load().to(BFloat16))
cute.copy(stsm_atom, tmp, sAi_ldsm[3, None, 0])
##### Phase 4: Ab = Ai * beta (KDA: no Abg) #####
if warp_id_ == 3:
cute.arch.mbarrier_wait(mma_u_mbar + stage_id, parity ^ 1)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
for i in cutlass.range_constexpr(BT // 16):
cute.copy(ldsm_atom, sAi_ldsm[warp_id_, None, i], Ai_bf16)
col_coord = (None, lane_id % 4, None, i)
s_beta_view = cute.make_tensor(s_beta, (2, 4, 2, BT // 16))
beta_col = s_beta_view[col_coord].load().reshape((2, 1, 2))
Ai_f32 = cvt.bf16x2_to_fp32x2(Ai).load().reshape((2, 2, 2))
Ab_f32 = Ai_f32 * beta_col
Ab = Ab_f32.to(BFloat16)
Ab_tmem = Ab_tmem_base + (BT // 2) * stage_id + i * 8
_tcgen05.st(warp_id_ * 32, Ab_tmem, "16x128b", 2, Ab)
# KDA: Abg == Ab (no per-chunk g on the matrix). Duplicate into the
# +16 lane region so the W MMA (reads Abg_tmem) sees valid data,
# matching GDN's tmem layout exactly.
_tcgen05.st(warp_id_ * 32 + 16, Ab_tmem, "16x128b", 2, Ab)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(inv_mbar + stage_id)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
elif warp_id < 4:
# epi warps (store U, W) -- VERBATIM from GDN
stage_id = 0
parity = 0
gU_tiles = cute.logical_divide(tmaU[None, head_id, None], (BT, None))
gW_tiles = cute.logical_divide(tmaW[None, head_id, None], (BT, None))
s_row = warp_id * 16 + lane_id % 16
sW_view = cute.zipped_divide(
sW[s_row, None],
tiler=cute.make_layout((8, 2)),
)
sU_view = cute.zipped_divide(
sU[s_row, None],
tiler=cute.make_layout((8, 2)),
)
sW_view = sW_view[(None, lane_id // 16), None]
sU_view = sU_view[(None, lane_id // 16), None]
for global_chunk_id in range(bid, num_global_chunks, grid_x):
U_tmem = U_tmem_base + V_dim * stage_id
if warp_id == 0:
cute.arch.mbarrier_wait(mma_w_mbar + stage_id, parity)
elif warp_id == 1:
with cute.arch.elect_one():
cute.arch.cp_async_bulk_wait_group(0, read=True)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
w_f32 = _tcgen05.ld(warp_id * 32 + 16, U_tmem, "16x256b", K_dim // 8)
_tcgen05.wait_ld()
w_bf16 = cute.make_rmem_tensor((8, K_dim // 16), BFloat16)
w_bf16.store(w_f32.to(BFloat16))
cute.copy(stsm_atom, w_bf16, sW_view)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
fence_before_tma_store()
if warp_id == 0:
cute.arch.mbarrier_wait(mma_u_mbar + stage_id, parity)
elif warp_id == 1:
simple_tma_copy(
W_tma_atom, sW, gW_tiles[(None, global_chunk_id), None]
)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
u_f32 = _tcgen05.ld(warp_id * 32, U_tmem, "16x256b", V_dim // 8)
_tcgen05.wait_ld()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(epi_mbar + stage_id)
u_bf16 = cute.make_rmem_tensor((8, V_dim // 16), BFloat16)
u_bf16.store(u_f32.to(BFloat16))
cute.copy(stsm_atom, u_bf16, sU_view)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
fence_before_tma_store()
if warp_id == 1:
simple_tma_copy(
U_tma_atom, sU, gU_tiles[(None, global_chunk_id), None]
)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
@cache
@staticmethod
def compile(H: int, Hv: int, K_dim: int, V_dim: int, num_stages: int = 2):
total_t = cute.sym_int()
pad_t = cute.sym_int()
total_chunks_n = cute.sym_int()
num_sequences = cute.sym_int()
KL = make_fake_tensor(BFloat16, (total_t, Hv, K_dim), divisibility=16)
KR = make_fake_tensor(BFloat16, (total_t, Hv, K_dim), divisibility=16)
KG = make_fake_tensor(BFloat16, (total_t, Hv, K_dim), divisibility=16)
V = make_fake_tensor(BFloat16, (total_t, Hv, V_dim), divisibility=16)
U = make_fake_tensor(BFloat16, (pad_t, Hv, V_dim), divisibility=16)
W = make_fake_tensor(BFloat16, (pad_t, Hv, K_dim), divisibility=16)
beta = make_fake_tensor(Float32, (total_t, Hv), divisibility=4)
cu_seqlens = make_fake_tensor(Int32, (num_sequences,), divisibility=1)
chunk_indices = make_fake_tensor(Int32, (total_chunks_n, 2), divisibility=2)
total_chunks = make_fake_tensor(Int32, (1,), divisibility=1)
kernel = Sm100KdaChunkUWKernel(H, Hv, K_dim, V_dim, num_stages)
stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True)
return cute.compile(
kernel,
KL,
KR,
KG,
V,
U,
W,
beta,
cu_seqlens,
chunk_indices,
total_chunks,
Int32(148),
stream,
options="--enable-tvm-ffi",
)
def kkt_inv_uw_cutedsl(
KL: torch.Tensor,
KR: torch.Tensor,
KG: torch.Tensor,
V: torch.Tensor,
U: torch.Tensor,
W: torch.Tensor,
beta: torch.Tensor,
cu_seqlens: torch.Tensor,
chunk_indices: torch.Tensor,
total_chunks: torch.Tensor,
num_sms: int = 148,
) -> None:
"""KDA KKT-inverse + U/W. KL/KR/KG are the pre-scaled keys (see module doc)."""
_, Hv, K_dim = KL.shape
_, _, V_dim = V.shape
Sm100KdaChunkUWKernel.compile(Hv, Hv, K_dim, V_dim)(
KL, KR, KG, V, U, W, beta, cu_seqlens, chunk_indices, total_chunks, num_sms
)
@@ -0,0 +1,584 @@
# SPDX-License-Identifier: Apache-2.0
# KDA (Kimi Delta Attention) SM100 output kernel.
#
# Adapted from gdn_blackwell/kernel_o.py. KDA's decay is PER-CHANNEL, so the
# decay cannot be applied as a post-MMA scalar Gamma. Instead all gate + scale
# factors are folded OUTSIDE this kernel into three pre-scaled tensors:
#
# qg [c,d] = scale * q[c,d] * exp(g_cu[c,d]) -> Q @ H.T term
# qg2[c,d] = scale * q[c,d] * exp(g_cu[c,d] - g_cu_last[d]) -> Aqk Q operand
# kg [j,d] = k[j,d] * exp(g_cu_last[d] - g_cu[j,d]) -> Aqk K operand
# (== kernel_h's kg, bounded <=|k|)
#
# Then:
# Aqk = strictLowerIncl(qg2 @ kg.T) (masking warp: causal mask only, NO Gamma)
# QH = qg @ H.T (scale + exp(g_cu) already baked)
# O = QH + Aqk @ v_new (epilogue: NO scale, NO exp(g_cu))
#
# Net effect: g_cu is NOT needed inside this kernel at all.
from functools import cache
import cutlass
import torch
from cuda.bindings.driver import CUstream
from cutlass import BFloat16, Int32, Int64, Uint32, cute
from cutlass.cute.nvgpu import cpasync, warp
from quack.compile_utils import make_fake_tensor
from sglang.srt.layers.attention.cute_utils import (
EVICT_FIRST,
_tcgen05,
cvt,
fence_before_tma_store,
simple_tma_copy,
)
class Sm100KdaChunkOKernel:
"""KDA per-token output (see module docstring)."""
def __init__(
self,
H: int,
Hv: int,
K_dim: int,
V_dim: int,
BT: int = 64,
num_stages: int = 2,
) -> None:
assert Hv % H == 0
assert K_dim == 128
assert V_dim == 128
assert BT == 64
self.H = H
self.Hv = Hv
self.K_dim = K_dim
self.V_dim = V_dim
self.BT = BT
self.num_stages = num_stages
self.num_warps = 10
@cute.jit
def _make_bf16_tma_args(
self,
tensor: cute.Tensor,
dim: cutlass.Constexpr[int],
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(self.BT, 1, (64, dim // 64), stages),
stride=(64, 0, (1, self.BT * 64), self.BT * dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, 64)),
slayout,
cta_tiler=(self.BT, 1, dim),
)
return atom, tma_tensor, slayout
@cute.jit
def _make_h_tma_args(
self,
tensor: cute.Tensor,
op: cpasync.TmaCopyOp,
stages: cutlass.Constexpr[int],
):
num_elems = 128 // (tensor.element_type.width // 8)
swizzle_128B = cute.make_swizzle(3, 4, 3)
slayout = cute.make_layout(
(1, self.V_dim, (num_elems, self.K_dim // num_elems), stages),
stride=(0, num_elems, (1, self.V_dim * num_elems), self.V_dim * self.K_dim),
)
slayout = cute.make_composed_layout(swizzle_128B, 0, slayout)
atom, tma_tensor = cpasync.make_tiled_tma_atom(
op,
cute.logical_divide(tensor, (None, None, num_elems)),
slayout,
cta_tiler=(1, self.V_dim, self.K_dim),
)
return atom, tma_tensor, slayout
@cute.jit
def __call__(
self,
qg: cute.Tensor, # scale*q*exp(g_cu) [T, Hv, K]
qg2: cute.Tensor, # scale*q*exp(g_cu-g_cu_last) [T, Hv, K]
kg: cute.Tensor, # k*exp(g_cu_last-g_cu) [T, Hv, K]
v_new_chunks: cute.Tensor,
h: cute.Tensor,
o: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_indices: cute.Tensor,
total_chunks: cute.Tensor,
num_sms: Int32,
stream: CUstream,
):
grid = (num_sms // self.Hv, self.Hv, 1)
block = (self.num_warps * 32, 1, 1)
tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
Q_args = self._make_bf16_tma_args(qg2, self.K_dim, tma_g2s, self.num_stages)
Q2_args = self._make_bf16_tma_args(qg, self.K_dim, tma_g2s, self.num_stages)
K_args = self._make_bf16_tma_args(kg, self.K_dim, tma_g2s, self.num_stages)
V_args = self._make_bf16_tma_args(
v_new_chunks, self.V_dim, tma_g2s, self.num_stages
)
H_args = self._make_h_tma_args(h, tma_g2s, self.num_stages)
O_args = self._make_bf16_tma_args(o, self.V_dim, tma_s2g, 1)
self.kernel(
Q_args,
Q2_args,
K_args,
V_args,
H_args,
O_args,
o,
cu_seqlens,
chunk_indices,
total_chunks,
).launch(grid=grid, block=block, stream=stream)
@cute.kernel
def kernel(
self,
Q_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
Q2_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
K_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
V_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
H_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
O_args: tuple[cute.CopyAtom, cute.Tensor, cute.ComposedLayout],
o: cute.Tensor,
cu_seqlens: cute.Tensor,
chunk_indices: cute.Tensor,
total_chunks: cute.Tensor,
):
tid, _, _ = cute.arch.thread_idx()
bid, v_head_id, _ = cute.arch.block_idx()
grid_x, _, _ = cute.arch.grid_dim()
warp_id = cute.arch.make_warp_uniform(tid // 32)
lane_id = tid % 32
BT = self.BT
K_dim = self.K_dim
V_dim = self.V_dim
num_stages = self.num_stages
num_global_chunks = total_chunks[0]
Q_tma_atom, tmaQ, sQ_layout = Q_args
Q2_tma_atom, tmaQ2, sQ2_layout = Q2_args
K_tma_atom, tmaK, sK_layout = K_args
V_tma_atom, tmaV, sV_layout = V_args
H_tma_atom, tmaH, sH_layout = H_args
O_tma_atom, tmaO, sO_layout = O_args
def allocate_tensor(smem, dtype, layout):
return smem.allocate_tensor(
dtype, layout.outer, byte_alignment=128, swizzle=layout.inner
)
smem = cutlass.utils.SmemAllocator()
sQ = allocate_tensor(smem, BFloat16, sQ_layout)[None, 0, None, None]
sQ2 = allocate_tensor(smem, BFloat16, sQ2_layout)[None, 0, None, None]
sK = allocate_tensor(smem, BFloat16, sK_layout)[None, 0, None, None]
sV = allocate_tensor(smem, BFloat16, sV_layout)[None, 0, None, None]
sH = allocate_tensor(smem, BFloat16, sH_layout)[0, None, None, None]
sO = allocate_tensor(smem, BFloat16, sO_layout)[None, 0, None, 0]
qk_full_mbar = smem.allocate_array(Int64, num_stages)
hv_full_mbar = smem.allocate_array(Int64, num_stages)
qk_empty_mbar = smem.allocate_array(Int64, num_stages)
pv_mma_mbar = smem.allocate_array(Int64, num_stages)
qk_mbar = smem.allocate_array(Int64, 1)
mask_mbar = smem.allocate_array(Int64, 1)
epi_mbar = smem.allocate_array(Int64, 1)
taddr = smem.allocate(Int32, 4)
qk_tmem = 0
p_tmem = 64
out_tmem = 128
qh_tmem = 256
if warp_id == 0:
with cute.arch.elect_one():
for i in cutlass.range_constexpr(num_stages):
cute.arch.mbarrier_init(qk_full_mbar + i, 1)
cute.arch.mbarrier_init(qk_empty_mbar + i, 1)
cute.arch.mbarrier_init(hv_full_mbar + i, 1)
cute.arch.mbarrier_init(pv_mma_mbar + i, 1)
cute.arch.mbarrier_init(qk_mbar, 1)
cute.arch.mbarrier_init(mask_mbar, 128)
cute.arch.mbarrier_init(epi_mbar, 128)
cute.arch.mbarrier_init_fence()
elif warp_id == 9:
cpasync.prefetch_descriptor(Q_tma_atom)
cpasync.prefetch_descriptor(Q2_tma_atom)
cpasync.prefetch_descriptor(K_tma_atom)
cpasync.prefetch_descriptor(V_tma_atom)
cpasync.prefetch_descriptor(H_tma_atom)
cute.arch.sync_threads()
if warp_id == 9:
# TMA warp
stage_id = 0
parity = 1
for global_chunk_id in range(bid, num_global_chunks, grid_x):
seq_id = chunk_indices[global_chunk_id, 0]
chunk_id = chunk_indices[global_chunk_id, 1]
bos = cu_seqlens[seq_id]
# copy qg2 (Q for Aqk), qg (Q for QH), kg (K for Aqk).
# KDA: per v-head tensors, index by v_head_id.
q_tile = cute.local_tile(
cute.domain_offset((bos, 0), tmaQ[None, v_head_id, None]),
tiler=(BT, K_dim),
coord=(chunk_id, 0),
)
q2_tile = cute.local_tile(
cute.domain_offset((bos, 0), tmaQ2[None, v_head_id, None]),
tiler=(BT, K_dim),
coord=(chunk_id, 0),
)
k_tile = cute.local_tile(
cute.domain_offset((bos, 0), tmaK[None, v_head_id, None]),
tiler=(BT, K_dim),
coord=(chunk_id, 0),
)
mbar = qk_full_mbar + stage_id
cute.arch.mbarrier_wait(qk_empty_mbar + stage_id, parity)
with cute.arch.elect_one():
STAGE_SIZE = BT * (K_dim + K_dim + K_dim) * 2
cute.arch.mbarrier_arrive_and_expect_tx(mbar, STAGE_SIZE)
simple_tma_copy(Q_tma_atom, q_tile, sQ[None, None, stage_id], mbar)
simple_tma_copy(Q2_tma_atom, q2_tile, sQ2[None, None, stage_id], mbar)
simple_tma_copy(K_tma_atom, k_tile, sK[None, None, stage_id], mbar)
# copy H and V
gH = tmaH[global_chunk_id * self.Hv + v_head_id, None, None]
gV = cute.local_tile(
tmaV[None, v_head_id, None],
tiler=(BT, V_dim),
coord=(global_chunk_id, 0),
)
mbar = hv_full_mbar + stage_id
cute.arch.mbarrier_wait(pv_mma_mbar + stage_id, parity)
with cute.arch.elect_one():
H_STAGE_SIZE = V_dim * K_dim * 2
V_STAGE_SIZE = BT * V_dim * 2
cute.arch.mbarrier_arrive_and_expect_tx(
mbar, H_STAGE_SIZE + V_STAGE_SIZE
)
simple_tma_copy(
H_tma_atom, gH, sH[None, None, stage_id], mbar, EVICT_FIRST
)
simple_tma_copy(
V_tma_atom, gV, sV[None, None, stage_id], mbar, EVICT_FIRST
)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
parity ^= 1
elif warp_id == 8:
# MMA warp
_tcgen05.alloc(taddr)
sdesc_template = _tcgen05.make_sdesc_128B_swizzle(BT * 128)
qk_idesc = _tcgen05.make_bf16_idesc(BT, BT)
qh_idesc = _tcgen05.make_bf16_idesc(BT, V_dim)
pv_idesc = _tcgen05.make_bf16_idesc(BT, V_dim, transpose_B=True)
stage_id = 0
tma_parity = 0
mask_parity = 0
for global_chunk_id in range(bid, num_global_chunks, grid_x):
qaddr = sQ[None, None, stage_id].iterator.toint()
q2addr = sQ2[None, None, stage_id].iterator.toint()
kaddr = sK[None, None, stage_id].iterator.toint()
haddr = sH[None, None, stage_id].iterator.toint()
vaddr = sV[None, None, stage_id].iterator.toint()
qdesc_base = sdesc_template | (qaddr >> 4)
q2desc_base = sdesc_template | (q2addr >> 4)
kdesc_base = sdesc_template | (kaddr >> 4)
hdesc_base = sdesc_template | (haddr >> 4)
vdesc_base = sdesc_template | (vaddr >> 4)
##### 1st MMA: Aqk = qg2 @ kg.T #####
cute.arch.mbarrier_wait(epi_mbar, mask_parity ^ 1)
cute.arch.mbarrier_wait(qk_full_mbar + stage_id, tma_parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(K_dim // BT):
for j in cutlass.range_constexpr(BT // 16):
qdesc = qdesc_base | ((i * BT * 128 + j * 32) >> 4)
kdesc = kdesc_base | ((i * BT * 128 + j * 32) >> 4)
_tcgen05.mma_f16(
qk_tmem, qdesc, kdesc, qk_idesc, (i > 0) or (j > 0)
)
_tcgen05.commit(qk_mbar)
##### 2nd MMA: QH = qg @ H.T #####
cute.arch.mbarrier_wait(hv_full_mbar + stage_id, tma_parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(K_dim // BT):
for j in cutlass.range_constexpr(BT // 16):
q2desc = q2desc_base | ((i * BT * 128 + j * 32) >> 4)
hdesc = hdesc_base | ((i * V_dim * 128 + j * 32) >> 4)
_tcgen05.mma_f16(
qh_tmem, q2desc, hdesc, qh_idesc, (i > 0) or (j > 0)
)
_tcgen05.commit(qk_empty_mbar + stage_id)
##### 3rd MMA: P @ V #####
cute.arch.mbarrier_wait(mask_mbar, mask_parity)
_tcgen05.fence_after_thread_sync()
with cute.arch.elect_one():
for i in cutlass.range_constexpr(BT // 16):
vdesc = vdesc_base | ((i * 16 * 128) >> 4)
_tcgen05.mma_ts_f16(
out_tmem, p_tmem + i * 8, vdesc, pv_idesc, i > 0
)
_tcgen05.commit(pv_mma_mbar + stage_id)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
tma_parity ^= 1
mask_parity ^= 1
cute.arch.mbarrier_wait(epi_mbar, mask_parity ^ 1)
_tcgen05.dealloc()
elif warp_id >= 4:
# masking warps -- KDA: causal mask only, decay is baked into operands.
warp_id_ = warp_id % 4
parity = 0
row_indices = cute.make_rmem_tensor(2, Int32)
row_indices[0] = warp_id_ * 16 + lane_id // 4
row_indices[1] = warp_id_ * 16 + lane_id // 4 + 8
row_indices = row_indices.load().reshape((1, 2))
col_indices = cute.make_rmem_tensor(2, Int32)
col_indices[0] = (lane_id % 4) * 2
col_indices[1] = (lane_id % 4) * 2 + 1
col_indices = col_indices.load().reshape((2, 1))
for global_chunk_id in range(bid, num_global_chunks, grid_x):
if warp_id_ == 0:
cute.arch.mbarrier_wait(qk_mbar, parity)
cute.arch.barrier(barrier_id=1, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
qk = _tcgen05.ld(warp_id_ * 32, qk_tmem, "16x256b", BT // 8)
qk = qk.reshape((2, 2, BT // 8))
_tcgen05.wait_ld()
for i in cutlass.range_constexpr(BT // 8):
# KDA: Aqk already carries the per-channel decay (no Gamma).
tmp = qk[None, None, i]
tmp = cute.where(row_indices >= col_indices + i * 8, tmp, 0.0)
attn_lo = cute.make_rmem_tensor(2, Uint32)
attn_lo[0] = cvt.fp32x2_to_bf16x2(tmp[0, 0], tmp[1, 0])
attn_lo[1] = cvt.fp32x2_to_bf16x2(tmp[0, 1], tmp[1, 1])
_tcgen05.st(warp_id_ * 32, p_tmem + i * 4, "16x128b", 1, attn_lo)
_tcgen05.wait_st()
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(mask_mbar)
parity ^= 1
else:
# epilogue warps -- KDA: O = QH + P@V (scale & exp(g_cu) baked into qg).
row0 = warp_id * 16 + lane_id // 4
row1 = row0 + 8
stage_id = 0
mma_parity = 0
op = cute.nvgpu.CopyUniversalOp()
cp_4B = cute.make_copy_atom(op, BFloat16, num_bits_per_copy=32)
stsm_op = warp.StMatrix8x8x16bOp(num_matrices=4, transpose=False)
stsm_atom = cute.make_copy_atom(stsm_op, BFloat16)
WIDTH = 64
o_view = cute.logical_divide(
o[None, v_head_id, None],
(None, cute.make_layout((2, 4, WIDTH // 8))),
)
o_view = o_view[None, ((None, lane_id % 4, None), None)]
for global_chunk_id in range(bid, num_global_chunks, grid_x):
seq_id = chunk_indices[global_chunk_id, 0]
chunk_id = chunk_indices[global_chunk_id, 1]
bos = cu_seqlens[seq_id]
eos = cu_seqlens[seq_id + 1]
chunk_start = bos + chunk_id * BT
full_chunk = chunk_start + BT <= eos
if warp_id == 0:
cute.arch.mbarrier_wait(pv_mma_mbar + stage_id, mma_parity)
elif warp_id == 3 and full_chunk:
cute.arch.cp_async_bulk_wait_group(0, read=True)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
_tcgen05.fence_after_thread_sync()
if full_chunk:
for i in cutlass.range_constexpr(V_dim // WIDTH):
qh = _tcgen05.ld(
warp_id * 32, qh_tmem + i * WIDTH, "16x256b", WIDTH // 8
)
pv = _tcgen05.ld(
warp_id * 32, out_tmem + i * WIDTH, "16x256b", WIDTH // 8
)
_tcgen05.wait_ld()
if i == V_dim // WIDTH - 1:
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(epi_mbar)
qh = qh.reshape((2, 2, WIDTH // 8))
pv = pv.reshape((2, 2, WIDTH // 8))
out_f32 = qh + pv
out_bf16 = cute.make_rmem_tensor((8, WIDTH // 16), BFloat16)
out_bf16.store(out_f32.to(BFloat16).reshape((8, WIDTH // 16)))
for j in cutlass.range_constexpr(WIDTH // 16):
s_row = warp_id * 16 + lane_id % 16
s_col = i * (WIDTH // 8) + j * 2 + lane_id // 16
sO_tile = cute.local_tile(sO[s_row, None], (8,), (s_col,))
cute.copy(stsm_atom, out_bf16[None, j], sO_tile)
cute.arch.barrier(barrier_id=2, number_of_threads=128)
fence_before_tma_store()
if warp_id == 3:
gO = cute.local_tile(
cute.domain_offset((bos, 0), tmaO[None, v_head_id, None]),
tiler=(BT, V_dim),
coord=(chunk_id, 0),
)
simple_tma_copy(O_tma_atom, sO, gO)
with cute.arch.elect_one():
cute.arch.cp_async_bulk_commit_group()
else:
for i in cutlass.range_constexpr(V_dim // WIDTH):
qh = _tcgen05.ld(
warp_id * 32, qh_tmem + i * WIDTH, "16x256b", WIDTH // 8
)
pv = _tcgen05.ld(
warp_id * 32, out_tmem + i * WIDTH, "16x256b", WIDTH // 8
)
_tcgen05.wait_ld()
if i == V_dim // WIDTH - 1:
_tcgen05.fence_before_thread_sync()
cute.arch.mbarrier_arrive(epi_mbar)
qh = qh.reshape((2, 2, WIDTH // 8))
pv = pv.reshape((2, 2, WIDTH // 8))
out_f32 = qh + pv
out_bf16 = cute.make_rmem_tensor((2, 2, WIDTH // 8), BFloat16)
out_bf16.store(out_f32.to(BFloat16))
if chunk_start + row0 < eos:
cute.copy(
cp_4B,
out_bf16[None, 0, None],
o_view[chunk_start + row0, None, None, i],
)
if chunk_start + row1 < eos:
cute.copy(
cp_4B,
out_bf16[None, 1, None],
o_view[chunk_start + row1, None, None, i],
)
stage_id = (stage_id + 1) % num_stages
if stage_id == 0:
mma_parity ^= 1
@cache
@staticmethod
def compile(
H: int,
Hv: int,
K_dim: int,
V_dim: int,
BT: int = 64,
num_stages: int = 2,
):
total_t = cute.sym_int()
pad_t = cute.sym_int()
total_chunks_n = cute.sym_int()
h_outer_n = cute.sym_int()
cu_entries = cute.sym_int()
qg = make_fake_tensor(BFloat16, (total_t, Hv, K_dim), divisibility=16)
qg2 = make_fake_tensor(BFloat16, (total_t, Hv, K_dim), divisibility=16)
kg = make_fake_tensor(BFloat16, (total_t, Hv, K_dim), divisibility=16)
v_new = make_fake_tensor(BFloat16, (pad_t, Hv, V_dim), divisibility=16)
h_flat = make_fake_tensor(BFloat16, (h_outer_n, V_dim, K_dim), divisibility=16)
o = make_fake_tensor(BFloat16, (total_t, Hv, V_dim), divisibility=16)
cu_seqlens = make_fake_tensor(Int32, (cu_entries,), divisibility=1)
chunk_indices = make_fake_tensor(Int32, (total_chunks_n, 2), divisibility=2)
total_chunks = make_fake_tensor(Int32, (1,), divisibility=1)
kernel = Sm100KdaChunkOKernel(H, Hv, K_dim, V_dim, BT, num_stages)
stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True)
return cute.compile(
kernel,
qg,
qg2,
kg,
v_new,
h_flat,
o,
cu_seqlens,
chunk_indices,
total_chunks,
Int32(148),
stream,
options="--enable-tvm-ffi",
)
def kda_o_cutedsl(
qg: torch.Tensor,
qg2: torch.Tensor,
kg: torch.Tensor,
v_new_chunks: torch.Tensor,
h: torch.Tensor,
o: torch.Tensor,
cu_seqlens: torch.Tensor,
chunk_indices: torch.Tensor,
total_chunks: torch.Tensor,
num_sms: int = 148,
) -> None:
"""KDA output kernel. qg/qg2/kg are the pre-scaled tensors (see module doc)."""
_, Hv, K_dim = qg.shape
_, _, V_dim = o.shape
Sm100KdaChunkOKernel.compile(Hv, Hv, K_dim, V_dim)(
qg,
qg2,
kg,
v_new_chunks.view(-1, Hv, V_dim),
h.view(-1, V_dim, K_dim),
o,
cu_seqlens,
chunk_indices,
total_chunks,
num_sms,
)
@@ -0,0 +1,102 @@
# SPDX-License-Identifier: Apache-2.0
# Fused Triton prologue for the KDA Blackwell pipeline.
#
# In ONE pass per (chunk, head) it computes the per-chunk cumsum g_cu and the five
# pre-scaled key/query tensors the cutedsl kernels consume, replacing ~30 separate
# PyTorch elementwise ops + copies:
#
# g_cu = cumsum_within_chunk(g) [T, Hv, K] (fp32, for kernel_h decay)
# g_last[d] = g_cu at the chunk's last token (= total sum over the chunk)
# kL = k * exp(g_cu - g_last) (kkt KKT-left)
# kR = k * exp(g_last - g_cu) (kkt KKT-right == kernel_h kg == kernel_o Aqk-K)
# kgw = k * exp(g_cu) (kkt W operand)
# qg = scale * q * exp(g_cu) (kernel_o Q@H)
# qg2 = scale * q * exp(g_cu - g_last) (kernel_o Aqk-Q)
import torch
import triton
import triton.language as tl
@triton.jit
def _kda_prologue_kernel(
q_ptr,
k_ptr,
g_ptr,
kL_ptr,
kR_ptr,
kgw_ptr,
qg_ptr,
qg2_ptr,
gcu_ptr,
cu_seqlens_ptr,
chunk_indices_ptr,
scale,
Hv: tl.constexpr,
K: tl.constexpr,
BT: tl.constexpr,
):
chunk = tl.program_id(0)
head = tl.program_id(1)
seq_id = tl.load(chunk_indices_ptr + chunk * 2 + 0)
chunk_id = tl.load(chunk_indices_ptr + chunk * 2 + 1)
bos = tl.load(cu_seqlens_ptr + seq_id)
eos = tl.load(cu_seqlens_ptr + seq_id + 1)
off_t = bos + chunk_id * BT
row = off_t + tl.arange(0, BT)
col = tl.arange(0, K)
mask_row = row < eos
offs = row[:, None] * (Hv * K) + head * K + col[None, :]
mask = mask_row[:, None]
g = tl.load(g_ptr + offs, mask=mask, other=0.0).to(tl.float32)
q = tl.load(q_ptr + offs, mask=mask, other=0.0).to(tl.float32)
k = tl.load(k_ptr + offs, mask=mask, other=0.0).to(tl.float32)
g_cu = tl.cumsum(g, axis=0) # [BT, K]
g_last = tl.sum(g, axis=0) # [K] (OOB rows contributed 0)
gml = g_cu - g_last[None, :] # g_cu - g_last (>= 0, since g_cu>=g_last)
e_gcu = tl.exp(g_cu) # <= 1
e_gml = tl.exp(gml) # >= 1 (kL side; huge entries get masked)
e_lmg = tl.exp(-gml) # <= 1 (bounded: kR / kg)
tl.store(gcu_ptr + offs, g_cu, mask=mask)
tl.store(kL_ptr + offs, (k * e_gml).to(kL_ptr.dtype.element_ty), mask=mask)
tl.store(kR_ptr + offs, (k * e_lmg).to(kR_ptr.dtype.element_ty), mask=mask)
tl.store(kgw_ptr + offs, (k * e_gcu).to(kgw_ptr.dtype.element_ty), mask=mask)
tl.store(qg_ptr + offs, (scale * q * e_gcu).to(qg_ptr.dtype.element_ty), mask=mask)
tl.store(
qg2_ptr + offs, (scale * q * e_gml).to(qg2_ptr.dtype.element_ty), mask=mask
)
def kda_prologue(q, k, g_act, scale, cu_seqlens, chunk_indices, num_chunks):
"""q/k/g_act: [T, Hv, K]. Returns (kL, kR, kgw, qg, qg2) bf16 + g_cu fp32."""
T, Hv, K = q.shape
kL = torch.empty_like(q, dtype=torch.bfloat16)
kR = torch.empty_like(q, dtype=torch.bfloat16)
kgw = torch.empty_like(q, dtype=torch.bfloat16)
qg = torch.empty_like(q, dtype=torch.bfloat16)
qg2 = torch.empty_like(q, dtype=torch.bfloat16)
g_cu = torch.empty_like(q, dtype=torch.float32)
grid = (num_chunks, Hv)
_kda_prologue_kernel[grid](
q,
k,
g_act,
kL,
kR,
kgw,
qg,
qg2,
g_cu,
cu_seqlens,
chunk_indices,
scale,
Hv=Hv,
K=K,
BT=64,
num_warps=8,
)
return kL, kR, kgw, qg, qg2, g_cu
@@ -1,3 +1,6 @@
import logging
from typing import Optional
import torch
from sglang.jit_kernel.cutedsl_kda import cutedsl_fused_sigmoid_gating_kda_update
@@ -5,9 +8,55 @@ from sglang.srt.layers.attention.linear.kernels.kernel_backend import (
LinearAttnKernelBase,
)
logger = logging.getLogger(__name__)
def _is_blackwell() -> bool:
"""True iff running on SM100+ (Blackwell), where the chunk prefill kernels run."""
if not torch.cuda.is_available():
return False
major, _ = torch.cuda.get_device_capability()
return major >= 10
class CuteDSLKDAKernel(LinearAttnKernelBase):
"""CuTe DSL kernel for KDA decode (CUDA only)."""
"""CuTe DSL kernel for KDA.
Decode: ``cutedsl_fused_sigmoid_gating_kda_update`` (SM90+).
Extend (prefill): SM100 chunk pipeline ``chunk_kda_cutedsl`` (SM100+ only,
``head_k_dim`` must be 128). On SM90 the prefill path is unsupported; callers
query :attr:`supports_prefill` and fall back to Triton.
"""
def __init__(self):
self.supports_prefill = _is_blackwell()
self._extend_fn: Optional[callable] = None
self._l2norm_fn: Optional[callable] = None
def _ensure_extend_loaded(self, head_k_dim: int) -> None:
if self._extend_fn is not None:
return
if not self.supports_prefill:
major = (
torch.cuda.get_device_capability()[0]
if torch.cuda.is_available()
else -1
)
raise RuntimeError(
f"CuTe DSL KDA prefill requires SM100+ (Blackwell); got SM{major}."
)
if head_k_dim != 128:
raise RuntimeError(
f"CuTe DSL KDA prefill requires head_k_dim=128, got {head_k_dim}."
)
from sglang.srt.layers.attention.fla.l2norm import l2norm_fwd
from sglang.srt.layers.attention.linear.kernels.kda_blackwell import (
chunk_kda_cutedsl,
)
self._extend_fn = chunk_kda_cutedsl
self._l2norm_fn = l2norm_fwd
logger.info("Using CuTe DSL KDA prefill (Blackwell)")
def decode(
self,
@@ -40,8 +89,60 @@ class CuteDSLKDAKernel(LinearAttnKernelBase):
softplus_threshold=20.0,
)
def extend(self, *args, **kwargs):
raise NotImplementedError("CuteDSLKDAKernel only supports decode")
def extend(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
*,
ssm_states: torch.Tensor,
cache_indices: torch.Tensor,
query_start_loc: torch.Tensor,
A_log: Optional[torch.Tensor] = None,
dt_bias: Optional[torch.Tensor] = None,
lower_bound: Optional[float] = None,
**kwargs,
) -> torch.Tensor:
head_k_dim = k.shape[-1]
self._ensure_extend_loaded(head_k_dim)
# [1, T, HV, D] -> [T, HV, D]; L2-norm Q/K outside the kernel.
q_n = self._l2norm_fn(q[0].contiguous()).to(torch.bfloat16)
k_n = self._l2norm_fn(k[0].contiguous()).to(torch.bfloat16)
v_in = v[0].contiguous().to(torch.bfloat16)
# Trim g/beta to q's real token count: the [:real_num_tokens] slice in
# unified_linear_attention_with_output narrows their batch dim (a no-op),
# not tokens, so padded rows survive and break the kernel's shape check.
num_tokens = q_n.shape[0]
g_in = g[0][:num_tokens] # raw forget gate; activated inside chunk_kda_cutedsl
beta_in = beta[0][:num_tokens].to(torch.float32)
cu_seqlens = query_start_loc.to(torch.int32)
# Pool gather: remap padding (-1) to the last (sentinel) slot. State is
# [slots, HV, V, K] == cutedsl [V,K] layout, no transpose needed.
ssm_cache_indices = torch.where(
cache_indices >= 0, cache_indices, ssm_states.shape[0] - 1
).to(torch.long)
initial_state = ssm_states[ssm_cache_indices].contiguous()
o, final_state = self._extend_fn(
q_n,
k_n,
v_in,
g_in,
beta_in,
initial_state,
cu_seqlens,
A_log=A_log,
dt_bias=dt_bias,
lower_bound=lower_bound,
)
ssm_states.index_copy_(0, ssm_cache_indices, final_state.to(ssm_states.dtype))
# Match chunk_kda's output layout [1, T, HV, V].
return o.unsqueeze(0)
def target_verify(self, *args, **kwargs):
raise NotImplementedError("CuteDSLKDAKernel only supports decode")
raise NotImplementedError("CuteDSLKDAKernel does not support target_verify")
@@ -0,0 +1,226 @@
"""Correctness test for the SM100 CuTe DSL KDA prefill pipeline.
Validates ``chunk_kda_cutedsl`` (the ``kda_blackwell`` package: fused Triton
prologue -> kkt_inv_uw -> h -> o) against the token-by-token
``fused_recurrent_kda`` Triton reference. Mirrors ``test_gdn_prefill_cutedsl.py``.
KDA differs from GDN by a PER-CHANNEL decay gate (g is [T, H, K], not scalar
[T, H]); otherwise the chunk pipeline, head dims, and recurrent-state layout
[N, H, V, K] are identical.
"""
import pytest
import torch
import torch.nn.functional as F
from sglang.test.ci.ci_register import register_cuda_ci
# CuteDSL prefill kernel only exists on Blackwell. Single-GPU kernel-unit suite,
# same slot as the GDN prefill test.
register_cuda_ci(est_time=60, suite="base-b-kernel-unit-1-gpu-b200")
if not (torch.cuda.is_available() and torch.cuda.get_device_capability()[0] >= 10):
pytest.skip(
"KDA CuteDSL prefill requires CUDA SM10x (Blackwell).",
allow_module_level=True,
)
from sglang.srt.layers.attention.fla.index import ( # noqa: E402
prepare_chunk_indices,
prepare_chunk_offsets,
)
from sglang.srt.layers.attention.fla.kda import fused_recurrent_kda # noqa: E402
from sglang.srt.layers.attention.linear.kernels.kda_blackwell import ( # noqa: E402
chunk_kda_cutedsl,
prepare_metadata,
)
def _l2norm(x: torch.Tensor) -> torch.Tensor:
return F.normalize(x.float(), p=2, dim=-1)
@pytest.mark.parametrize("num_seqs", [1, 5, 257])
def test_kda_chunk_cutedsl_correctness(num_seqs: int):
torch.manual_seed(num_seqs)
seq_lens = torch.randint(1, 130, (num_seqs,), dtype=torch.int32)
cu_seqlens = torch.zeros(num_seqs + 1, device="cuda", dtype=torch.int32)
cu_seqlens[1:] = seq_lens.to("cuda").cumsum(0)
total_tokens = int(cu_seqlens[-1].item())
# KDA shares the head count across q/k/v (G=1). CuteDSL prefill hard-requires
# head_k_dim == head_v_dim == 128.
num_heads = 8
head_dim = 128
scale = head_dim**-0.5
q = _l2norm(torch.randn(1, total_tokens, num_heads, head_dim, device="cuda"))
k = _l2norm(torch.randn(1, total_tokens, num_heads, head_dim, device="cuda"))
v = torch.randn(1, total_tokens, num_heads, head_dim, device="cuda")
# Per-channel KDA gate. Mild gates keep the kernel's externalized per-channel
# pre-scaling (qg2 = scale*q*exp(g_cu - g_last), unbounded >= 1) inside fp32
# range -- this matches real Kimi-Linear retention gates. Large per-chunk gate
# spans are a known limitation (B2 TODO: clamp / sub-chunk normalize).
A_log = torch.randn(num_heads, device="cuda") * 0.5 - 1.5
dt_bias = torch.randn(num_heads, head_dim, device="cuda") * 0.1
g_raw = torch.randn(1, total_tokens, num_heads, head_dim, device="cuda")
g_act = -A_log.exp().view(1, 1, num_heads, 1) * F.softplus(
g_raw + dt_bias.view(1, 1, num_heads, head_dim)
)
beta = torch.sigmoid(torch.randn(1, total_tokens, num_heads, device="cuda")).float()
# Recurrent-state layout [N, H, V, K] (V-major) -- identical for the recurrent
# reference and the cutedsl ht output (no transpose needed).
initial_state = (
torch.randn(num_seqs, num_heads, head_dim, head_dim, device="cuda") * 0.05
).float()
# --- metadata helper must match the FLA chunkers the Triton path uses ---
chunk_indices, chunk_offsets, _, total_chunks = prepare_metadata(cu_seqlens)
torch.cuda.synchronize()
expected_indices = prepare_chunk_indices(cu_seqlens, 64)
expected_offsets = prepare_chunk_offsets(cu_seqlens, 64)
torch.testing.assert_close(chunk_offsets, expected_offsets.to(torch.int32))
torch.testing.assert_close(chunk_indices[:total_chunks], expected_indices)
# --- reference: token-by-token recurrent kernel (ground truth) ---
ref_o, ref_state = fused_recurrent_kda(
q=q,
k=k,
v=v,
g=g_act,
beta=beta,
scale=scale,
initial_state=initial_state.clone(),
inplace_final_state=False,
use_qk_l2norm_in_kernel=False,
cu_seqlens=cu_seqlens.long(),
)
# --- cutedsl chunk prefill (q/k already L2-normed; pass g_act directly) ---
o, ht = chunk_kda_cutedsl(
q[0].bfloat16(),
k[0].bfloat16(),
v[0].bfloat16(),
g_act[0].float(),
beta[0].float(),
initial_state.clone(),
cu_seqlens,
scale,
)
torch.cuda.synchronize()
assert torch.isfinite(o).all(), "cutedsl output has non-finite values"
assert torch.isfinite(ht).all(), "cutedsl final state has non-finite values"
o_error = (o.float() - ref_o[0].float()).abs()
state_error = (ht.float() - ref_state.float()).abs()
# bf16 MMA + Newton-Schulz inverse noise; matches the B200 e2e validation
# (o ~5e-4, ht ~4e-3) with margin.
assert o_error.max().item() < 1e-2
assert o_error.mean().item() < 1e-3
assert state_error.max().item() < 5e-2
assert state_error.mean().item() < 5e-3
def test_kda_chunk_cutedsl_internal_gate_activation():
"""The A_log/dt_bias gate activation inside chunk_kda_cutedsl must match
feeding a pre-activated gate."""
torch.manual_seed(0)
T = 256
num_heads = 8
head_dim = 128
scale = head_dim**-0.5
cu_seqlens = torch.tensor([0, T], dtype=torch.int32, device="cuda")
q = _l2norm(torch.randn(1, T, num_heads, head_dim, device="cuda")).bfloat16()
k = _l2norm(torch.randn(1, T, num_heads, head_dim, device="cuda")).bfloat16()
v = torch.randn(1, T, num_heads, head_dim, device="cuda").bfloat16()
A_log = torch.randn(num_heads, device="cuda") * 0.5 - 1.5
dt_bias = torch.randn(num_heads, head_dim, device="cuda") * 0.1
g_raw = torch.randn(1, T, num_heads, head_dim, device="cuda")
beta = torch.sigmoid(torch.randn(1, T, num_heads, device="cuda")).float()
h0 = torch.zeros(1, num_heads, head_dim, head_dim, device="cuda")
g_pre = -A_log.exp().view(1, num_heads, 1) * F.softplus(
g_raw[0] + dt_bias.view(1, num_heads, head_dim)
)
o_pre, ht_pre = chunk_kda_cutedsl(
q[0], k[0], v[0], g_pre.float(), beta[0], h0.clone(), cu_seqlens, scale
)
o_int, ht_int = chunk_kda_cutedsl(
q[0],
k[0],
v[0],
g_raw[0].float(),
beta[0],
h0.clone(),
cu_seqlens,
scale,
A_log=A_log,
dt_bias=dt_bias,
)
torch.cuda.synchronize()
torch.testing.assert_close(o_int.float(), o_pre.float(), atol=2e-3, rtol=1e-2)
torch.testing.assert_close(ht_int.float(), ht_pre.float(), atol=2e-3, rtol=1e-2)
def test_kda_chunk_cutedsl_realistic_gate():
"""Real Kimi-Linear retention gates are far stronger than the mild gates in
the tests above (exp(A_log) ~ 0.2 there vs up to ~200 in the model). Strong
gates make the per-chunk cumulative decay span large; a chunk-global g_last
reference would overflow the externalized exp(g_cu - g_last) pre-scaling.
The intra-chunk matrices use a sub-chunk-normalized path instead, so this
must stay finite and match the recurrent reference."""
torch.manual_seed(0)
num_heads = 8
head_dim = 128
scale = head_dim**-0.5
T = 128
cu_seqlens = torch.tensor([0, T], dtype=torch.int32, device="cuda")
q = _l2norm(torch.randn(1, T, num_heads, head_dim, device="cuda"))
k = _l2norm(torch.randn(1, T, num_heads, head_dim, device="cuda"))
v = torch.randn(1, T, num_heads, head_dim, device="cuda")
# exp(A_log) ~ exp(1.5) ~ 4.5 mean (real model reaches ~200), vs ~0.22 above.
A_log = torch.randn(num_heads, device="cuda") * 0.5 + 1.5
dt_bias = torch.randn(num_heads, head_dim, device="cuda") * 0.1
g_raw = torch.randn(1, T, num_heads, head_dim, device="cuda")
g_act = -A_log.exp().view(1, 1, num_heads, 1) * F.softplus(
g_raw + dt_bias.view(1, 1, num_heads, head_dim)
)
beta = torch.sigmoid(torch.randn(1, T, num_heads, device="cuda")).float()
h0 = torch.zeros(1, num_heads, head_dim, head_dim, device="cuda").float()
ref_o, _ = fused_recurrent_kda(
q=q,
k=k,
v=v,
g=g_act,
beta=beta,
scale=scale,
initial_state=h0.clone(),
inplace_final_state=False,
use_qk_l2norm_in_kernel=False,
cu_seqlens=cu_seqlens.long(),
)
o, ht = chunk_kda_cutedsl(
q[0].bfloat16(),
k[0].bfloat16(),
v[0].bfloat16(),
g_act[0].float(),
beta[0].float(),
h0.clone(),
cu_seqlens,
scale,
)
torch.cuda.synchronize()
assert torch.isfinite(o).all() and torch.isfinite(ht).all()
assert (o.float() - ref_o[0].float()).abs().max().item() < 1e-2
if __name__ == "__main__":
import sys
sys.exit(pytest.main([__file__, "-v"]))