From 518e35fae7ae881c4a2f2e249a737bb48a6caeb8 Mon Sep 17 00:00:00 2001 From: Yuan Luo Date: Wed, 10 Jun 2026 21:25:19 +0800 Subject: [PATCH] [KDA] Add CuteDSL Prefill Kernel on SM100 (#27488) Co-authored-by: luoyuan.luo --- .../bench_kda_prefill_cutedsl.py | 357 +++++++++ .../layers/attention/linear/kda_backend.py | 20 +- .../linear/kernels/kda_blackwell/__init__.py | 221 ++++++ .../linear/kernels/kda_blackwell/kernel_h.py | 690 ++++++++++++++++ .../kda_blackwell/kernel_kkt_inv_uw.py | 741 ++++++++++++++++++ .../linear/kernels/kda_blackwell/kernel_o.py | 584 ++++++++++++++ .../linear/kernels/kda_blackwell/prologue.py | 102 +++ .../attention/linear/kernels/kda_cutedsl.py | 109 ++- .../attention/test_kda_prefill_cutedsl.py | 226 ++++++ 9 files changed, 3045 insertions(+), 5 deletions(-) create mode 100644 benchmark/bench_linear_attention/bench_kda_prefill_cutedsl.py create mode 100644 python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/__init__.py create mode 100644 python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_h.py create mode 100644 python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_kkt_inv_uw.py create mode 100644 python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_o.py create mode 100644 python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/prologue.py create mode 100644 test/registered/attention/test_kda_prefill_cutedsl.py diff --git a/benchmark/bench_linear_attention/bench_kda_prefill_cutedsl.py b/benchmark/bench_linear_attention/bench_kda_prefill_cutedsl.py new file mode 100644 index 000000000..08fe40003 --- /dev/null +++ b/benchmark/bench_linear_attention/bench_kda_prefill_cutedsl.py @@ -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()) diff --git a/python/sglang/srt/layers/attention/linear/kda_backend.py b/python/sglang/srt/layers/attention/linear/kda_backend.py index d4e63cac8..065dc1110 100644 --- a/python/sglang/srt/layers/attention/linear/kda_backend.py +++ b/python/sglang/srt/layers/attention/linear/kda_backend.py @@ -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( diff --git a/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/__init__.py b/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/__init__.py new file mode 100644 index 000000000..92399d7eb --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/__init__.py @@ -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 diff --git a/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_h.py b/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_h.py new file mode 100644 index 000000000..a06de16be --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_h.py @@ -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 + ) diff --git a/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_kkt_inv_uw.py b/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_kkt_inv_uw.py new file mode 100644 index 000000000..7900edefb --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_kkt_inv_uw.py @@ -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 + ) diff --git a/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_o.py b/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_o.py new file mode 100644 index 000000000..70c4634c3 --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_o.py @@ -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, + ) diff --git a/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/prologue.py b/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/prologue.py new file mode 100644 index 000000000..7ffe7027a --- /dev/null +++ b/python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/prologue.py @@ -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 diff --git a/python/sglang/srt/layers/attention/linear/kernels/kda_cutedsl.py b/python/sglang/srt/layers/attention/linear/kernels/kda_cutedsl.py index c91cf691c..b8583ada8 100644 --- a/python/sglang/srt/layers/attention/linear/kernels/kda_cutedsl.py +++ b/python/sglang/srt/layers/attention/linear/kernels/kda_cutedsl.py @@ -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") diff --git a/test/registered/attention/test_kda_prefill_cutedsl.py b/test/registered/attention/test_kda_prefill_cutedsl.py new file mode 100644 index 000000000..2efefdc76 --- /dev/null +++ b/test/registered/attention/test_kda_prefill_cutedsl.py @@ -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"]))