[KDA] Add FlashInfer SM100 KDA decode + MTP (target_verify) backend (#30113)
Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
@@ -0,0 +1,280 @@
|
|||||||
|
"""
|
||||||
|
Benchmark & Correctness: FlashInfer KDA (SM100) vs Triton KDA — decode & MTP verify.
|
||||||
|
|
||||||
|
Exercises the two real backend wrappers used by ``KDAKernelDispatcher``:
|
||||||
|
- ``FlashInferKDAKernel`` — wraps ``flashinfer.kda_decode.recurrent_kda``
|
||||||
|
(CuTe DSL, SM100/Blackwell only). Provides ``decode`` + ``target_verify``.
|
||||||
|
- ``TritonKDAKernel`` — wraps ``fused_sigmoid_gating_delta_rule_update``
|
||||||
|
(IS_KDA=True). Reference for both ``decode`` and ``target_verify``.
|
||||||
|
|
||||||
|
Two modes:
|
||||||
|
- decode : single-token decode (T=1), in-place SSM update.
|
||||||
|
- verify : MTP / speculative-decode ``target_verify`` over T=1+num_spec draft
|
||||||
|
tokens per sequence, writing per-token states into the speculative
|
||||||
|
``intermediate_ssm`` scratch (the recurrent_kda adapter / the Triton
|
||||||
|
intermediate_states_buffer path).
|
||||||
|
|
||||||
|
Reports correctness (output vs the Triton reference) and performance (us, speedup).
|
||||||
|
Requires an SM100 GPU + a FlashInfer build exposing ``recurrent_kda``; on other
|
||||||
|
GPUs the FlashInfer side is skipped and only the Triton path is timed.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python bench_kda_flashinfer_mtp.py # decode+verify, correctness+bench
|
||||||
|
python bench_kda_flashinfer_mtp.py --mode bench --task verify
|
||||||
|
python bench_kda_flashinfer_mtp.py --num-spec 7 # 8 draft tokens / verify step
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.attention.linear.kernels.kda_triton import TritonKDAKernel
|
||||||
|
|
||||||
|
|
||||||
|
def _make_flashinfer_kernel():
|
||||||
|
"""Instantiate FlashInferKDAKernel, or None if unavailable (non-SM100)."""
|
||||||
|
try:
|
||||||
|
from sglang.srt.layers.attention.linear.kernels.kda_flashinfer import (
|
||||||
|
FlashInferKDAKernel,
|
||||||
|
)
|
||||||
|
|
||||||
|
return FlashInferKDAKernel()
|
||||||
|
except Exception as e: # noqa: BLE001 - report and degrade gracefully
|
||||||
|
print(f" [skip flashinfer] {type(e).__name__}: {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Input construction
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
def make_decode_inputs(B, H, HV, K, V, pool_size, device, dtype, seed=42):
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
q = torch.randn(1, B, H, K, device=device, dtype=dtype) * 0.5
|
||||||
|
k = torch.randn(1, B, H, K, device=device, dtype=dtype) * 0.5
|
||||||
|
v = torch.randn(1, B, HV, V, device=device, dtype=dtype) * 0.5
|
||||||
|
a = torch.randn(B, HV * K, device=device, dtype=dtype) * 0.5 - 1.0 # raw per-K gate
|
||||||
|
b = torch.randn(B, HV, device=device, dtype=dtype) * 0.5 # beta LOGIT
|
||||||
|
A_log = torch.randn(HV, device=device, dtype=torch.float32) * 0.2
|
||||||
|
dt_bias = torch.randn(HV * K, device=device, dtype=torch.float32) * 0.1
|
||||||
|
ssm = torch.randn(pool_size, HV, V, K, device=device, dtype=dtype) * 0.01
|
||||||
|
cache_indices = torch.arange(B, device=device, dtype=torch.int32)
|
||||||
|
qsl = torch.arange(B + 1, device=device, dtype=torch.int32)
|
||||||
|
return dict(
|
||||||
|
q=q.contiguous(),
|
||||||
|
k=k.contiguous(),
|
||||||
|
v=v.contiguous(),
|
||||||
|
a=a.contiguous(),
|
||||||
|
b=b.contiguous(),
|
||||||
|
A_log=A_log,
|
||||||
|
dt_bias=dt_bias,
|
||||||
|
ssm=ssm.contiguous(),
|
||||||
|
cache_indices=cache_indices,
|
||||||
|
qsl=qsl,
|
||||||
|
B=B,
|
||||||
|
H=H,
|
||||||
|
HV=HV,
|
||||||
|
K=K,
|
||||||
|
V=V,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def make_verify_inputs(B, T, H, HV, K, V, pool_size, device, dtype, seed=42):
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
seq = B * T
|
||||||
|
q = torch.randn(1, seq, H, K, device=device, dtype=dtype) * 0.5
|
||||||
|
k = torch.randn(1, seq, H, K, device=device, dtype=dtype) * 0.5
|
||||||
|
v = torch.randn(1, seq, HV, V, device=device, dtype=dtype) * 0.5
|
||||||
|
a = torch.randn(seq, HV * K, device=device, dtype=dtype) * 0.5 - 1.0
|
||||||
|
b = torch.randn(seq, HV, device=device, dtype=dtype) * 0.5
|
||||||
|
A_log = torch.randn(HV, device=device, dtype=torch.float32) * 0.2
|
||||||
|
dt_bias = torch.randn(HV * K, device=device, dtype=torch.float32) * 0.1
|
||||||
|
ssm = torch.randn(pool_size, HV, V, K, device=device, dtype=dtype) * 0.01
|
||||||
|
cache_indices = torch.arange(B, device=device, dtype=torch.int32)
|
||||||
|
qsl = torch.arange(0, seq + 1, T, device=device, dtype=torch.int32)
|
||||||
|
# speculative intermediate_ssm scratch: [n_scratch, T, HV, V, K]; per-request row.
|
||||||
|
intermediate_states = torch.zeros(B, T, HV, V, K, device=device, dtype=dtype)
|
||||||
|
intermediate_indices = torch.arange(B, device=device, dtype=torch.int32)
|
||||||
|
return dict(
|
||||||
|
q=q.contiguous(),
|
||||||
|
k=k.contiguous(),
|
||||||
|
v=v.contiguous(),
|
||||||
|
a=a.contiguous(),
|
||||||
|
b=b.contiguous(),
|
||||||
|
A_log=A_log,
|
||||||
|
dt_bias=dt_bias,
|
||||||
|
ssm=ssm.contiguous(),
|
||||||
|
cache_indices=cache_indices,
|
||||||
|
qsl=qsl,
|
||||||
|
intermediate_states=intermediate_states.contiguous(),
|
||||||
|
intermediate_indices=intermediate_indices,
|
||||||
|
B=B,
|
||||||
|
T=T,
|
||||||
|
H=H,
|
||||||
|
HV=HV,
|
||||||
|
K=K,
|
||||||
|
V=V,
|
||||||
|
seq=seq,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Calls (fresh state clone each time so timing/correctness are independent)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
def call_decode(kernel, inp, ssm):
|
||||||
|
# `ssm` is the (mutable, updated in-place) committed-state buffer the caller owns
|
||||||
|
# — cloned fresh for correctness, reused across timed iters (latency is unchanged
|
||||||
|
# by accumulated state; cloning a ~100s-of-MB pool every call would dominate).
|
||||||
|
out = kernel.decode(
|
||||||
|
inp["q"],
|
||||||
|
inp["k"],
|
||||||
|
inp["v"],
|
||||||
|
inp["a"],
|
||||||
|
inp["b"],
|
||||||
|
A_log=inp["A_log"],
|
||||||
|
dt_bias=inp["dt_bias"],
|
||||||
|
ssm_states=ssm,
|
||||||
|
cache_indices=inp["cache_indices"],
|
||||||
|
query_start_loc=inp["qsl"],
|
||||||
|
)
|
||||||
|
return out.reshape(inp["B"], inp["HV"], inp["V"]).float()
|
||||||
|
|
||||||
|
|
||||||
|
def call_verify(kernel, inp, ssm, intermediate_states):
|
||||||
|
out = kernel.target_verify(
|
||||||
|
A_log=inp["A_log"],
|
||||||
|
dt_bias=inp["dt_bias"],
|
||||||
|
q=inp["q"],
|
||||||
|
k=inp["k"],
|
||||||
|
v=inp["v"],
|
||||||
|
a=inp["a"],
|
||||||
|
b=inp["b"],
|
||||||
|
ssm_states=ssm,
|
||||||
|
cache_indices=inp["cache_indices"],
|
||||||
|
query_start_loc=inp["qsl"],
|
||||||
|
intermediate_states_buffer=intermediate_states,
|
||||||
|
intermediate_state_indices=inp["intermediate_indices"],
|
||||||
|
cache_steps=inp["T"],
|
||||||
|
retrieve_parent_token=None,
|
||||||
|
)
|
||||||
|
return out.reshape(inp["seq"], inp["HV"], inp["V"]).float()
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Timing
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
def _time(fn, warmup=20, iters=100):
|
||||||
|
for _ in range(warmup):
|
||||||
|
fn()
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
start = torch.cuda.Event(enable_timing=True)
|
||||||
|
end = torch.cuda.Event(enable_timing=True)
|
||||||
|
start.record()
|
||||||
|
for _ in range(iters):
|
||||||
|
fn()
|
||||||
|
end.record()
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
return start.elapsed_time(end) / iters # ms
|
||||||
|
|
||||||
|
|
||||||
|
def run(task, fi, tri, device, dtype, args):
|
||||||
|
is_verify = task == "verify"
|
||||||
|
T = 1 + args.num_spec if is_verify else 1
|
||||||
|
title = f"target_verify (MTP, T={T})" if is_verify else "decode (T=1)"
|
||||||
|
print("=" * 92)
|
||||||
|
print(
|
||||||
|
f"KDA {title}: FlashInfer (SM100) vs Triton | K={args.head_k} V={args.head_v} dtype={dtype}"
|
||||||
|
)
|
||||||
|
print("=" * 92)
|
||||||
|
hdr = "B" if not is_verify else "B(xT)"
|
||||||
|
print(
|
||||||
|
f" {hdr:>6} {'H':>3} {'HV':>3} | {'triton(us)':>11} | "
|
||||||
|
f"{'flashinfer(us)':>14} | {'speedup':>8} | {'out_max_diff':>12}"
|
||||||
|
)
|
||||||
|
print(" " + "-" * 86)
|
||||||
|
|
||||||
|
for B in args.batch_sizes:
|
||||||
|
for H in args.num_q_heads:
|
||||||
|
for HV in args.num_v_heads:
|
||||||
|
if HV % H != 0:
|
||||||
|
continue
|
||||||
|
K, V = args.head_k, args.head_v
|
||||||
|
pool = max(args.pool_size, B + 16)
|
||||||
|
if is_verify:
|
||||||
|
inp = make_verify_inputs(B, T, H, HV, K, V, pool, device, dtype)
|
||||||
|
corr = lambda kern: call_verify( # noqa: E731
|
||||||
|
kern,
|
||||||
|
inp,
|
||||||
|
inp["ssm"].clone(),
|
||||||
|
inp["intermediate_states"].clone(),
|
||||||
|
)
|
||||||
|
ssm_t, intermediate_states_t = (
|
||||||
|
inp["ssm"].clone(),
|
||||||
|
inp["intermediate_states"].clone(),
|
||||||
|
)
|
||||||
|
timed = lambda kern: call_verify(
|
||||||
|
kern, inp, ssm_t, intermediate_states_t
|
||||||
|
) # noqa: E731
|
||||||
|
else:
|
||||||
|
inp = make_decode_inputs(B, H, HV, K, V, pool, device, dtype)
|
||||||
|
corr = lambda kern: call_decode(
|
||||||
|
kern, inp, inp["ssm"].clone()
|
||||||
|
) # noqa: E731
|
||||||
|
ssm_t = inp["ssm"].clone()
|
||||||
|
timed = lambda kern: call_decode(kern, inp, ssm_t) # noqa: E731
|
||||||
|
|
||||||
|
o_tri = corr(tri)
|
||||||
|
diff = "n/a"
|
||||||
|
if fi is not None:
|
||||||
|
o_fi = corr(fi)
|
||||||
|
diff = f"{(o_fi - o_tri).abs().max().item():.2e}"
|
||||||
|
|
||||||
|
ms_tri = _time(lambda: timed(tri))
|
||||||
|
ms_fi = _time(lambda: timed(fi)) if fi is not None else float("nan")
|
||||||
|
speed = (
|
||||||
|
(ms_tri / ms_fi) if fi is not None and ms_fi > 0 else float("nan")
|
||||||
|
)
|
||||||
|
fi_us = f"{ms_fi * 1000:>14.1f}" if fi is not None else f"{'skip':>14}"
|
||||||
|
sp = f"{speed:>7.2f}x" if fi is not None else f"{'-':>8}"
|
||||||
|
print(
|
||||||
|
f" {B:>6} {H:>3} {HV:>3} | {ms_tri * 1000:>11.1f} | "
|
||||||
|
f"{fi_us} | {sp} | {diff:>12}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
p = argparse.ArgumentParser(
|
||||||
|
description="Benchmark FlashInfer vs Triton KDA decode/verify"
|
||||||
|
)
|
||||||
|
p.add_argument("--task", choices=["decode", "verify", "all"], default="all")
|
||||||
|
p.add_argument(
|
||||||
|
"--mode", choices=["all", "bench"], default="all"
|
||||||
|
) # correctness inlined
|
||||||
|
p.add_argument("--dtype", choices=["bfloat16", "float16"], default="bfloat16")
|
||||||
|
p.add_argument("--head-k", type=int, default=128)
|
||||||
|
p.add_argument("--head-v", type=int, default=128)
|
||||||
|
p.add_argument("--pool-size", type=int, default=512)
|
||||||
|
p.add_argument(
|
||||||
|
"--num-spec", type=int, default=7, help="draft tokens = 1 + num_spec"
|
||||||
|
)
|
||||||
|
p.add_argument(
|
||||||
|
"--batch-sizes", type=int, nargs="+", default=[1, 4, 16, 32, 64, 128]
|
||||||
|
)
|
||||||
|
p.add_argument("--num-q-heads", type=int, nargs="+", default=[16])
|
||||||
|
p.add_argument("--num-v-heads", type=int, nargs="+", default=[16])
|
||||||
|
args = p.parse_args()
|
||||||
|
|
||||||
|
device, dtype = "cuda", getattr(torch, args.dtype)
|
||||||
|
cap = torch.cuda.get_device_capability()
|
||||||
|
print(f"Device: {torch.cuda.get_device_name()} (SM {cap[0]}{cap[1]})")
|
||||||
|
|
||||||
|
fi = _make_flashinfer_kernel()
|
||||||
|
tri = TritonKDAKernel()
|
||||||
|
|
||||||
|
tasks = ["decode", "verify"] if args.task == "all" else [args.task]
|
||||||
|
for t in tasks:
|
||||||
|
run(t, fi, tri, device, dtype, args)
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
raise SystemExit(main())
|
||||||
@@ -52,12 +52,32 @@ class KDAKernelDispatcher:
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.decode_kernel = CuteDSLKDAKernel()
|
self.decode_kernel = CuteDSLKDAKernel()
|
||||||
|
elif decode_backend.is_flashinfer():
|
||||||
|
# FlashInfer recurrent_kda: SM100 decode + MTP (target_verify).
|
||||||
|
# Prefill stays on Triton / CuTe DSL (FlashInfer has no KDA chunk kernel).
|
||||||
|
if not is_cuda():
|
||||||
|
raise ValueError("KDA FlashInfer backend requires CUDA")
|
||||||
|
from sglang.srt.layers.attention.linear.kernels.kda_flashinfer import (
|
||||||
|
FlashInferKDAKernel,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.decode_kernel = FlashInferKDAKernel()
|
||||||
else:
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported KDA decode backend: {decode_backend}. "
|
f"Unsupported KDA decode backend: {decode_backend}. "
|
||||||
"KDA currently only supports 'triton'."
|
"KDA supports 'triton', 'cutedsl', or 'flashinfer'."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# target_verify (MTP / speculative decode) kernel: each decode backend
|
||||||
|
# verifies with its own kernel. FlashInfer decode uses recurrent_kda (SM100,
|
||||||
|
# chain only); Triton -- and CuTe DSL, which has no verify of its own -- use
|
||||||
|
# the Triton fused KDA verify, which handles chain + tree
|
||||||
|
# (retrieve_parent_token) and per-step checkpointing and is the reference the
|
||||||
|
# KDA backend correctness tests assert against.
|
||||||
|
self.verify_kernel = (
|
||||||
|
self.decode_kernel if decode_backend.is_flashinfer() else triton_kernel
|
||||||
|
)
|
||||||
|
|
||||||
if prefill_backend.is_triton():
|
if prefill_backend.is_triton():
|
||||||
self.extend_kernel = triton_kernel
|
self.extend_kernel = triton_kernel
|
||||||
elif prefill_backend.is_flashkda():
|
elif prefill_backend.is_flashkda():
|
||||||
@@ -97,6 +117,7 @@ class KDAKernelDispatcher:
|
|||||||
|
|
||||||
rank0_log(
|
rank0_log(
|
||||||
f"KDA kernel dispatcher: decode={self.decode_kernel.__class__.__name__}, "
|
f"KDA kernel dispatcher: decode={self.decode_kernel.__class__.__name__}, "
|
||||||
|
f"verify={self.verify_kernel.__class__.__name__}, "
|
||||||
f"extend={self.extend_kernel.__class__.__name__} "
|
f"extend={self.extend_kernel.__class__.__name__} "
|
||||||
f"packed_decode={self.supports_packed_decode}"
|
f"packed_decode={self.supports_packed_decode}"
|
||||||
)
|
)
|
||||||
@@ -163,6 +184,45 @@ class KDAKernelDispatcher:
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def target_verify(
|
||||||
|
self,
|
||||||
|
A_log: torch.Tensor,
|
||||||
|
dt_bias: torch.Tensor,
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
a: torch.Tensor,
|
||||||
|
b: torch.Tensor,
|
||||||
|
*,
|
||||||
|
ssm_states: torch.Tensor,
|
||||||
|
cache_indices: torch.Tensor,
|
||||||
|
query_start_loc: torch.Tensor,
|
||||||
|
intermediate_states_buffer: torch.Tensor,
|
||||||
|
intermediate_state_indices: torch.Tensor,
|
||||||
|
cache_steps: int,
|
||||||
|
retrieve_parent_token: torch.Tensor,
|
||||||
|
**kwargs,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""MTP / speculative-decode verify, routed to ``self.verify_kernel``
|
||||||
|
(FlashInfer decode -> recurrent_kda; Triton / CuTe DSL decode -> the Triton
|
||||||
|
fused KDA verify)."""
|
||||||
|
return self.verify_kernel.target_verify(
|
||||||
|
A_log=A_log,
|
||||||
|
dt_bias=dt_bias,
|
||||||
|
q=q,
|
||||||
|
k=k,
|
||||||
|
v=v,
|
||||||
|
a=a,
|
||||||
|
b=b,
|
||||||
|
ssm_states=ssm_states,
|
||||||
|
cache_indices=cache_indices,
|
||||||
|
query_start_loc=query_start_loc,
|
||||||
|
intermediate_states_buffer=intermediate_states_buffer,
|
||||||
|
intermediate_state_indices=intermediate_state_indices,
|
||||||
|
cache_steps=cache_steps,
|
||||||
|
retrieve_parent_token=retrieve_parent_token,
|
||||||
|
)
|
||||||
|
|
||||||
def extend(
|
def extend(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
@@ -196,7 +256,23 @@ class KDAAttnBackend(MambaAttnBackendBase):
|
|||||||
super().__init__(model_runner)
|
super().__init__(model_runner)
|
||||||
decode_backend = get_linear_attn_decode_backend()
|
decode_backend = get_linear_attn_decode_backend()
|
||||||
prefill_backend = get_linear_attn_prefill_backend()
|
prefill_backend = get_linear_attn_prefill_backend()
|
||||||
|
# KDA FlashInfer speculative decode (target_verify) is linear-chain only --
|
||||||
|
# recurrent_kda has no tree-ancestor traversal. Reject EAGLE tree verify
|
||||||
|
# (topk > 1) early at setup instead of deep in the per-step verify call.
|
||||||
|
# (The kernel keeps a per-call retrieve_parent_token guard as a backstop; it
|
||||||
|
# also covers ngram tree, which this topk field does not.)
|
||||||
|
speculative_topk = model_runner.server_args.speculative_eagle_topk or 1
|
||||||
|
if decode_backend.is_flashinfer() and speculative_topk > 1:
|
||||||
|
raise ValueError(
|
||||||
|
"KDA FlashInfer speculative decoding only supports topk=1 "
|
||||||
|
"(EAGLE tree verify / retrieve_parent_token is unsupported)."
|
||||||
|
)
|
||||||
self.kernel_dispatcher = KDAKernelDispatcher(decode_backend, prefill_backend)
|
self.kernel_dispatcher = KDAKernelDispatcher(decode_backend, prefill_backend)
|
||||||
|
# Per-request row index into the speculative `intermediate_ssm` scratch,
|
||||||
|
# used by the MTP / target_verify path (mirrors GDNAttnBackend).
|
||||||
|
self.verify_intermediate_state_indices = torch.arange(
|
||||||
|
self.req_to_token_pool.size, dtype=torch.int32, device=model_runner.device
|
||||||
|
)
|
||||||
|
|
||||||
def forward_decode(
|
def forward_decode(
|
||||||
self,
|
self,
|
||||||
@@ -212,18 +288,8 @@ class KDAAttnBackend(MambaAttnBackendBase):
|
|||||||
query_start_loc = self.forward_metadata.query_start_loc
|
query_start_loc = self.forward_metadata.query_start_loc
|
||||||
cache_indices = self.forward_metadata.mamba_cache_indices
|
cache_indices = self.forward_metadata.mamba_cache_indices
|
||||||
|
|
||||||
# ReplaySSM ring: per-layer ring slices + the once-per-forward per-row
|
# ReplaySSM is mostly a GDN bandwidth optimization. It remains wired for
|
||||||
# write cursor. All None unless --enable-linear-replayssm, so packed_decode
|
# KDA correctness paths, but packed decode is faster for KDA today.
|
||||||
# falls through to the byte-identical legacy KDA path. KDA ships WITHOUT
|
|
||||||
# radix coordination for now, so force_flush is None/zeroed (the ring
|
|
||||||
# flushes only at the natural write_pos == L-1 wrap; set in the shared
|
|
||||||
# HybridLinearAttn metadata, which zeroes force_flush for KDA models).
|
|
||||||
# NOTE: ReplaySSM decode is a GDN (scalar-gate) bandwidth win; on KDA the
|
|
||||||
# per-K g_cache is K x larger and the reconstruction refolds the per-K
|
|
||||||
# decay every step, so it is correct but SLOWER than packed (a measured
|
|
||||||
# decode regression). Kept wired for correctness + the spec-decode path;
|
|
||||||
# not recommended for KDA decode. Revisit on Blackwell (more tensor-core
|
|
||||||
# throughput may flip the compute/bandwidth tradeoff).
|
|
||||||
replayssm_write_pos = getattr(
|
replayssm_write_pos = getattr(
|
||||||
self.forward_metadata, "replayssm_write_pos", None
|
self.forward_metadata, "replayssm_write_pos", None
|
||||||
)
|
)
|
||||||
@@ -243,16 +309,8 @@ class KDAAttnBackend(MambaAttnBackendBase):
|
|||||||
conv_state_indices=cache_indices,
|
conv_state_indices=cache_indices,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Skip split + reshape by consuming the packed mixed_qkv directly in a
|
# The packed kernel assumes one token per request. Assert the dispatch
|
||||||
# single fused Triton kernel (KDA per-K gate variant of GDN PR #20627).
|
# invariant before taking the fused path.
|
||||||
#
|
|
||||||
# The packed kernel hard-assumes one token per sequence (T=1): it has no
|
|
||||||
# query_start_loc / per-sequence loop. forward_decode is only entered in
|
|
||||||
# decode mode (see HybridLinearAttnBackend.forward dispatch), where each
|
|
||||||
# request contributes exactly one token, so #tokens == #requests. Multi-
|
|
||||||
# token-per-seq speculative paths (target_verify / draft_extend) go
|
|
||||||
# through forward_extend instead. Assert the invariant so a future
|
|
||||||
# routing change fails loudly rather than silently corrupting state.
|
|
||||||
if self.kernel_dispatcher.supports_packed_decode:
|
if self.kernel_dispatcher.supports_packed_decode:
|
||||||
assert qkv.shape[0] == cache_indices.shape[0], (
|
assert qkv.shape[0] == cache_indices.shape[0], (
|
||||||
"KDA packed decode requires one token per sequence (T=1): "
|
"KDA packed decode requires one token per sequence (T=1): "
|
||||||
@@ -303,6 +361,11 @@ class KDAAttnBackend(MambaAttnBackendBase):
|
|||||||
b: torch.Tensor,
|
b: torch.Tensor,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
|
# MTP / speculative-decode verify is a multi-token-per-seq path with
|
||||||
|
# per-step state checkpointing + central rollback; handled separately.
|
||||||
|
if forward_batch.forward_mode.is_target_verify():
|
||||||
|
return self._forward_target_verify(layer, forward_batch, mixed_qkv, a, b)
|
||||||
|
|
||||||
query_start_loc = self.forward_metadata.query_start_loc
|
query_start_loc = self.forward_metadata.query_start_loc
|
||||||
cache_indices = self.forward_metadata.mamba_cache_indices
|
cache_indices = self.forward_metadata.mamba_cache_indices
|
||||||
|
|
||||||
@@ -375,13 +438,89 @@ class KDAAttnBackend(MambaAttnBackendBase):
|
|||||||
dt_bias=layer.dt_bias,
|
dt_bias=layer.dt_bias,
|
||||||
lower_bound=getattr(layer, "lower_bound", None),
|
lower_bound=getattr(layer, "lower_bound", None),
|
||||||
extend_seq_lens_cpu=forward_batch.extend_seq_lens_cpu,
|
extend_seq_lens_cpu=forward_batch.extend_seq_lens_cpu,
|
||||||
# target_verify / draft_extend_v2 also reach forward_extend; they must
|
# draft_extend_v2 must stay rollback-able, so kernels that commit state
|
||||||
# stay rollback-able, so a kernel that commits state in place (e.g.
|
# in place (e.g. FlashKDA) must not run for it.
|
||||||
# FlashKDA) must not run for them.
|
is_spec_decode=forward_batch.forward_mode.is_draft_extend_v2(),
|
||||||
is_spec_decode=(
|
|
||||||
forward_batch.forward_mode.is_target_verify()
|
|
||||||
or forward_batch.forward_mode.is_draft_extend_v2()
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
return core_attn_out
|
return core_attn_out
|
||||||
|
|
||||||
|
def _forward_target_verify(
|
||||||
|
self,
|
||||||
|
layer: RadixLinearAttention,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
mixed_qkv: torch.Tensor,
|
||||||
|
a: torch.Tensor,
|
||||||
|
b: torch.Tensor,
|
||||||
|
):
|
||||||
|
"""MTP / speculative-decode verify (topk=1), mirroring the GDN backend.
|
||||||
|
|
||||||
|
Conv1d runs per draft token with intermediate-window checkpointing; the
|
||||||
|
SSM verify kernel writes each draft token's post-state into the
|
||||||
|
speculative `intermediate_ssm` scratch so the central post-verify rollback
|
||||||
|
(update_mamba_state_after_mtp_verify) can commit the accepted-length state.
|
||||||
|
"""
|
||||||
|
fm = self.forward_metadata
|
||||||
|
seq_len = mixed_qkv.shape[0]
|
||||||
|
query_start_loc = fm.query_start_loc
|
||||||
|
cache_indices = fm.mamba_cache_indices
|
||||||
|
retrieve_next_token = fm.retrieve_next_token
|
||||||
|
retrieve_next_sibling = fm.retrieve_next_sibling
|
||||||
|
retrieve_parent_token = fm.retrieve_parent_token
|
||||||
|
|
||||||
|
mamba_cache_params = self.req_to_token_pool.mamba2_layer_cache(layer.layer_id)
|
||||||
|
conv_states = mamba_cache_params.conv[0]
|
||||||
|
ssm_states = mamba_cache_params.temporal
|
||||||
|
intermediate_state_cache = getattr(mamba_cache_params, "intermediate_ssm", None)
|
||||||
|
if intermediate_state_cache is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"KDA target_verify requires a speculative mamba cache "
|
||||||
|
"(MambaPool.SpeculativeState); none found."
|
||||||
|
)
|
||||||
|
intermediate_conv_window_cache = mamba_cache_params.intermediate_conv_window[0]
|
||||||
|
intermediate_state_indices = self.verify_intermediate_state_indices
|
||||||
|
|
||||||
|
draft_token_num = forward_batch.spec_info.draft_token_num
|
||||||
|
batch_size = seq_len // draft_token_num
|
||||||
|
|
||||||
|
# causal_conv1d_update expects [.., dim, width]. KDA keeps dense conv-window
|
||||||
|
# scratch because the deduplicated overlapping layout cannot be transposed.
|
||||||
|
mixed_qkv_reshaped = mixed_qkv.view(batch_size, draft_token_num, -1).transpose(
|
||||||
|
1, 2
|
||||||
|
)
|
||||||
|
mixed_qkv_processed = causal_conv1d_update(
|
||||||
|
mixed_qkv_reshaped,
|
||||||
|
conv_states.transpose(-1, -2),
|
||||||
|
layer.conv_weights,
|
||||||
|
layer.bias,
|
||||||
|
activation="silu",
|
||||||
|
conv_state_indices=cache_indices[:batch_size],
|
||||||
|
intermediate_conv_window=intermediate_conv_window_cache.transpose(-1, -2),
|
||||||
|
intermediate_state_indices=intermediate_state_indices[:batch_size],
|
||||||
|
retrieve_next_token=retrieve_next_token,
|
||||||
|
retrieve_next_sibling=retrieve_next_sibling,
|
||||||
|
retrieve_parent_token=retrieve_parent_token,
|
||||||
|
)
|
||||||
|
mixed_qkv = mixed_qkv_processed.transpose(1, 2).reshape(seq_len, -1)
|
||||||
|
|
||||||
|
q, k, v = mixed_qkv.split([layer.q_dim, layer.k_dim, layer.v_dim], dim=-1)
|
||||||
|
q = q.unflatten(-1, (-1, layer.head_q_dim)).unsqueeze(0) # n (h d) -> 1 n h d
|
||||||
|
k = k.unflatten(-1, (-1, layer.head_k_dim)).unsqueeze(0)
|
||||||
|
v = v.unflatten(-1, (-1, layer.head_v_dim)).unsqueeze(0)
|
||||||
|
|
||||||
|
return self.kernel_dispatcher.target_verify(
|
||||||
|
A_log=layer.A_log,
|
||||||
|
dt_bias=layer.dt_bias,
|
||||||
|
q=q,
|
||||||
|
k=k,
|
||||||
|
v=v,
|
||||||
|
a=a,
|
||||||
|
b=b,
|
||||||
|
ssm_states=ssm_states,
|
||||||
|
cache_indices=cache_indices,
|
||||||
|
query_start_loc=query_start_loc,
|
||||||
|
intermediate_states_buffer=intermediate_state_cache,
|
||||||
|
intermediate_state_indices=intermediate_state_indices,
|
||||||
|
cache_steps=draft_token_num,
|
||||||
|
retrieve_parent_token=retrieve_parent_token,
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,288 @@
|
|||||||
|
"""FlashInfer KDA decode/verify wrapper.
|
||||||
|
|
||||||
|
Wraps ``flashinfer.kda_decode.recurrent_kda`` (SM100 / Blackwell). FlashInfer has
|
||||||
|
no KDA prefill kernel, so ``extend`` stays on Triton / CuTe DSL.
|
||||||
|
|
||||||
|
Contract with the Triton KDA reference:
|
||||||
|
- raw per-K gate ``a`` is activated in-kernel as
|
||||||
|
``-exp(A_log) * softplus(a + dt_bias)``;
|
||||||
|
- beta ``b`` is a logit, so this wrapper passes ``sigmoid(b)``;
|
||||||
|
- q/k are L2-normalized in-kernel;
|
||||||
|
- state layout is ``[N, HV, V, K]`` for committed and speculative state.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.attention.linear.kernels.kernel_backend import (
|
||||||
|
LinearAttnKernelBase,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils import is_cuda
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Lazy import for the FlashInfer KDA kernel
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
_flashinfer_kda_available: Optional[bool] = None
|
||||||
|
_flashinfer_recurrent_kda = None
|
||||||
|
|
||||||
|
|
||||||
|
def _get_flashinfer_kda_kernel():
|
||||||
|
"""Lazy import for FlashInfer ``recurrent_kda`` (decode + MTP).
|
||||||
|
|
||||||
|
Returns (available, recurrent_kda_fn).
|
||||||
|
"""
|
||||||
|
global _flashinfer_kda_available, _flashinfer_recurrent_kda
|
||||||
|
if _flashinfer_kda_available is None:
|
||||||
|
try:
|
||||||
|
os.environ.setdefault("FLASHINFER_DISABLE_VERSION_CHECK", "1")
|
||||||
|
|
||||||
|
from flashinfer.kda_decode import recurrent_kda
|
||||||
|
|
||||||
|
_flashinfer_recurrent_kda = recurrent_kda
|
||||||
|
# recurrent_kda is SM100-only (CuTe DSL, Blackwell).
|
||||||
|
_flashinfer_kda_available = (
|
||||||
|
is_cuda() and torch.cuda.get_device_capability()[0] >= 10
|
||||||
|
)
|
||||||
|
if _flashinfer_kda_available:
|
||||||
|
logger.info("FlashInfer KDA kernel (recurrent_kda) loaded successfully")
|
||||||
|
except (ImportError, RuntimeError) as e:
|
||||||
|
logger.warning(f"FlashInfer KDA kernel not available: {e}")
|
||||||
|
_flashinfer_kda_available = False
|
||||||
|
_flashinfer_recurrent_kda = None
|
||||||
|
return _flashinfer_kda_available, _flashinfer_recurrent_kda
|
||||||
|
|
||||||
|
|
||||||
|
class FlashInferKDAKernel(LinearAttnKernelBase):
|
||||||
|
"""FlashInfer KDA kernel: SM100 decode + MTP (target_verify), topk=1.
|
||||||
|
|
||||||
|
Prefill (``extend``) is intentionally not implemented -- FlashInfer ships no
|
||||||
|
KDA chunk kernel; the dispatcher keeps prefill on Triton / CuTe DSL.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
available, self._recurrent_kda = _get_flashinfer_kda_kernel()
|
||||||
|
if not available or self._recurrent_kda is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"FlashInfer KDA kernel (recurrent_kda) is not available. "
|
||||||
|
"Requires SM100 (Blackwell) and a FlashInfer build with KDA support."
|
||||||
|
)
|
||||||
|
# Cache the per-layer constant gate-param prep (A_log/dt_bias reshape+cast),
|
||||||
|
# keyed by tensor identity. Layer params are persistent weights so id() is
|
||||||
|
# stable; this removes the per-call reshape/float/contiguous work.
|
||||||
|
self._gate_cache: dict = {}
|
||||||
|
# Cache the constant per-(row-map, batch, T) verify scatter indices
|
||||||
|
# (ssm_state_indices), which never change across verify calls.
|
||||||
|
self._verify_idx_cache: dict = {}
|
||||||
|
logger.info("Using FlashInfer KDA kernel")
|
||||||
|
|
||||||
|
# ---- gate / beta normalization (shared by decode + verify) ----
|
||||||
|
|
||||||
|
def _prep_gate_params(self, A_log: torch.Tensor, dt_bias: torch.Tensor):
|
||||||
|
# A_log: [1, 1, H, 1] -> [H] fp32; dt_bias: [H*K] (1D) -> fp32. Cached per
|
||||||
|
# layer (constant weights) so this is a dict lookup on the hot path.
|
||||||
|
key = (id(A_log), id(dt_bias))
|
||||||
|
cached = self._gate_cache.get(key)
|
||||||
|
if cached is not None:
|
||||||
|
return cached
|
||||||
|
A_log_fi = A_log.reshape(-1).float().contiguous()
|
||||||
|
dt_bias_fi = (
|
||||||
|
dt_bias.reshape(-1).float().contiguous() if dt_bias is not None else None
|
||||||
|
)
|
||||||
|
self._gate_cache[key] = (A_log_fi, dt_bias_fi)
|
||||||
|
return A_log_fi, dt_bias_fi
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _beta_logit_to_prob(b: torch.Tensor) -> torch.Tensor:
|
||||||
|
# Triton KDA does beta = sigmoid(b); recurrent_kda wants beta pre-sigmoided.
|
||||||
|
# torch.sigmoid computes in fp32 internally, so a single sigmoid on the bf16
|
||||||
|
# logit is enough (avoids an explicit fp32 upcast + downcast = 2 extra kernels).
|
||||||
|
return torch.sigmoid(b).to(torch.bfloat16)
|
||||||
|
|
||||||
|
# ---- decode ----
|
||||||
|
|
||||||
|
def decode(
|
||||||
|
self,
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
a: torch.Tensor,
|
||||||
|
b: torch.Tensor,
|
||||||
|
*,
|
||||||
|
A_log: torch.Tensor,
|
||||||
|
dt_bias: torch.Tensor,
|
||||||
|
ssm_states: torch.Tensor,
|
||||||
|
cache_indices: torch.Tensor,
|
||||||
|
query_start_loc: torch.Tensor,
|
||||||
|
**kwargs,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
batch_size = cache_indices.shape[0]
|
||||||
|
num_heads = q.shape[2]
|
||||||
|
head_k_dim = q.shape[3]
|
||||||
|
num_v_heads = v.shape[2]
|
||||||
|
head_v_dim = v.shape[3]
|
||||||
|
|
||||||
|
# Pack each request as a length-1 sequence ([1, B, ...] + cu_seqlens) so
|
||||||
|
# recurrent_kda indexes the committed pool IN-KERNEL via ssm_state_indices.
|
||||||
|
# The plain [B, 1, ...] path (no cu_seqlens) instead python-gathers
|
||||||
|
# initial_state[indices] and scatters it back with index_put around the
|
||||||
|
# kernel (~141us at B=64 in ncu); the cu_seqlens path skips both. q/k/v
|
||||||
|
# already arrive as [1, B, H, D] from forward_decode, so the reshape is a
|
||||||
|
# no-op view. recurrent_kda's cp.async + shared-mem staging are hardwired to
|
||||||
|
# bf16 (2-byte elements) for q/k/v/g/beta and the state, so every input is
|
||||||
|
# cast to bf16 -- a no-op for the common bf16 KDA model, a correct downcast
|
||||||
|
# otherwise (float16 bits would be reinterpreted as bf16 without the cast).
|
||||||
|
query_fi = q.reshape(1, batch_size, num_heads, head_k_dim).to(torch.bfloat16)
|
||||||
|
key_fi = k.reshape(1, batch_size, num_heads, head_k_dim).to(torch.bfloat16)
|
||||||
|
value_fi = v.reshape(1, batch_size, num_v_heads, head_v_dim).to(torch.bfloat16)
|
||||||
|
g_fi = a.reshape(1, batch_size, num_v_heads, head_k_dim).to(torch.bfloat16)
|
||||||
|
beta_fi = self._beta_logit_to_prob(b).reshape(1, batch_size, num_v_heads)
|
||||||
|
|
||||||
|
A_log_fi, dt_bias_fi = self._prep_gate_params(A_log, dt_bias)
|
||||||
|
|
||||||
|
# Softplus gate (lower_bound=None) to match the Triton KDA decode path;
|
||||||
|
# in-place state update into the committed pool (no rollback for decode).
|
||||||
|
# query_start_loc is the decode cu_seqlens (one token per request).
|
||||||
|
output_fi, _ = self._recurrent_kda(
|
||||||
|
q=query_fi,
|
||||||
|
k=key_fi,
|
||||||
|
v=value_fi,
|
||||||
|
g=g_fi,
|
||||||
|
beta=beta_fi,
|
||||||
|
A_log=A_log_fi,
|
||||||
|
dt_bias=dt_bias_fi,
|
||||||
|
scale=None,
|
||||||
|
initial_state=ssm_states,
|
||||||
|
output_final_state=False,
|
||||||
|
use_qk_l2norm_in_kernel=True,
|
||||||
|
use_gate_in_kernel=True,
|
||||||
|
lower_bound=None,
|
||||||
|
cu_seqlens=query_start_loc.to(torch.int32),
|
||||||
|
ssm_state_indices=cache_indices.to(torch.int32),
|
||||||
|
)
|
||||||
|
|
||||||
|
return output_fi.view(1, batch_size, num_v_heads, head_v_dim)
|
||||||
|
|
||||||
|
# ---- target_verify (MTP, topk=1) ----
|
||||||
|
|
||||||
|
def target_verify(
|
||||||
|
self,
|
||||||
|
A_log: torch.Tensor,
|
||||||
|
dt_bias: torch.Tensor,
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
a: torch.Tensor,
|
||||||
|
b: torch.Tensor,
|
||||||
|
*,
|
||||||
|
ssm_states: torch.Tensor,
|
||||||
|
cache_indices: torch.Tensor,
|
||||||
|
query_start_loc: torch.Tensor,
|
||||||
|
intermediate_states_buffer: torch.Tensor,
|
||||||
|
intermediate_state_indices: torch.Tensor,
|
||||||
|
cache_steps: int,
|
||||||
|
retrieve_parent_token: torch.Tensor,
|
||||||
|
**kwargs,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if retrieve_parent_token is not None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"FlashInfer KDA verify kernel only supports topk=1 "
|
||||||
|
"(retrieve_parent_token must be None)."
|
||||||
|
)
|
||||||
|
|
||||||
|
seq_len = q.shape[1]
|
||||||
|
batch_size = query_start_loc.shape[0] - 1
|
||||||
|
draft_token_num = cache_steps # T = 1 + num_spec_tokens
|
||||||
|
num_spec_tokens = draft_token_num - 1
|
||||||
|
num_heads = q.shape[2]
|
||||||
|
head_k_dim = q.shape[3]
|
||||||
|
num_v_heads = v.shape[2]
|
||||||
|
head_v_dim = v.shape[3]
|
||||||
|
|
||||||
|
# Packed [1, N*T, ...] inputs, cu_seqlens = query_start_loc (draft stride).
|
||||||
|
# recurrent_kda is bf16-only (see decode), so cast every input to bf16.
|
||||||
|
q_fi = q.reshape(1, seq_len, num_heads, head_k_dim).to(torch.bfloat16)
|
||||||
|
k_fi = k.reshape(1, seq_len, num_heads, head_k_dim).to(torch.bfloat16)
|
||||||
|
v_fi = v.reshape(1, seq_len, num_v_heads, head_v_dim).to(torch.bfloat16)
|
||||||
|
g_fi = a.reshape(1, seq_len, num_v_heads, head_k_dim).to(torch.bfloat16)
|
||||||
|
beta_fi = self._beta_logit_to_prob(b).reshape(1, seq_len, num_v_heads)
|
||||||
|
|
||||||
|
A_log_fi, dt_bias_fi = self._prep_gate_params(A_log, dt_bias)
|
||||||
|
|
||||||
|
# recurrent_kda indexes a flat state pool. Map each request/step to the
|
||||||
|
# matching slot in SGLang's [scratch_row, allocated_step, HV, V, K] buffer.
|
||||||
|
scratch = intermediate_states_buffer # [N_scratch, T, HV, V, K]
|
||||||
|
scratch_steps = scratch.shape[1]
|
||||||
|
if draft_token_num > scratch_steps:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"KDA verify needs {draft_token_num} scratch steps, "
|
||||||
|
f"but intermediate_ssm only has {scratch_steps}."
|
||||||
|
)
|
||||||
|
|
||||||
|
base_rows = intermediate_state_indices[:batch_size]
|
||||||
|
cache_key = (
|
||||||
|
id(intermediate_state_indices),
|
||||||
|
batch_size,
|
||||||
|
draft_token_num,
|
||||||
|
scratch_steps,
|
||||||
|
)
|
||||||
|
ssm_state_indices = self._verify_idx_cache.get(cache_key)
|
||||||
|
if ssm_state_indices is None:
|
||||||
|
# The fast seed copy below assumes row n in scratch belongs to request n.
|
||||||
|
expected = torch.arange(
|
||||||
|
batch_size, device=base_rows.device, dtype=base_rows.dtype
|
||||||
|
)
|
||||||
|
if not torch.equal(base_rows, expected):
|
||||||
|
raise RuntimeError(
|
||||||
|
"FlashInfer KDA verify requires an identity intermediate row-map "
|
||||||
|
"(verify_intermediate_state_indices must be arange)."
|
||||||
|
)
|
||||||
|
step = torch.arange(draft_token_num, device=q.device, dtype=torch.int32)
|
||||||
|
ssm_state_indices = (
|
||||||
|
base_rows.to(torch.int32)[:, None] * scratch_steps + step[None, :]
|
||||||
|
).contiguous() # [N, T]
|
||||||
|
self._verify_idx_cache[cache_key] = ssm_state_indices
|
||||||
|
|
||||||
|
# Seed step 0 from committed state, then recurrent_kda overwrites it with
|
||||||
|
# token-0 post-state. Padded graph rows clamp to slot 0; their output is ignored.
|
||||||
|
base_state = ssm_states.index_select(
|
||||||
|
0, cache_indices[:batch_size].clamp(min=0).to(torch.int64)
|
||||||
|
)
|
||||||
|
scratch[:batch_size, 0].copy_(base_state)
|
||||||
|
|
||||||
|
# Same storage as scratch, flattened over the allocated step stride.
|
||||||
|
state_pool = scratch.view(
|
||||||
|
scratch.shape[0] * scratch_steps, num_v_heads, head_v_dim, head_k_dim
|
||||||
|
)
|
||||||
|
|
||||||
|
output_fi, _ = self._recurrent_kda(
|
||||||
|
q=q_fi,
|
||||||
|
k=k_fi,
|
||||||
|
v=v_fi,
|
||||||
|
g=g_fi,
|
||||||
|
beta=beta_fi,
|
||||||
|
A_log=A_log_fi,
|
||||||
|
dt_bias=dt_bias_fi,
|
||||||
|
scale=None,
|
||||||
|
initial_state=state_pool,
|
||||||
|
output_final_state=False,
|
||||||
|
use_qk_l2norm_in_kernel=True,
|
||||||
|
use_gate_in_kernel=True,
|
||||||
|
lower_bound=None,
|
||||||
|
cu_seqlens=query_start_loc.to(torch.int32),
|
||||||
|
ssm_state_indices=ssm_state_indices,
|
||||||
|
num_spec_tokens=num_spec_tokens,
|
||||||
|
)
|
||||||
|
|
||||||
|
return output_fi.view(1, seq_len, num_v_heads, head_v_dim)
|
||||||
|
|
||||||
|
# ---- extend (prefill): not provided by FlashInfer ----
|
||||||
|
|
||||||
|
def extend(self, *args, **kwargs):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"FlashInferKDAKernel has no prefill kernel; keep prefill on Triton / CuTe DSL."
|
||||||
|
)
|
||||||
@@ -141,6 +141,53 @@ class TritonKDAKernel(LinearAttnKernelBase):
|
|||||||
is_kda=True,
|
is_kda=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def target_verify(
|
||||||
|
self,
|
||||||
|
A_log: torch.Tensor,
|
||||||
|
dt_bias: torch.Tensor,
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
a: torch.Tensor,
|
||||||
|
b: torch.Tensor,
|
||||||
|
*,
|
||||||
|
ssm_states: torch.Tensor,
|
||||||
|
cache_indices: torch.Tensor,
|
||||||
|
query_start_loc: torch.Tensor,
|
||||||
|
intermediate_states_buffer: torch.Tensor,
|
||||||
|
intermediate_state_indices: torch.Tensor,
|
||||||
|
cache_steps: int,
|
||||||
|
retrieve_parent_token: torch.Tensor,
|
||||||
|
**kwargs,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
# KDA MTP / speculative-decode verify via the fused KDA kernel (IS_KDA=True),
|
||||||
|
# mirroring the GDN triton verify path. Reads the committed state, writes
|
||||||
|
# per-draft-token intermediate states to the scratch buffer, does NOT mutate
|
||||||
|
# the committed pool (disable_state_update=True), and handles chain + tree
|
||||||
|
# (retrieve_parent_token). The verify kernel for the Triton / CuTe DSL KDA
|
||||||
|
# decode backends, and the reference the KDA correctness tests assert against.
|
||||||
|
return fused_sigmoid_gating_delta_rule_update(
|
||||||
|
A_log=A_log,
|
||||||
|
dt_bias=dt_bias,
|
||||||
|
q=q,
|
||||||
|
k=k,
|
||||||
|
v=v,
|
||||||
|
a=a,
|
||||||
|
b=b,
|
||||||
|
initial_state_source=ssm_states,
|
||||||
|
initial_state_indices=cache_indices,
|
||||||
|
cu_seqlens=query_start_loc,
|
||||||
|
use_qk_l2norm_in_kernel=True,
|
||||||
|
softplus_beta=1.0,
|
||||||
|
softplus_threshold=20.0,
|
||||||
|
is_kda=True,
|
||||||
|
disable_state_update=True,
|
||||||
|
intermediate_states_buffer=intermediate_states_buffer,
|
||||||
|
intermediate_state_indices=intermediate_state_indices,
|
||||||
|
cache_steps=cache_steps,
|
||||||
|
retrieve_parent_token=retrieve_parent_token,
|
||||||
|
)
|
||||||
|
|
||||||
def extend(
|
def extend(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
|
|||||||
@@ -102,21 +102,19 @@ _use_aiter = bool(envs.SGLANG_USE_AITER.get()) and _is_hip
|
|||||||
|
|
||||||
|
|
||||||
def conv_window_dedup_enabled(
|
def conv_window_dedup_enabled(
|
||||||
is_npu: bool, is_cpu: bool, speculative_eagle_topk: Optional[int]
|
is_npu: bool, is_cpu: bool, speculative_eagle_topk: Optional[int], is_kda: bool
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Whether the deduplicated sliding-window conv-intermediate layout is safe.
|
"""Whether the deduplicated sliding-window conv-intermediate layout is safe.
|
||||||
|
|
||||||
It is only correct for a *linear* draft chain (``speculative_eagle_topk <= 1``,
|
It is safe for CUDA linear draft chains whose kernels consume the window raw.
|
||||||
i.e. NEXTN / MTP): consecutive draft tokens then form a true sliding window, so
|
Tree verify, NPU/CPU, and KDA keep dense windows: tree ancestors need independent
|
||||||
the overlapping physical columns hold identical values. Under EAGLE *tree*
|
windows, platform kernels expect contiguous steps, and KDA transposes the window
|
||||||
verify (``topk > 1``) the conv kernel walks per-token tree ancestors, so aliased
|
before conv so the overlapping ``as_strided`` layout would corrupt stores.
|
||||||
columns can need different values from different parent chains -> fall back to
|
|
||||||
the dense layout. NPU/CPU also keep the dense layout (their kernels assume
|
|
||||||
contiguous per-step windows). See ``MambaPool.__init__``.
|
|
||||||
"""
|
"""
|
||||||
return (
|
return (
|
||||||
not is_npu
|
not is_npu
|
||||||
and not is_cpu
|
and not is_cpu
|
||||||
|
and not is_kda
|
||||||
and (speculative_eagle_topk is None or speculative_eagle_topk <= 1)
|
and (speculative_eagle_topk is None or speculative_eagle_topk <= 1)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -576,7 +574,7 @@ class MambaPool:
|
|||||||
# `fused_conv_window_scatter_with_mask` scatter is layout-agnostic,
|
# `fused_conv_window_scatter_with_mask` scatter is layout-agnostic,
|
||||||
# so the dense fallback reads correctly through the same code path.
|
# so the dense fallback reads correctly through the same code path.
|
||||||
dedup_conv_window = conv_window_dedup_enabled(
|
dedup_conv_window = conv_window_dedup_enabled(
|
||||||
_is_npu, _is_cpu, speculative_eagle_topk
|
_is_npu, _is_cpu, speculative_eagle_topk, cache_params.is_kda
|
||||||
)
|
)
|
||||||
self._intermediate_conv_window_phys = []
|
self._intermediate_conv_window_phys = []
|
||||||
if dedup_conv_window:
|
if dedup_conv_window:
|
||||||
|
|||||||
@@ -0,0 +1,286 @@
|
|||||||
|
"""Correctness tests for the FlashInfer SM100 KDA decode + MTP backend.
|
||||||
|
|
||||||
|
Compares ``FlashInferKDAKernel`` with the Triton KDA reference for decode output,
|
||||||
|
state updates, and topk=1 target_verify checkpoints. ``recurrent_kda`` is
|
||||||
|
SM100-only and requires a FlashInfer build that exposes it.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
# SM100 single-GPU kernel-unit suite, same slot as the CuteDSL KDA prefill test.
|
||||||
|
# Disabled in public CI until the B200 runner image ships recurrent_kda.
|
||||||
|
register_cuda_ci(
|
||||||
|
est_time=60,
|
||||||
|
stage="base-b-kernel-unit",
|
||||||
|
runner_config="4-gpu-b200",
|
||||||
|
disabled="recurrent_kda (SM100 KDA decode) not guaranteed in public CI FlashInfer build",
|
||||||
|
)
|
||||||
|
|
||||||
|
if not (torch.cuda.is_available() and torch.cuda.get_device_capability()[0] >= 10):
|
||||||
|
pytest.skip(
|
||||||
|
"FlashInfer KDA (recurrent_kda) requires CUDA SM10x (Blackwell).",
|
||||||
|
allow_module_level=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
from sglang.srt.layers.attention.linear.kernels.kda_flashinfer import ( # noqa: E402
|
||||||
|
FlashInferKDAKernel,
|
||||||
|
_get_flashinfer_kda_kernel,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.attention.linear.kernels.kda_triton import ( # noqa: E402
|
||||||
|
TritonKDAKernel,
|
||||||
|
)
|
||||||
|
|
||||||
|
_available, _ = _get_flashinfer_kda_kernel()
|
||||||
|
if not _available:
|
||||||
|
pytest.skip(
|
||||||
|
"FlashInfer build does not expose recurrent_kda (KDA decode).",
|
||||||
|
allow_module_level=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# KDA: head_k_dim == head_v_dim == 128; single q/v head group (HV == H) here.
|
||||||
|
H, HV, K, V = 16, 16, 128, 128
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Inputs (matched to the sglang KDA decode/verify contract: raw per-K gate `a`,
|
||||||
|
# beta logit `b`, SSM pool [N, HV, V, K], decode cu_seqlens = query_start_loc).
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
def _make_decode_inputs(batch_size, device="cuda", dtype=torch.bfloat16):
|
||||||
|
B, pool = batch_size, batch_size + 16
|
||||||
|
return dict(
|
||||||
|
B=B,
|
||||||
|
q=(torch.randn(1, B, H, K, device=device, dtype=dtype) * 0.5).contiguous(),
|
||||||
|
k=(torch.randn(1, B, H, K, device=device, dtype=dtype) * 0.5).contiguous(),
|
||||||
|
v=(torch.randn(1, B, HV, V, device=device, dtype=dtype) * 0.5).contiguous(),
|
||||||
|
a=(torch.randn(B, HV * K, device=device, dtype=dtype) * 0.5 - 1.0).contiguous(),
|
||||||
|
b=(torch.randn(B, HV, device=device, dtype=dtype) * 0.5).contiguous(),
|
||||||
|
A_log=torch.randn(HV, device=device, dtype=torch.float32) * 0.2,
|
||||||
|
dt_bias=torch.randn(HV * K, device=device, dtype=torch.float32) * 0.1,
|
||||||
|
ssm=(
|
||||||
|
torch.randn(pool, HV, V, K, device=device, dtype=dtype) * 0.01
|
||||||
|
).contiguous(),
|
||||||
|
cache_indices=torch.arange(B, device=device, dtype=torch.int32),
|
||||||
|
qsl=torch.arange(B + 1, device=device, dtype=torch.int32),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_verify_inputs(
|
||||||
|
batch_size,
|
||||||
|
cache_steps,
|
||||||
|
allocated_steps=None,
|
||||||
|
device="cuda",
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
):
|
||||||
|
B, T = batch_size, cache_steps
|
||||||
|
S = allocated_steps or T
|
||||||
|
assert S >= T
|
||||||
|
seq, pool = B * T, B + 16
|
||||||
|
return dict(
|
||||||
|
B=B,
|
||||||
|
T=T,
|
||||||
|
allocated_steps=S,
|
||||||
|
seq=seq,
|
||||||
|
q=(torch.randn(1, seq, H, K, device=device, dtype=dtype) * 0.5).contiguous(),
|
||||||
|
k=(torch.randn(1, seq, H, K, device=device, dtype=dtype) * 0.5).contiguous(),
|
||||||
|
v=(torch.randn(1, seq, HV, V, device=device, dtype=dtype) * 0.5).contiguous(),
|
||||||
|
a=(
|
||||||
|
torch.randn(seq, HV * K, device=device, dtype=dtype) * 0.5 - 1.0
|
||||||
|
).contiguous(),
|
||||||
|
b=(torch.randn(seq, HV, device=device, dtype=dtype) * 0.5).contiguous(),
|
||||||
|
A_log=torch.randn(HV, device=device, dtype=torch.float32) * 0.2,
|
||||||
|
dt_bias=torch.randn(HV * K, device=device, dtype=torch.float32) * 0.1,
|
||||||
|
ssm=(
|
||||||
|
torch.randn(pool, HV, V, K, device=device, dtype=dtype) * 0.01
|
||||||
|
).contiguous(),
|
||||||
|
cache_indices=torch.arange(B, device=device, dtype=torch.int32),
|
||||||
|
qsl=torch.arange(0, seq + 1, T, device=device, dtype=torch.int32),
|
||||||
|
intermediate_states=torch.zeros(
|
||||||
|
B, S, HV, V, K, device=device, dtype=dtype
|
||||||
|
).contiguous(),
|
||||||
|
intermediate_indices=torch.arange(B, device=device, dtype=torch.int32),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _decode(kern, d, ssm):
|
||||||
|
# `ssm` is updated in place (committed-pool decode step); pass a fresh clone.
|
||||||
|
return kern.decode(
|
||||||
|
d["q"],
|
||||||
|
d["k"],
|
||||||
|
d["v"],
|
||||||
|
d["a"],
|
||||||
|
d["b"],
|
||||||
|
A_log=d["A_log"],
|
||||||
|
dt_bias=d["dt_bias"],
|
||||||
|
ssm_states=ssm,
|
||||||
|
cache_indices=d["cache_indices"],
|
||||||
|
query_start_loc=d["qsl"],
|
||||||
|
).reshape(d["B"], HV, V)
|
||||||
|
|
||||||
|
|
||||||
|
def _verify(kern, d, ssm, intermediate_states):
|
||||||
|
return kern.target_verify(
|
||||||
|
A_log=d["A_log"],
|
||||||
|
dt_bias=d["dt_bias"],
|
||||||
|
q=d["q"],
|
||||||
|
k=d["k"],
|
||||||
|
v=d["v"],
|
||||||
|
a=d["a"],
|
||||||
|
b=d["b"],
|
||||||
|
ssm_states=ssm,
|
||||||
|
cache_indices=d["cache_indices"],
|
||||||
|
query_start_loc=d["qsl"],
|
||||||
|
intermediate_states_buffer=intermediate_states,
|
||||||
|
intermediate_state_indices=d["intermediate_indices"],
|
||||||
|
cache_steps=d["T"],
|
||||||
|
retrieve_parent_token=None,
|
||||||
|
).reshape(d["seq"], HV, V)
|
||||||
|
|
||||||
|
|
||||||
|
def _sequential_decode_states(kern, d):
|
||||||
|
"""Ground truth for verify checkpoints: single-token decode over each step."""
|
||||||
|
B, T = d["B"], d["T"]
|
||||||
|
st = d["ssm"].clone() # committed pool [pool, HV, V, K], updated in place by decode
|
||||||
|
ci = d["cache_indices"].long()
|
||||||
|
qsl_dec = torch.arange(B + 1, device=st.device, dtype=torch.int32)
|
||||||
|
ref = torch.zeros(B, T, HV, V, K, device=st.device, dtype=st.dtype)
|
||||||
|
for t in range(T):
|
||||||
|
pos = torch.arange(B, device=st.device) * T + t # token t of each request
|
||||||
|
kern.decode(
|
||||||
|
d["q"][:, pos].contiguous(),
|
||||||
|
d["k"][:, pos].contiguous(),
|
||||||
|
d["v"][:, pos].contiguous(),
|
||||||
|
d["a"][pos].contiguous(),
|
||||||
|
d["b"][pos].contiguous(),
|
||||||
|
A_log=d["A_log"],
|
||||||
|
dt_bias=d["dt_bias"],
|
||||||
|
ssm_states=st,
|
||||||
|
cache_indices=d["cache_indices"],
|
||||||
|
query_start_loc=qsl_dec,
|
||||||
|
)
|
||||||
|
ref[:, t] = st[ci] # post-token-t state for each request
|
||||||
|
return ref
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Tests
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
@pytest.mark.parametrize("batch_size", [1, 8, 64, 128])
|
||||||
|
def test_kda_decode_flashinfer_matches_triton(batch_size):
|
||||||
|
"""FlashInfer decode output + committed-pool state update match the Triton
|
||||||
|
KDA decode reference."""
|
||||||
|
torch.manual_seed(batch_size)
|
||||||
|
d = _make_decode_inputs(batch_size)
|
||||||
|
fi, tri = FlashInferKDAKernel(), TritonKDAKernel()
|
||||||
|
|
||||||
|
st_ref = d["ssm"].clone()
|
||||||
|
ref_out = _decode(tri, d, st_ref).float()
|
||||||
|
st_fi = d["ssm"].clone()
|
||||||
|
out = _decode(fi, d, st_fi).float()
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
assert torch.isfinite(out).all(), "FlashInfer decode output has non-finite values"
|
||||||
|
assert torch.isfinite(st_fi).all(), "FlashInfer decode state has non-finite values"
|
||||||
|
|
||||||
|
o_err = (out - ref_out).abs()
|
||||||
|
# bf16 recurrent step; B200 kernel-unit measured out max-abs-diff ~1e-4.
|
||||||
|
assert o_err.max().item() < 1e-2, f"decode out max diff {o_err.max().item():.2e}"
|
||||||
|
assert o_err.mean().item() < 1e-3, f"decode out mean diff {o_err.mean().item():.2e}"
|
||||||
|
|
||||||
|
# Updated committed-pool slots (SSM state [HV, V, K]) must match too.
|
||||||
|
idx = d["cache_indices"].long()
|
||||||
|
s_err = (st_fi[idx].float() - st_ref[idx].float()).abs()
|
||||||
|
assert s_err.max().item() < 1e-1, f"decode state max diff {s_err.max().item():.2e}"
|
||||||
|
assert (
|
||||||
|
s_err.mean().item() < 1e-2
|
||||||
|
), f"decode state mean diff {s_err.mean().item():.2e}"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("batch_size,num_spec", [(1, 7), (8, 7), (32, 3)])
|
||||||
|
def test_kda_target_verify_flashinfer_matches_triton(batch_size, num_spec):
|
||||||
|
"""FlashInfer MTP / target_verify (topk=1) per-draft-token output matches the
|
||||||
|
Triton KDA verify reference over T = 1 + num_spec draft tokens per sequence."""
|
||||||
|
torch.manual_seed(batch_size + num_spec)
|
||||||
|
d = _make_verify_inputs(batch_size, 1 + num_spec)
|
||||||
|
fi, tri = FlashInferKDAKernel(), TritonKDAKernel()
|
||||||
|
|
||||||
|
ref_out = _verify(
|
||||||
|
tri, d, d["ssm"].clone(), d["intermediate_states"].clone()
|
||||||
|
).float()
|
||||||
|
out = _verify(fi, d, d["ssm"].clone(), d["intermediate_states"].clone()).float()
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
assert torch.isfinite(out).all(), "FlashInfer verify output has non-finite values"
|
||||||
|
o_err = (out - ref_out).abs()
|
||||||
|
# B200 kernel-unit measured verify out max-abs-diff ~2e-4.
|
||||||
|
assert o_err.max().item() < 1e-2, f"verify out max diff {o_err.max().item():.2e}"
|
||||||
|
assert o_err.mean().item() < 1e-3, f"verify out mean diff {o_err.mean().item():.2e}"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"batch_size,num_spec,extra_steps",
|
||||||
|
[(1, 7, 0), (8, 7, 0), (32, 3, 2)],
|
||||||
|
)
|
||||||
|
def test_kda_target_verify_flashinfer_checkpoint_states(
|
||||||
|
batch_size, num_spec, extra_steps
|
||||||
|
):
|
||||||
|
"""Checkpoint states must match true sequential decode states."""
|
||||||
|
torch.manual_seed(1000 + batch_size + num_spec)
|
||||||
|
cache_steps = 1 + num_spec
|
||||||
|
d = _make_verify_inputs(
|
||||||
|
batch_size,
|
||||||
|
cache_steps,
|
||||||
|
allocated_steps=cache_steps + extra_steps,
|
||||||
|
)
|
||||||
|
fi = FlashInferKDAKernel()
|
||||||
|
|
||||||
|
ref_states = _sequential_decode_states(fi, d).float()
|
||||||
|
|
||||||
|
intermediate_states = d["intermediate_states"].clone()
|
||||||
|
_verify(
|
||||||
|
fi, d, d["ssm"].clone(), intermediate_states
|
||||||
|
) # fills intermediate_states[n, t] in place
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
got = intermediate_states[:, : d["T"]].float() # [B, T, HV, V, K] checkpoint states
|
||||||
|
assert torch.isfinite(got).all(), "verify checkpoint states have non-finite values"
|
||||||
|
s_err = (got - ref_states).abs()
|
||||||
|
# bf16 recurrent state; same tolerance as the decode committed-state check.
|
||||||
|
assert (
|
||||||
|
s_err.max().item() < 1e-1
|
||||||
|
), f"checkpoint state max diff {s_err.max().item():.2e}"
|
||||||
|
assert (
|
||||||
|
s_err.mean().item() < 1e-2
|
||||||
|
), f"checkpoint state mean diff {s_err.mean().item():.2e}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_kda_target_verify_flashinfer_rejects_tree_spec():
|
||||||
|
"""Tree speculation (retrieve_parent_token != None) is unsupported (topk=1
|
||||||
|
linear chain only) and must raise, not silently miscompute."""
|
||||||
|
d = _make_verify_inputs(2, 4)
|
||||||
|
parent = torch.zeros(d["seq"], device="cuda", dtype=torch.int32)
|
||||||
|
with pytest.raises(RuntimeError, match="topk=1"):
|
||||||
|
FlashInferKDAKernel().target_verify(
|
||||||
|
A_log=d["A_log"],
|
||||||
|
dt_bias=d["dt_bias"],
|
||||||
|
q=d["q"],
|
||||||
|
k=d["k"],
|
||||||
|
v=d["v"],
|
||||||
|
a=d["a"],
|
||||||
|
b=d["b"],
|
||||||
|
ssm_states=d["ssm"].clone(),
|
||||||
|
cache_indices=d["cache_indices"],
|
||||||
|
query_start_loc=d["qsl"],
|
||||||
|
intermediate_states_buffer=d["intermediate_states"].clone(),
|
||||||
|
intermediate_state_indices=d["intermediate_indices"],
|
||||||
|
cache_steps=d["T"],
|
||||||
|
retrieve_parent_token=parent,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
import sys
|
||||||
|
|
||||||
|
sys.exit(pytest.main([__file__, "-v"]))
|
||||||
Reference in New Issue
Block a user