[DeepSeek-V4] Add Q8KV8 sparse MLA prefill runtime backend (#32327)

Co-authored-by: Ho-Ren (Jack) Chuang <horenchuang@bytedance.com>
Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
shiyang814-cpu
2026-08-20 10:23:49 +08:00
committed by GitHub
co-authored by Ho-Ren Chuang Xiaoyu Zhang
parent a49560ce50
commit 9db4ba8da1
7 changed files with 1337 additions and 13 deletions
@@ -85,6 +85,140 @@ def dequantize_k_cache_paged(
return out return out
def gather_dequant_requant_fp8_paged(
quant_k_cache: torch.Tensor,
page_table_1_flattened: torch.Tensor,
page_size: int,
extra_rows: int = 0,
out: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Gather DeepSeek-V4 paged KV cache into a flat FP8 workspace.
This is the Q8KV8 sparse-prefill adapter for the DeepSeek-V4 packed layout.
It gathers token IDs from the existing paged cache, dequantizes the 448-dim
nope region with its UE8M0 per-64 scales, casts the 64-dim BF16 rope tail to
FP8, and writes the result as ``(num_tokens + extra_rows, 1, 512)`` FP8.
``extra_rows`` appends zero rows for kernels that map masked sparse indices
to a valid zero landing pad.
"""
assert quant_k_cache.is_contiguous()
assert page_table_1_flattened.dtype in (torch.int32, torch.int64)
assert extra_rows >= 0
quant_k_cache_u8 = quant_k_cache.view(torch.uint8)
num_tokens = page_table_1_flattened.shape[0]
total_rows = num_tokens + extra_rows
bytes_per_page = quant_k_cache_u8.shape[-1]
s_offset_bytes = page_size * NOPE_ROPE_BYTES
buf_fp8 = quant_k_cache_u8.view(fp8_dtype).reshape(-1)
buf_bf16 = quant_k_cache_u8.view(torch.bfloat16).reshape(-1)
buf_uint8 = quant_k_cache_u8.reshape(-1)
if out is None:
out = torch.zeros(
(total_rows, 1, DIM_NOPE + DIM_ROPE),
dtype=fp8_dtype,
device=quant_k_cache.device,
)
else:
assert out.shape == (total_rows, 1, DIM_NOPE + DIM_ROPE)
assert out.dtype == fp8_dtype
if extra_rows:
out[num_tokens:].zero_()
if num_tokens == 0:
return out
_gather_dequant_requant_fp8_paged_kernel[(num_tokens,)](
out,
buf_fp8,
buf_bf16,
buf_uint8,
page_table_1_flattened,
out.stride(0),
BYTES_PER_PAGE=bytes_per_page,
PAGE_SIZE=page_size,
DIM_NOPE=DIM_NOPE,
DIM_ROPE=DIM_ROPE,
TILE_SIZE=TILE_SIZE,
NUM_SCALE_TILES=NUM_SCALE_TILES,
NOPE_ROPE_BYTES=NOPE_ROPE_BYTES,
PADDED_SCALE_PER_TOKEN=PADDED_SCALE_PER_TOKEN,
S_OFFSET_BYTES=s_offset_bytes,
)
return out
def q8kv8_padded_num_heads(num_heads: int) -> int:
"""Return a Q-head count supported by the SM90 Q8KV8 kernel."""
if num_heads <= 0:
raise ValueError(f"num_heads must be positive, got {num_heads}")
if num_heads <= 64:
return 64
if num_heads <= 128:
return 128
raise ValueError(
"DeepSeek-V4 Q8KV8 sparse prefill supports at most 128 local "
f"query heads, got {num_heads}"
)
def cast_q_fp8_for_q8kv8_prefill(
q: torch.Tensor,
padded_num_heads: Optional[int] = None,
out: Optional[torch.Tensor] = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Cast DeepSeek-V4 sparse-prefill Q to the Q8KV8 kernel format.
The incoming Q is the model-produced BF16/FP16 tensor already shaped as
``(num_tokens, num_heads, 512)`` after removing the singleton MQA axis.
The SM90 kernel processes query heads in 64-head blocks. Tensor parallelism
commonly leaves fewer than 64 local heads, so the active heads are copied
into a zero-padded 64/128-head FP8 tensor.
"""
assert q.ndim == 3
assert q.shape[-1] == DIM_NOPE + DIM_ROPE
num_tokens, num_heads, head_dim = q.shape
if padded_num_heads is None:
padded_num_heads = q8kv8_padded_num_heads(num_heads)
if padded_num_heads not in (64, 128) or padded_num_heads < num_heads:
raise ValueError(
f"invalid padded_num_heads={padded_num_heads} for num_heads={num_heads}"
)
expected_shape = (num_tokens, padded_num_heads, head_dim)
if out is None:
q_fp8 = torch.zeros(
expected_shape,
dtype=fp8_dtype,
device=q.device,
)
else:
if (
out.shape != expected_shape
or out.dtype != fp8_dtype
or out.device != q.device
):
raise ValueError(
"Q8KV8 Q output must have shape/dtype/device "
f"{expected_shape}/{fp8_dtype}/{q.device}, got "
f"{tuple(out.shape)}/{out.dtype}/{out.device}"
)
q_fp8 = out
if padded_num_heads > num_heads:
q_fp8[:, num_heads:].zero_()
q_fp8[:, :num_heads].copy_(q)
q_scale = torch.ones((), dtype=torch.float32, device=q.device)
return q_fp8, q_scale
@triton.jit @triton.jit
def _dequantize_k_cache_paged_kernel( def _dequantize_k_cache_paged_kernel(
output_ptr, output_ptr,
@@ -136,6 +270,58 @@ def _dequantize_k_cache_paged_kernel(
tl.store(output_ptr + out_row_base + DIM_NOPE + rope_offs, rope_data) tl.store(output_ptr + out_row_base + DIM_NOPE + rope_offs, rope_data)
@triton.jit
def _gather_dequant_requant_fp8_paged_kernel(
output_ptr,
buf_fp8_ptr,
buf_bf16_ptr,
buf_uint8_ptr,
page_table_ptr,
output_stride_0,
BYTES_PER_PAGE: tl.constexpr,
PAGE_SIZE: tl.constexpr,
DIM_NOPE: tl.constexpr,
DIM_ROPE: tl.constexpr,
TILE_SIZE: tl.constexpr,
NUM_SCALE_TILES: tl.constexpr,
NOPE_ROPE_BYTES: tl.constexpr,
PADDED_SCALE_PER_TOKEN: tl.constexpr,
S_OFFSET_BYTES: tl.constexpr,
):
token_id = tl.program_id(0)
loc = tl.load(page_table_ptr + token_id).to(tl.int64)
page_idx = loc // PAGE_SIZE
in_page = loc % PAGE_SIZE
page_byte_base = page_idx * BYTES_PER_PAGE
token_data_base = page_byte_base + in_page * NOPE_ROPE_BYTES
token_scale_base = (
page_byte_base + S_OFFSET_BYTES + in_page * PADDED_SCALE_PER_TOKEN
)
out_row_base = token_id * output_stride_0
nope_offs = tl.arange(0, TILE_SIZE)
for tile_id in tl.static_range(NUM_SCALE_TILES):
fp8_off = token_data_base + tile_id * TILE_SIZE + nope_offs
fp8_vals = tl.load(buf_fp8_ptr + fp8_off).to(tl.float32)
scale_u8 = tl.load(buf_uint8_ptr + token_scale_base + tile_id).to(tl.int32)
scale_pow2 = tl.exp2((scale_u8 - 127).to(tl.float32))
out_off = out_row_base + tile_id * TILE_SIZE + nope_offs
tl.store(
output_ptr + out_off,
(fp8_vals * scale_pow2).to(output_ptr.dtype.element_ty),
)
rope_offs = tl.arange(0, DIM_ROPE)
bf16_off = (token_data_base + DIM_NOPE) // 2 + rope_offs
rope_data = tl.load(buf_bf16_ptr + bf16_off)
tl.store(
output_ptr + out_row_base + DIM_NOPE + rope_offs,
rope_data.to(output_ptr.dtype.element_ty),
)
def dequantize_k_cache_paged_ref( def dequantize_k_cache_paged_ref(
quant_k_cache: torch.Tensor, quant_k_cache: torch.Tensor,
page_table_1_flattened: torch.Tensor, page_table_1_flattened: torch.Tensor,
@@ -196,6 +382,29 @@ def dequantize_k_cache_paged_ref(
return out return out
def gather_dequant_requant_fp8_paged_ref(
quant_k_cache: torch.Tensor,
page_table_1_flattened: torch.Tensor,
page_size: int,
extra_rows: int = 0,
) -> torch.Tensor:
"""Torch reference for :func:`gather_dequant_requant_fp8_paged`."""
active = dequantize_k_cache_paged_ref(
quant_k_cache,
page_table_1_flattened,
page_size,
).to(fp8_dtype)
if extra_rows == 0:
return active
out = torch.zeros(
(active.shape[0] + extra_rows, 1, DIM_NOPE + DIM_ROPE),
dtype=fp8_dtype,
device=active.device,
)
out[: active.shape[0]] = active
return out
if __name__ == "__main__": if __name__ == "__main__":
assert torch.cuda.is_available(), "this self-test needs a CUDA device" assert torch.cuda.is_available(), "this self-test needs a CUDA device"
torch.manual_seed(0) torch.manual_seed(0)
@@ -284,20 +284,116 @@ def sparse_mla_q8kv8_prefill_fwd(
"""Run Q8KV8 (FP8) sparse prefill attention on SM90. """Run Q8KV8 (FP8) sparse prefill attention on SM90.
The kernel writes into three output tensors. By default fresh tensors The kernel writes into three output tensors. By default fresh tensors
are allocated and returned; callers that want to reuse buffers (e.g. are allocated and returned; callers that want to reuse buffers may pass
for CUDA graph capture) may pass pre-allocated ``out`` / ``max_logits`` pre-allocated ``out`` / ``max_logits`` / ``lse`` tensors of the expected
/ ``lse`` tensors of the expected shape/dtype/device. The three output shape/dtype/device. The three output tensors must not alias each other.
tensors must not alias each other.
Returns: Returns:
out: [s_q, h_q, d_v], bfloat16 out: [s_q, h_q, d_v], bfloat16
max_logits: [s_q, h_q], float32 max_logits: [s_q, h_q], float32
lse: [s_q, h_q], float32 lse: [s_q, h_q], float32
""" """
# Validate ranks before unpacking shapes so malformed callers fail with a
# clear error instead of a Python unpacking/indexing exception.
if q.ndim != 3:
raise ValueError(f"q must have shape (s_q, h_q, d_qk), got {tuple(q.shape)}")
if kv.ndim != 3:
raise ValueError(
f"kv must have shape (s_kv, h_kv, d_qk), got {tuple(kv.shape)}"
)
if indices.ndim != 3:
raise ValueError(
"indices must have shape (s_q, h_kv, topk), " f"got {tuple(indices.shape)}"
)
s_q, h_q, d_qk = q.shape s_q, h_q, d_qk = q.shape
s_kv = kv.shape[0] s_kv, h_kv, kv_d_qk = kv.shape
h_kv = kv.shape[1]
topk = indices.shape[2] topk = indices.shape[2]
device = q.device
# entry.cuh interprets q/kv as contiguous FP8 buffers and launches all
# accesses on q's CUDA device. Reject contract violations before launch.
if not q.is_cuda:
raise ValueError("q must be a CUDA tensor")
if not kv.is_cuda:
raise ValueError("kv must be a CUDA tensor")
if not indices.is_cuda:
raise ValueError("indices must be a CUDA tensor")
if kv.device != device:
raise ValueError(f"kv must be on q's device {device}, got {kv.device}")
if indices.device != device:
raise ValueError(
f"indices must be on q's device {device}, got {indices.device}"
)
if q.dtype != torch.float8_e4m3fn:
raise ValueError(f"q must be torch.float8_e4m3fn, got {q.dtype}")
if kv.dtype != torch.float8_e4m3fn:
raise ValueError(f"kv must be torch.float8_e4m3fn, got {kv.dtype}")
if not q.is_contiguous():
raise ValueError("q must be contiguous")
if not kv.is_contiguous():
raise ValueError("kv must be contiguous")
if not indices.is_contiguous():
raise ValueError("indices must be contiguous")
if kv_d_qk != d_qk:
raise ValueError(f"kv d_qk must match q d_qk={d_qk}, got {kv_d_qk}")
# The CUDA implementation uses B_H=64 and launches h_q / B_H CTAs.
# Reject unpadded TP-local head counts instead of launching zero CTAs and
# returning uninitialized outputs, which can appear to callers as a hang or
# a later collective failure.
if h_q == 0 or h_q % 64 != 0:
raise ValueError(
"sparse_mla_q8kv8_prefill_fwd requires h_q padded to a positive "
f"multiple of 64, got {h_q}"
)
if h_kv != 1:
raise ValueError(f"sparse_mla_q8kv8_prefill_fwd requires h_kv=1, got {h_kv}")
if d_qk not in (512, 576):
raise ValueError(
f"sparse_mla_q8kv8_prefill_fwd supports d_qk=512/576, got {d_qk}"
)
if indices.shape[:2] != (s_q, h_kv):
raise ValueError(
"indices must have shape "
f"({s_q}, {h_kv}, topk), got {tuple(indices.shape)}"
)
if indices.dtype != torch.int32:
raise ValueError(f"indices must be int32, got {indices.dtype}")
if topk == 0 or topk % 128 != 0:
raise ValueError(
"Q8KV8 sparse-prefill topk width must be a positive multiple of 128, "
f"got {topk}"
)
if topk_length is not None:
if topk_length.shape != (s_q,) or topk_length.dtype != torch.int32:
raise ValueError(
f"topk_length must be int32 with shape ({s_q},), got "
f"{tuple(topk_length.shape)}/{topk_length.dtype}"
)
if not topk_length.is_cuda:
raise ValueError("topk_length must be a CUDA tensor")
if topk_length.device != device:
raise ValueError(
"topk_length must be on q's device "
f"{device}, got {topk_length.device}"
)
if not topk_length.is_contiguous():
raise ValueError("topk_length must be contiguous")
if torch.any(topk_length < 0).item() or torch.any(topk_length > topk).item():
raise ValueError(
"topk_length values must satisfy " f"0 <= topk_length <= topk ({topk})"
)
if d_v != 512: if d_v != 512:
raise ValueError( raise ValueError(
@@ -307,15 +403,49 @@ def sparse_mla_q8kv8_prefill_fwd(
if attn_sink is not None and topk_length is None: if attn_sink is not None and topk_length is None:
raise ValueError("attn_sink requires topk_length to be provided as well") raise ValueError("attn_sink requires topk_length to be provided as well")
device = q.device if attn_sink is not None:
if attn_sink.shape != (h_q,) or attn_sink.dtype != torch.float32:
raise ValueError(
f"attn_sink must be float32 with shape ({h_q},), got "
f"{tuple(attn_sink.shape)}/{attn_sink.dtype}"
)
if not attn_sink.is_cuda:
raise ValueError("attn_sink must be a CUDA tensor")
if attn_sink.device != device:
raise ValueError(
f"attn_sink must be on q's device {device}, got {attn_sink.device}"
)
if not attn_sink.is_contiguous():
raise ValueError("attn_sink must be contiguous")
for name, scale in (("q_scale", q_scale), ("kv_scale", kv_scale)):
if not isinstance(scale, torch.Tensor):
raise ValueError(f"{name} must be a torch.Tensor")
if not scale.is_cuda:
raise ValueError(f"{name} must be a CUDA tensor")
if scale.device != device:
raise ValueError(
f"{name} must be on q's device {device}, got {scale.device}"
)
if scale.dtype != torch.float32:
raise ValueError(f"{name} must be float32, got {scale.dtype}")
if scale.numel() != 1:
raise ValueError(
f"{name} must be a scalar tensor, got shape {tuple(scale.shape)}"
)
if not scale.is_contiguous():
raise ValueError(f"{name} must be contiguous")
if out is None: if out is None:
out = torch.empty(s_q, h_q, d_v, dtype=torch.bfloat16, device=device) out = torch.empty(s_q, h_q, d_v, dtype=torch.bfloat16, device=device)
else: else:
_check_out_buffer(out, "out", (s_q, h_q, d_v), torch.bfloat16, device) _check_out_buffer(out, "out", (s_q, h_q, d_v), torch.bfloat16, device)
if max_logits is None: if max_logits is None:
max_logits = torch.empty(s_q, h_q, dtype=torch.float32, device=device) max_logits = torch.empty(s_q, h_q, dtype=torch.float32, device=device)
else: else:
_check_out_buffer(max_logits, "max_logits", (s_q, h_q), torch.float32, device) _check_out_buffer(max_logits, "max_logits", (s_q, h_q), torch.float32, device)
if lse is None: if lse is None:
lse = torch.empty(s_q, h_q, dtype=torch.float32, device=device) lse = torch.empty(s_q, h_q, dtype=torch.float32, device=device)
else: else:
@@ -19,7 +19,11 @@ import torch
import torch.nn.functional as F import torch.nn.functional as F
from sglang.kernels.ops.attention.dsv4.dequant_k_cache import ( from sglang.kernels.ops.attention.dsv4.dequant_k_cache import (
cast_q_fp8_for_q8kv8_prefill,
dequantize_k_cache_paged, dequantize_k_cache_paged,
fp8_dtype,
gather_dequant_requant_fp8_paged,
q8kv8_padded_num_heads,
) )
from sglang.kernels.ops.attention.dsv4.metadata_kernel import ( from sglang.kernels.ops.attention.dsv4.metadata_kernel import (
init_compression_metadata as _init_compression_metadata_triton, init_compression_metadata as _init_compression_metadata_triton,
@@ -56,8 +60,12 @@ from sglang.srt.layers.attention.dsv4.metadata import (
from sglang.srt.layers.attention.dsv4.sparse_prefill_utils import ( from sglang.srt.layers.attention.dsv4.sparse_prefill_utils import (
SparsePrefillChunkCache, SparsePrefillChunkCache,
SparsePrefillWorkspace, SparsePrefillWorkspace,
use_dsv4_q8kv8_sparse_prefill,
)
from sglang.srt.layers.attention.verify_mask import (
VerifyMask,
maybe_create_verify_mask,
) )
from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask
from sglang.srt.layers.cp.utils import is_cp_v2_active from sglang.srt.layers.cp.utils import is_cp_v2_active
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
@@ -70,7 +78,7 @@ from sglang.srt.speculative.ragged_verify import (
read_ragged_verify_mode, read_ragged_verify_mode,
resolve_ragged_verify_layout, resolve_ragged_verify_layout,
) )
from sglang.srt.utils import ceil_align, is_cuda, is_xpu from sglang.srt.utils import ceil_align, is_cuda, is_sm90_supported, is_xpu
from sglang.srt.utils.common import is_sm120_supported from sglang.srt.utils.common import is_sm120_supported
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -552,6 +560,22 @@ class DeepseekV4AttnBackend(
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend( self.dsa_topk_backend: DSATopKBackend = DSATopKBackend(
model_runner.server_args.dsa_topk_backend model_runner.server_args.dsa_topk_backend
) )
self.dsv4_prefill_backend: str = getattr(
model_runner.server_args, "dsv4_prefill_backend", "auto"
)
if use_dsv4_q8kv8_sparse_prefill(self.dsv4_prefill_backend):
if not is_sm90_supported():
raise ValueError(
"DeepSeek-V4 flashmla_sparse_q8 prefill requires SM90 CUDA GPUs."
)
if self.head_dim_v != 512:
raise ValueError(
"DeepSeek-V4 flashmla_sparse_q8 prefill requires d_v=512, "
f"got {self.head_dim_v}."
)
self._q8kv8_qpad_buf = None
self._q8kv8_attn_sink_pad = None
self._q8kv8_identity_scale = None
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = model_runner.server_args.speculative_eagle_topk or 0
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4" assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
self.mtp_enabled = self.topk > 0 self.mtp_enabled = self.topk > 0
@@ -1673,6 +1697,16 @@ class DeepseekV4AttnBackend(
or envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.get() or envs.SGLANG_OPT_FLASHMLA_SPARSE_PREFILL.get()
) )
): ):
if use_dsv4_q8kv8_sparse_prefill(self.dsv4_prefill_backend):
return self._forward_prefill_sparse_q8kv8(
q=q,
layer_id=layer_id,
compress_ratio=compress_ratio,
forward_batch=forward_batch,
token_to_kv_pool=token_to_kv_pool,
core_attn_metadata=core_attn_metadata,
attn_sink=attn_sink,
)
return self._forward_prefill_sparse( return self._forward_prefill_sparse(
q=q, q=q,
layer_id=layer_id, layer_id=layer_id,
@@ -1847,6 +1881,216 @@ class DeepseekV4AttnBackend(
) )
return o return o
def _prepare_q8kv8_q_and_sink(
self,
q: torch.Tensor,
attn_sink: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int]:
"""Pad TP-local heads to the SM90 kernel's 64-head CTA granularity."""
num_tokens, num_heads, head_dim = q.shape
padded_heads = q8kv8_padded_num_heads(num_heads)
qpad = getattr(self, "_q8kv8_qpad_buf", None)
if (
qpad is None
or qpad.shape[0] < num_tokens
or qpad.shape[1] != padded_heads
or qpad.shape[2] != head_dim
or qpad.device != q.device
):
qpad = torch.empty(
(num_tokens, padded_heads, head_dim),
dtype=fp8_dtype,
device=q.device,
)
self._q8kv8_qpad_buf = qpad
qpad = qpad[:num_tokens]
q_fp8, _ = cast_q_fp8_for_q8kv8_prefill(
q,
padded_num_heads=padded_heads,
out=qpad,
)
sink_pad = getattr(self, "_q8kv8_attn_sink_pad", None)
if (
sink_pad is None
or sink_pad.shape != (padded_heads,)
or sink_pad.device != q.device
):
sink_pad = torch.zeros(padded_heads, dtype=torch.float32, device=q.device)
self._q8kv8_attn_sink_pad = sink_pad
sink_pad[:num_heads].copy_(attn_sink.reshape(-1)[:num_heads])
if padded_heads > num_heads:
sink_pad[num_heads:].zero_()
scale = getattr(self, "_q8kv8_identity_scale", None)
if scale is None or scale.device != q.device:
scale = torch.ones((), dtype=torch.float32, device=q.device)
self._q8kv8_identity_scale = scale
return q_fp8, sink_pad, scale, num_heads
def _forward_prefill_sparse_q8kv8(
self,
q: torch.Tensor,
layer_id: int,
compress_ratio: Literal[0, 4, 128],
forward_batch: ForwardBatch,
token_to_kv_pool: DeepSeekV4TokenToKVPool,
core_attn_metadata: DSV4AttnMetadata,
attn_sink: torch.Tensor,
) -> torch.Tensor:
"""Experimental DeepSeek-V4 sparse prefill path using Q8KV8 kernels.
This mirrors ``_forward_prefill_sparse``'s cache/index construction, but
writes the gathered KV workspace as FP8 and calls the SM90 Q8KV8 sparse
prefill kernel. The path is selected by ``--dsv4-prefill-backend
flashmla_sparse_q8``; ``SGLANG_DSV4_Q8KV8_PREFILL`` remains as a debug
override for focused runtime validation.
"""
from sglang.kernels.ops.attention.sparse_mla_q8kv8_prefill_sm90 import (
sparse_mla_q8kv8_prefill_fwd,
)
q_flat = q.squeeze(1)
if q_flat.ndim != 3:
raise ValueError(
f"Q8KV8 sparse prefill expects 3D Q after squeeze, got {q_flat.shape}"
)
if attn_sink.numel() != q_flat.shape[1]:
raise ValueError(
f"attn_sink has {attn_sink.numel()} heads but Q has "
f"{q_flat.shape[1]} local heads"
)
q_fp8, attn_sink_pad, identity_scale, active_heads = (
self._prepare_q8kv8_q_and_sink(q_flat, attn_sink)
)
if not getattr(self, "_q8kv8_sparse_prefill_log_emitted", False):
logger.info(
"DSV4_Q8KV8_SPARSE_PREFILL_HIT layer_id=%s "
"compress_ratio=%s q_shape=%s padded_heads=%s d_v=%s",
layer_id,
compress_ratio,
tuple(q_flat.shape),
q_fp8.shape[1],
self.head_dim_v,
)
self._q8kv8_sparse_prefill_log_emitted = True
cache = self.forward_metadata.sparse_prefill_cache
if cache is None:
seq_lens_cpu = forward_batch.seq_lens_cpu
assert seq_lens_cpu is not None
extend_seq_lens_cpu = forward_batch.extend_seq_lens_cpu
assert extend_seq_lens_cpu is not None
total_swa = sum(
min(int(seq_len), int(extend_len) + SWA_WINDOW - 1)
for seq_len, extend_len in zip(
seq_lens_cpu.tolist(), extend_seq_lens_cpu, strict=True
)
)
cache = SparsePrefillChunkCache.build(
seq_lens=forward_batch.seq_lens.to(torch.int32),
extend_seq_lens=forward_batch.extend_seq_lens.to(torch.int32),
req_pool_indices=forward_batch.req_pool_indices.to(torch.int32),
req_to_token=self.req_to_token,
full_to_swa=token_to_kv_pool.full_to_swa_index_mapping,
swa_window_size=SWA_WINDOW,
swa_page_size=token_to_kv_pool.swa_window_size,
num_qo_tokens=q_flat.shape[0],
max_seq_len=int(seq_lens_cpu.max().item()),
total_swa=total_swa,
)
self.forward_metadata.sparse_prefill_cache = cache
compressed_slice = None
extra_k_cache = None
extra_page_size = None
flat_token_ids = None
if compress_ratio == 0:
workspace = self.sparse_prefill_workspace.get(
cache.swa_token_ids.shape[0] + 1,
dtype=fp8_dtype,
)
combined_indices = cache.c0_combined_indices
combined_lens = cache.c0_combined_lens
swa_slice = workspace
else:
extra_page_size = token_to_kv_pool.get_extra_key_page_size(layer_id)
extra_k_cache = token_to_kv_pool.get_extra_key_buffer(layer_id)
if compress_ratio == 128:
assert core_attn_metadata.c128_page_indices is not None
cache.ensure_c128(core_attn_metadata.c128_page_indices)
flat_token_ids = cache.c128_flat_token_ids
combined_indices = cache.c128_combined_indices
combined_lens = cache.c128_combined_lens
else:
assert core_attn_metadata.c4_sparse_raw_indices is not None, (
"Q8KV8 sparse-prefill c4 path requires c4_sparse_raw_indices "
"(allocated in init_flashmla_related when is_prefill=True)"
)
cache.ensure_c4(core_attn_metadata.page_table, extra_page_size)
flat_token_ids = cache.c4_flat_token_ids
combined_indices, combined_lens = cache.combine_c4_layer(
c4_sparse_raw_indices=core_attn_metadata.c4_sparse_raw_indices[
: cache.num_qo_tokens
],
)
n_compressed = flat_token_ids.shape[0]
workspace = self.sparse_prefill_workspace.get(
n_compressed + cache.swa_token_ids.shape[0] + 1,
dtype=fp8_dtype,
)
compressed_slice = workspace[:n_compressed]
swa_slice = workspace[n_compressed:]
if compressed_slice is not None:
gather_dequant_requant_fp8_paged(
extra_k_cache,
flat_token_ids,
page_size=extra_page_size,
out=compressed_slice,
)
gather_dequant_requant_fp8_paged(
token_to_kv_pool.get_swa_key_buffer_radix(layer_id),
cache.swa_token_ids,
page_size=cache.swa_page_size,
extra_rows=1,
out=swa_slice,
)
sentinel_row = workspace.shape[0] - 1
q8_indices = torch.where(
combined_indices < 0,
torch.full_like(combined_indices, sentinel_row),
combined_indices,
)
o, _, _ = sparse_mla_q8kv8_prefill_fwd(
q=q_fp8,
kv=workspace,
indices=q8_indices.unsqueeze(1),
sm_scale=self.softmax_scale,
q_scale=identity_scale,
kv_scale=identity_scale,
d_v=self.head_dim_v,
attn_sink=attn_sink_pad,
topk_length=combined_lens,
)
return o[:, :active_heads]
def expand_prefill_casually( def expand_prefill_casually(
self, self,
num_tokens: int, num_tokens: int,
@@ -32,6 +32,7 @@ For SWA-only layers callers pass ``topk=0``, ``compressed_base = 0`` (the
compressed branch becomes a no-op) and any ``compress_ratio >= 1``. compressed branch becomes a no-op) and any ``compress_ratio >= 1``.
""" """
import os
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Optional from typing import Optional
@@ -47,7 +48,8 @@ SPARSE_PREFILL_TOPK_ALIGNMENT = 128
# Bf16 workspace per-token width, matching ``dequantize_k_cache_paged``'s # Bf16 workspace per-token width, matching ``dequantize_k_cache_paged``'s
# output: 448 fp8 nope (dequanted) + 64 bf16 rope = 512. # output: 448 fp8 nope (dequanted) + 64 bf16 rope = 512.
WORKSPACE_DIM = DIM_NOPE + DIM_ROPE WORKSPACE_DIM = DIM_NOPE + DIM_ROPE
DSV4_Q8KV8_PREFILL_ENV = "SGLANG_DSV4_Q8KV8_PREFILL"
DSV4_Q8KV8_PREFILL_LOG_ENV = "SGLANG_DSV4_Q8KV8_PREFILL_LOG"
from sglang.kernels.ops.attention.dsv4.sparse_prefill_kernels import ( from sglang.kernels.ops.attention.dsv4.sparse_prefill_kernels import (
_build_swa_token_ids_kernel, _build_swa_token_ids_kernel,
@@ -55,6 +57,24 @@ from sglang.kernels.ops.attention.dsv4.sparse_prefill_kernels import (
) )
def use_dsv4_q8kv8_sparse_prefill(dsv4_prefill_backend: str = "auto") -> bool:
"""Return whether DeepSeek-V4 sparse prefill should use Q8KV8.
``dsv4_prefill_backend`` is the production configuration. The environment
variable remains as a debug override while the runtime path is being
hardened: truthy values force Q8 on, falsy values force it off.
"""
env_value = os.getenv(DSV4_Q8KV8_PREFILL_ENV)
if env_value is not None:
return env_value.lower() in {
"1",
"true",
"yes",
"on",
}
return dsv4_prefill_backend == "flashmla_sparse_q8"
class SparsePrefillWorkspace: class SparsePrefillWorkspace:
"""Backend-owned scratch storage for sparse prefill KV dequantization. """Backend-owned scratch storage for sparse prefill KV dequantization.
@@ -68,13 +88,18 @@ class SparsePrefillWorkspace:
self.device = device self.device = device
self._buffer: Optional[torch.Tensor] = None self._buffer: Optional[torch.Tensor] = None
def get(self, num_tokens: int) -> torch.Tensor: def get(
self,
num_tokens: int,
dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
assert num_tokens > 0 assert num_tokens > 0
current_capacity = self._buffer.shape[0] if self._buffer is not None else 0 current_capacity = self._buffer.shape[0] if self._buffer is not None else 0
if num_tokens > current_capacity: current_dtype = self._buffer.dtype if self._buffer is not None else None
if num_tokens > current_capacity or dtype != current_dtype:
self._buffer = torch.empty( self._buffer = torch.empty(
(num_tokens, 1, WORKSPACE_DIM), (num_tokens, 1, WORKSPACE_DIM),
dtype=torch.bfloat16, dtype=dtype,
device=self.device, device=self.device,
) )
return self._buffer[:num_tokens] return self._buffer[:num_tokens]
+18
View File
@@ -359,6 +359,12 @@ DSA_CHOICES = [
] ]
NSA_CHOICES = DSA_CHOICES # deprecated alias NSA_CHOICES = DSA_CHOICES # deprecated alias
DSV4_PREFILL_BACKEND_CHOICES = [
"auto",
"flashmla_sparse",
"flashmla_sparse_q8",
]
DSA_TOPK_BACKEND_CHOICES = ["sgl-kernel", "torch", "flashinfer"] DSA_TOPK_BACKEND_CHOICES = ["sgl-kernel", "torch", "flashinfer"]
DSA_PAGED_MQA_LOGITS_BACKEND_CHOICES = ["auto", "deepgemm", "cutedsl", "aiter"] DSA_PAGED_MQA_LOGITS_BACKEND_CHOICES = ["auto", "deepgemm", "cutedsl", "aiter"]
@@ -1805,6 +1811,18 @@ class ServerArgs:
), ),
NS("exec.kernel"), NS("exec.kernel"),
] = None ] = None
dsv4_prefill_backend: A[
str,
Arg(
help=(
"DeepSeek-V4 sparse prefill backend. 'auto' and "
"'flashmla_sparse' use the existing BF16 sparse prefill path; "
"'flashmla_sparse_q8' enables the Q8KV8 sparse prefill path."
),
choices=DSV4_PREFILL_BACKEND_CHOICES,
),
NS("exec.kernel"),
] = "auto"
dsa_decode_backend: A[ dsa_decode_backend: A[
Optional[str], Optional[str],
Arg( Arg(
@@ -0,0 +1,681 @@
"""DeepSeek-V4 Q8KV8 sparse-prefill backend helper tests.
These tests avoid starting a full server. They construct the minimum V4
metadata and token-pool surface consumed by the sparse-prefill helpers, then
compare the BF16 sparse path's gathered workspace against the Q8 path's FP8
workspace after dequantizing it back to BF16.
"""
from __future__ import annotations
import sys
import types
from contextlib import contextmanager
from types import SimpleNamespace
import pytest
import torch
from sglang.kernels.ops.attention.dsv4.index_buf_accessor import SetKAndS
from sglang.kernels.ops.attention.dsv4.quant_k_cache import (
quant_to_nope_fp8_rope_bf16_pack_triton,
)
from sglang.kernels.ops.attention.sparse_mla_q8kv8_prefill_sm90 import (
sparse_mla_q8kv8_prefill_fwd,
)
from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend
from sglang.srt.layers.attention.dsv4.sparse_prefill_utils import (
SparsePrefillChunkCache,
SparsePrefillWorkspace,
use_dsv4_q8kv8_sparse_prefill,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.utils import is_sm90_supported
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=120, stage="base-b-kernel-unit", runner_config="1-gpu-large")
def test_q8kv8_sparse_prefill_backend_selector_uses_cli_value():
assert not use_dsv4_q8kv8_sparse_prefill()
assert not use_dsv4_q8kv8_sparse_prefill("auto")
assert not use_dsv4_q8kv8_sparse_prefill("flashmla_sparse")
assert use_dsv4_q8kv8_sparse_prefill("flashmla_sparse_q8")
class _Pool:
def __init__(self, page_size: int):
self.page_size = page_size
class _Capture:
def __init__(self):
self.calls = []
def record(self, **kwargs):
cloned = {}
for name, value in kwargs.items():
if isinstance(value, torch.Tensor):
cloned[name] = value.detach().clone()
else:
cloned[name] = value
self.calls.append(cloned)
class _TokenToKVPool:
def __init__(
self,
*,
swa_key_buffer: torch.Tensor,
full_to_swa_index_mapping: torch.Tensor,
page_size: int,
extra_key_buffer: torch.Tensor | None = None,
):
self._swa_key_buffer = swa_key_buffer
self._extra_key_buffer = (
extra_key_buffer if extra_key_buffer is not None else swa_key_buffer
)
self.full_to_swa_index_mapping = full_to_swa_index_mapping
self.swa_window_size = page_size
def get_swa_key_buffer_radix(self, layer_id: int) -> torch.Tensor:
_ = layer_id
return self._swa_key_buffer
def get_extra_key_page_size(self, layer_id: int) -> int:
_ = layer_id
return self.swa_window_size
def get_extra_key_buffer(self, layer_id: int) -> torch.Tensor:
_ = layer_id
return self._extra_key_buffer
def _sm90_available() -> bool:
return torch.cuda.is_available() and is_sm90_supported()
def _make_v4_paged_kv_cache(
*,
total_slots: int,
page_size: int,
seed: int,
device: torch.device,
) -> torch.Tensor:
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
num_pages = (total_slots + page_size - 1) // page_size
total_slots = num_pages * page_size
bytes_per_token = 448 + 64 * 2 + 8
quant_k_cache = torch.zeros(
num_pages,
page_size * bytes_per_token,
dtype=torch.uint8,
device=device,
)
k_bf16 = (torch.randn(total_slots, 512, device=device) * 0.25).to(torch.bfloat16)
pack = quant_to_nope_fp8_rope_bf16_pack_triton(k_bf16)
loc = torch.arange(total_slots, dtype=torch.int32, device=device)
SetKAndS.torch(_Pool(page_size), quant_k_cache, loc, pack)
return quant_k_cache
def _make_forward_batch_and_mapping(
device: torch.device,
) -> tuple[ForwardBatch, torch.Tensor]:
seq_lens = torch.tensor([96, 144], dtype=torch.int32, device=device)
extend_seq_lens = torch.tensor([3, 2], dtype=torch.int32, device=device)
req_pool_indices = torch.tensor([0, 1], dtype=torch.int32, device=device)
seq0 = int(seq_lens[0].item())
seq1 = int(seq_lens[1].item())
req_to_token = torch.zeros(
(2, int(seq_lens.max().item())), dtype=torch.int32, device=device
)
req_to_token[0, :seq0] = torch.arange(seq0, dtype=torch.int32, device=device)
req1_base = 192
req_to_token[1, :seq1] = req1_base + torch.arange(
seq1, dtype=torch.int32, device=device
)
forward_batch = ForwardBatch(
forward_mode=ForwardMode.EXTEND,
batch_size=2,
input_ids=torch.zeros(
int(extend_seq_lens.sum().item()), dtype=torch.int32, device=device
),
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=torch.zeros(
int(extend_seq_lens.sum().item()), dtype=torch.int32, device=device
),
seq_lens_sum=int(seq_lens.sum().item()),
seq_lens_cpu=seq_lens.detach().cpu(),
extend_num_tokens=int(extend_seq_lens.sum().item()),
extend_seq_lens=extend_seq_lens,
extend_seq_lens_cpu=[int(x) for x in extend_seq_lens.detach().cpu().tolist()],
)
return forward_batch, req_to_token
def _make_backend(
device: torch.device,
req_to_token: torch.Tensor,
dsv4_prefill_backend: str = "auto",
) -> DeepseekV4AttnBackend:
backend = DeepseekV4AttnBackend.__new__(DeepseekV4AttnBackend)
backend.forward_metadata = SimpleNamespace(sparse_prefill_cache=None)
backend.req_to_token = req_to_token
backend.sparse_prefill_workspace = SparsePrefillWorkspace(device)
backend.softmax_scale = 512**-0.5
backend.head_dim_v = 512
backend.dsv4_prefill_backend = dsv4_prefill_backend
return backend
def _make_sparse_prefill_case(
device: torch.device,
local_heads: int = 64,
):
page_size = 64
total_slots = 384
forward_batch, req_to_token = _make_forward_batch_and_mapping(device)
backend = _make_backend(device, req_to_token)
quant_k_cache = _make_v4_paged_kv_cache(
total_slots=total_slots,
page_size=page_size,
seed=3,
device=device,
)
extra_k_cache = _make_v4_paged_kv_cache(
total_slots=total_slots,
page_size=page_size,
seed=7,
device=device,
)
token_to_kv_pool = _TokenToKVPool(
swa_key_buffer=quant_k_cache,
extra_key_buffer=extra_k_cache,
full_to_swa_index_mapping=torch.arange(
total_slots, dtype=torch.int64, device=device
),
page_size=page_size,
)
generator = torch.Generator(device=device)
generator.manual_seed(11)
q = (
torch.randn(
forward_batch.extend_num_tokens,
1,
local_heads,
512,
device=device,
generator=generator,
)
* 0.05
).to(torch.bfloat16)
attn_sink = torch.zeros(local_heads, dtype=torch.float32, device=device)
core_attn_metadata = SimpleNamespace()
return backend, forward_batch, token_to_kv_pool, q, attn_sink, core_attn_metadata
def _populate_compress_metadata(
core_attn_metadata: SimpleNamespace,
*,
compress_ratio: int,
device: torch.device,
) -> None:
if compress_ratio == 4:
core_attn_metadata.page_table = torch.zeros(
(2, 4), dtype=torch.int32, device=device
)
core_attn_metadata.c4_sparse_raw_indices = torch.zeros(
(16, 1), dtype=torch.int32, device=device
)
elif compress_ratio == 128:
core_attn_metadata.c128_page_indices = torch.zeros(
(16, 1), dtype=torch.int32, device=device
)
@contextmanager
def _patched_compressed_sparse_cache_paths(compress_ratio: int):
if compress_ratio == 0:
yield
return
old_ensure_c4 = SparsePrefillChunkCache.ensure_c4
old_ensure_c128 = SparsePrefillChunkCache.ensure_c128
old_combine_c4_layer = SparsePrefillChunkCache.combine_c4_layer
def _with_compressed_prefix(cache: SparsePrefillChunkCache, n_compressed: int):
shifted_swa = torch.where(
cache.c0_combined_indices >= 0,
cache.c0_combined_indices + n_compressed,
cache.c0_combined_indices,
)
n_prefix = min(n_compressed, shifted_swa.shape[1])
if n_prefix > 0:
shifted_swa[:, :n_prefix] = torch.arange(
n_prefix, dtype=shifted_swa.dtype, device=shifted_swa.device
)
combined_lens = torch.clamp(
cache.c0_combined_lens + n_prefix,
max=shifted_swa.shape[1],
)
return shifted_swa, combined_lens
def fake_ensure_c128(self, c128_page_indices):
_ = c128_page_indices
n_compressed = 8
self.c128_flat_token_ids = torch.arange(
n_compressed, dtype=torch.int64, device=self.swa_token_ids.device
)
self.c128_combined_indices, self.c128_combined_lens = _with_compressed_prefix(
self, n_compressed
)
def fake_ensure_c4(self, page_table, extra_page_size):
_ = page_table, extra_page_size
n_compressed = 8
self.c4_flat_token_ids = torch.arange(
n_compressed, dtype=torch.int64, device=self.swa_token_ids.device
)
def fake_combine_c4_layer(self, c4_sparse_raw_indices):
_ = c4_sparse_raw_indices
return _with_compressed_prefix(self, self.c4_flat_token_ids.shape[0])
SparsePrefillChunkCache.ensure_c128 = fake_ensure_c128
SparsePrefillChunkCache.ensure_c4 = fake_ensure_c4
SparsePrefillChunkCache.combine_c4_layer = fake_combine_c4_layer
try:
yield
finally:
SparsePrefillChunkCache.ensure_c4 = old_ensure_c4
SparsePrefillChunkCache.ensure_c128 = old_ensure_c128
SparsePrefillChunkCache.combine_c4_layer = old_combine_c4_layer
def _make_q8kv8_kernel_args(
*,
device: torch.device,
s_q: int = 4,
h_q: int = 64,
d_qk: int = 512,
s_kv: int = 256,
h_kv: int = 1,
topk: int = 128,
):
q = (torch.randn(s_q, h_q, d_qk, device=device) * 0.05).to(torch.float8_e4m3fn)
kv = (torch.randn(s_kv, h_kv, d_qk, device=device) * 0.05).to(torch.float8_e4m3fn)
indices = torch.randint(
0, s_kv, (s_q, h_kv, topk), dtype=torch.int32, device=device
)
topk_length = torch.full((s_q,), topk, dtype=torch.int32, device=device)
return {
"q": q.contiguous(),
"kv": kv.contiguous(),
"indices": indices.contiguous(),
"sm_scale": 512**-0.5,
"q_scale": torch.ones((), dtype=torch.float32, device=device),
"kv_scale": torch.ones((), dtype=torch.float32, device=device),
"d_v": 512,
"attn_sink": torch.zeros(h_q, dtype=torch.float32, device=device),
"topk_length": topk_length,
}
@contextmanager
def _patched_sparse_kernels(
bf16_capture: _Capture,
q8_capture: _Capture,
):
def fake_flash_mla_sparse_fwd(
*,
q,
kv,
indices,
sm_scale,
d_v,
attn_sink,
topk_length,
):
bf16_capture.record(
q=q,
kv=kv,
indices=indices,
sm_scale=sm_scale,
d_v=d_v,
attn_sink=attn_sink,
topk_length=topk_length,
)
out = torch.zeros(
(q.shape[0], q.shape[1], d_v), dtype=torch.bfloat16, device=q.device
)
meta = torch.zeros(
(q.shape[0], q.shape[1]), dtype=torch.float32, device=q.device
)
return out, meta, meta
def fake_sparse_mla_q8kv8_prefill_fwd(
*,
q,
kv,
indices,
sm_scale,
q_scale,
kv_scale,
d_v,
attn_sink,
topk_length,
):
q8_capture.record(
q=q,
kv=kv,
indices=indices,
sm_scale=sm_scale,
q_scale=q_scale,
kv_scale=kv_scale,
d_v=d_v,
attn_sink=attn_sink,
topk_length=topk_length,
)
out = torch.zeros(
(q.shape[0], q.shape[1], d_v), dtype=torch.bfloat16, device=q.device
)
meta = torch.zeros(
(q.shape[0], q.shape[1]), dtype=torch.float32, device=q.device
)
return out, meta, meta
sgl_kernel_pkg = types.ModuleType("sgl_kernel")
flash_mla_mod = types.ModuleType("sgl_kernel.flash_mla")
flash_mla_mod.flash_mla_sparse_fwd = fake_flash_mla_sparse_fwd
sgl_kernel_pkg.flash_mla = flash_mla_mod
q8_module_name = "sglang.kernels.ops.attention.sparse_mla_q8kv8_prefill_sm90"
q8_mod = types.ModuleType(q8_module_name)
q8_mod.sparse_mla_q8kv8_prefill_fwd = fake_sparse_mla_q8kv8_prefill_fwd
old_modules = {
name: sys.modules.get(name)
for name in (
"sgl_kernel",
"sgl_kernel.flash_mla",
q8_module_name,
)
}
sys.modules["sgl_kernel"] = sgl_kernel_pkg
sys.modules["sgl_kernel.flash_mla"] = flash_mla_mod
sys.modules[q8_module_name] = q8_mod
try:
yield
finally:
for name, old_value in old_modules.items():
if old_value is None:
sys.modules.pop(name, None)
else:
sys.modules[name] = old_value
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is not available")
@pytest.mark.parametrize("compress_ratio", [0, 4, 128])
def test_q8kv8_sparse_prefill_helper_builds_fp8_workspace_matching_bf16_path(
compress_ratio: int,
):
from sglang.kernels.ops.attention.dsv4.dequant_k_cache import fp8_dtype
device = torch.device("cuda")
backend, forward_batch, token_to_kv_pool, q, attn_sink, core_attn_metadata = (
_make_sparse_prefill_case(device, local_heads=16)
)
_populate_compress_metadata(
core_attn_metadata,
compress_ratio=compress_ratio,
device=device,
)
bf16_capture = _Capture()
q8_capture = _Capture()
with _patched_sparse_kernels(
bf16_capture, q8_capture
), _patched_compressed_sparse_cache_paths(compress_ratio):
bf16_out = backend._forward_prefill_sparse(
q=q,
layer_id=0,
compress_ratio=compress_ratio,
forward_batch=forward_batch,
token_to_kv_pool=token_to_kv_pool,
core_attn_metadata=core_attn_metadata,
attn_sink=attn_sink,
)
sparse_cache = backend.forward_metadata.sparse_prefill_cache
q8_out = backend._forward_prefill_sparse_q8kv8(
q=q,
layer_id=0,
compress_ratio=compress_ratio,
forward_batch=forward_batch,
token_to_kv_pool=token_to_kv_pool,
core_attn_metadata=core_attn_metadata,
attn_sink=attn_sink,
)
assert backend.forward_metadata.sparse_prefill_cache is sparse_cache
assert (
bf16_out.shape
== q8_out.shape
== (
forward_batch.extend_num_tokens,
16,
512,
)
)
assert len(bf16_capture.calls) == 1
assert len(q8_capture.calls) == 1
bf16_call = bf16_capture.calls[0]
q8_call = q8_capture.calls[0]
bf16_kv = bf16_call["kv"]
q8_kv = q8_call["kv"]
assert bf16_kv.dtype == torch.bfloat16
assert q8_kv.dtype == fp8_dtype
assert q8_call["q"].dtype == fp8_dtype
assert q8_call["q"].shape[1] == 64
assert torch.count_nonzero(q8_call["q"][:, 16:]).item() == 0
assert q8_call["attn_sink"].shape == (64,)
assert q8_kv.shape[0] == bf16_kv.shape[0] + 1
torch.testing.assert_close(
q8_kv[:-1].to(torch.bfloat16).float(),
bf16_kv.float(),
atol=3e-2,
rtol=2e-1,
)
assert torch.equal(
q8_kv[-1].to(torch.bfloat16),
torch.zeros_like(q8_kv[-1].to(torch.bfloat16)),
)
bf16_indices = bf16_call["indices"]
q8_indices = q8_call["indices"]
sentinel_row = q8_kv.shape[0] - 1
valid_mask = bf16_indices >= 0
assert torch.equal(q8_indices[valid_mask], bf16_indices[valid_mask])
assert torch.equal(
q8_indices[~valid_mask],
torch.full_like(q8_indices[~valid_mask], sentinel_row),
)
assert torch.equal(q8_call["topk_length"], bf16_call["topk_length"])
assert q8_call["q_scale"].item() == pytest.approx(1.0)
assert q8_call["kv_scale"].item() == pytest.approx(1.0)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is not available")
def test_q8kv8_sparse_prefill_rejects_topk_64_before_cuda_launch():
args = _make_q8kv8_kernel_args(device=torch.device("cuda"), topk=64)
with pytest.raises(ValueError, match="positive multiple of 128"):
sparse_mla_q8kv8_prefill_fwd(**args)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is not available")
@pytest.mark.parametrize(
("mutate", "error_match"),
[
(
lambda args: args.update(q=args["q"].float()),
"q must be torch.float8_e4m3fn",
),
(
lambda args: args.update(kv=args["kv"].float()),
"kv must be torch.float8_e4m3fn",
),
(
lambda args: args.update(
q=torch.empty(
args["q"].shape[0],
args["q"].shape[1],
args["q"].shape[2] + 1,
dtype=args["q"].dtype,
device=args["q"].device,
)[:, :, : args["q"].shape[2]]
),
"q must be contiguous",
),
(
lambda args: args.update(
q_scale=torch.ones(2, dtype=torch.float32, device=args["q"].device)
),
"q_scale must be a scalar tensor",
),
(
lambda args: args.update(
kv_scale=torch.ones((), dtype=torch.float16, device=args["q"].device)
),
"kv_scale must be float32",
),
],
)
def test_q8kv8_sparse_prefill_rejects_invalid_tensor_contracts(
mutate,
error_match: str,
):
args = _make_q8kv8_kernel_args(device=torch.device("cuda"), topk=128)
mutate(args)
with pytest.raises(ValueError, match=error_match):
sparse_mla_q8kv8_prefill_fwd(**args)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is not available")
@pytest.mark.parametrize("bad_length", [-1, 129])
def test_q8kv8_sparse_prefill_rejects_invalid_topk_length_bounds(
bad_length: int,
):
args = _make_q8kv8_kernel_args(device=torch.device("cuda"), topk=128)
args["topk_length"][0] = bad_length
with pytest.raises(ValueError, match="0 <= topk_length <= topk"):
sparse_mla_q8kv8_prefill_fwd(**args)
@pytest.mark.skipif(
not _sm90_available(), reason="Q8KV8 sparse prefill requires SM90 CUDA"
)
def test_q8kv8_sparse_prefill_real_kernel_matches_bf16_sparse_path():
device = torch.device("cuda")
backend, forward_batch, token_to_kv_pool, q, attn_sink, core_attn_metadata = (
_make_sparse_prefill_case(device, local_heads=64)
)
bf16_out = backend._forward_prefill_sparse(
q=q,
layer_id=0,
compress_ratio=0,
forward_batch=forward_batch,
token_to_kv_pool=token_to_kv_pool,
core_attn_metadata=core_attn_metadata,
attn_sink=attn_sink,
)
sparse_cache = backend.forward_metadata.sparse_prefill_cache
q8_out = backend._forward_prefill_sparse_q8kv8(
q=q,
layer_id=0,
compress_ratio=0,
forward_batch=forward_batch,
token_to_kv_pool=token_to_kv_pool,
core_attn_metadata=core_attn_metadata,
attn_sink=attn_sink,
)
torch.cuda.synchronize()
assert backend.forward_metadata.sparse_prefill_cache is sparse_cache
assert (
bf16_out.shape
== q8_out.shape
== (
forward_batch.extend_num_tokens,
64,
512,
)
)
assert bf16_out.dtype == torch.bfloat16
assert q8_out.dtype == torch.bfloat16
assert torch.isfinite(bf16_out.float()).all()
assert torch.isfinite(q8_out.float()).all()
abs_diff = (q8_out.float() - bf16_out.float()).abs()
assert abs_diff.mean().item() < 0.03
assert torch.quantile(abs_diff.flatten(), 0.99).item() < 0.2
torch.testing.assert_close(
q8_out.float(),
bf16_out.float(),
atol=2.5e-1,
rtol=3.0e-1,
)
@pytest.mark.skipif(
not _sm90_available(), reason="Q8KV8 sparse prefill requires SM90 CUDA"
)
def test_q8kv8_sparse_prefill_real_kernel_repeated_launch_stable():
args = _make_q8kv8_kernel_args(
device=torch.device("cuda"),
s_q=512,
h_q=64,
d_qk=512,
s_kv=1024,
h_kv=1,
topk=256,
)
baseline = None
for _ in range(10):
out, max_logits, lse = sparse_mla_q8kv8_prefill_fwd(**args)
torch.cuda.synchronize()
assert out.shape == (512, 64, 512)
assert max_logits.shape == (512, 64)
assert lse.shape == (512, 64)
assert torch.isfinite(out.float()).all()
assert torch.isfinite(max_logits).all()
assert torch.isfinite(lse).all()
current = out.float().detach().clone()
if baseline is None:
baseline = current
else:
torch.testing.assert_close(
current,
baseline,
atol=1e-2,
rtol=1e-2,
)
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"]))
@@ -52,6 +52,23 @@ class TestPrepareServerArgs(CustomTestCase):
): ):
ServerArgs(model_path="dummy", prefill_decode_interval=-1) ServerArgs(model_path="dummy", prefill_decode_interval=-1)
def test_dsv4_prefill_backend_cli_choices(self):
parser = server_args_module.argparse.ArgumentParser()
ServerArgs.add_cli_args(parser)
base_args = ["--model-path", "dummy-model"]
default_args = parser.parse_args(base_args)
self.assertEqual(default_args.dsv4_prefill_backend, "auto")
q8_args = parser.parse_args(
base_args + ["--dsv4-prefill-backend", "flashmla_sparse_q8"]
)
self.assertEqual(q8_args.dsv4_prefill_backend, "flashmla_sparse_q8")
with self.assertRaises(SystemExit):
parser.parse_args(base_args + ["--dsv4-prefill-backend", "flashmla_kv"])
def test_return_hidden_states_mode_configuration(self): def test_return_hidden_states_mode_configuration(self):
disabled = ServerArgs(model_path="dummy") disabled = ServerArgs(model_path="dummy")
self.assertFalse(disabled.enable_return_hidden_states) self.assertFalse(disabled.enable_return_hidden_states)