[DeepseekV32] Enable flashmla_prefill kernel with fp8 kvcache (#11655)

Signed-off-by: Hao Lu <14827759+hlu1@users.noreply.github.com>
This commit is contained in:
hlu1
2025-10-27 23:11:48 -07:00
committed by GitHub
parent 83b2240074
commit 81a632ace6
4 changed files with 367 additions and 44 deletions
+156 -20
View File
@@ -1,12 +1,14 @@
from __future__ import annotations
from dataclasses import dataclass
from enum import IntEnum, auto
from typing import TYPE_CHECKING, Dict, List, Literal, Optional, TypeAlias
import torch
from sglang.srt.configs.model_config import get_nsa_index_topk, is_deepseek_nsa
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.nsa.dequant_k_cache import dequantize_k_cache_paged
from sglang.srt.layers.attention.nsa.nsa_indexer import BaseIndexerMetadata
from sglang.srt.layers.attention.nsa.quant_k_cache import quantize_k_cache
from sglang.srt.layers.attention.nsa.transform_index import (
@@ -98,11 +100,27 @@ class NSAMetadata:
nsa_max_seqlen_q: Literal[1] = 1 # always 1 for decode, variable for extend
flashmla_metadata: Optional[NSAFlashMLAMetadata] = None
# The sum of sequence lengths for key, prefill only
seq_lens_sum: Optional[int] = None
# The flattened 1D page table with shape (seq_lens_sum,), prefill only
# this table is always with page_size = 1
page_table_1_flattened: Optional[torch.Tensor] = None
# The offset of topk indices in ragged kv, prefill only
# shape: (seq_lens_sum,)
topk_indices_offset: Optional[torch.Tensor] = None
class TopkTransformMethod(IntEnum):
# Transform topk indices to indices to the page table (page_size = 1)
PAGED = auto()
# Transform topk indices to indices to ragged kv (non-paged)
RAGGED = auto()
@dataclass(frozen=True)
class NSAIndexerMetadata(BaseIndexerMetadata):
attn_metadata: NSAMetadata
topk_transform_method: TopkTransformMethod
def get_seqlens_int32(self) -> torch.Tensor:
return self.attn_metadata.cache_seqlens_int32
@@ -118,23 +136,36 @@ class NSAIndexerMetadata(BaseIndexerMetadata):
logits: torch.Tensor,
topk: int,
) -> torch.Tensor:
from sgl_kernel import fast_topk_transform_fused, fast_topk_v2
from sgl_kernel import (
fast_topk_transform_fused,
fast_topk_transform_ragged_fused,
fast_topk_v2,
)
if not NSA_FUSE_TOPK:
return fast_topk_v2(logits, self.get_seqlens_expanded(), topk)
# NOTE(dark): if fused, we return a transformed page table directly
return fast_topk_transform_fused(
score=logits,
lengths=self.get_seqlens_expanded(),
page_table_size_1=self.attn_metadata.page_table_1,
cu_seqlens_q=self.attn_metadata.cu_seqlens_q,
topk=topk,
)
elif self.topk_transform_method == TopkTransformMethod.PAGED:
# NOTE(dark): if fused, we return a transformed page table directly
return fast_topk_transform_fused(
score=logits,
lengths=self.get_seqlens_expanded(),
page_table_size_1=self.attn_metadata.page_table_1,
cu_seqlens_q=self.attn_metadata.cu_seqlens_q,
topk=topk,
)
elif self.topk_transform_method == TopkTransformMethod.RAGGED:
return fast_topk_transform_ragged_fused(
score=logits,
lengths=self.get_seqlens_expanded(),
topk_indices_offset=self.attn_metadata.topk_indices_offset,
topk=topk,
)
else:
assert False, f"Unsupported {self.topk_transform_method = }"
def compute_cu_seqlens(seqlens: torch.Tensor) -> torch.Tensor:
assert seqlens.dtype == torch.int32 and seqlens.is_cuda
assert seqlens.dtype == torch.int32
return torch.nn.functional.pad(
torch.cumsum(seqlens, dim=0, dtype=torch.int32), (1, 0)
)
@@ -181,6 +212,7 @@ class NativeSparseAttnBackend(AttentionBackend):
global NSA_PREFILL_IMPL, NSA_DECODE_IMPL
NSA_PREFILL_IMPL = model_runner.server_args.nsa_prefill_backend
NSA_DECODE_IMPL = model_runner.server_args.nsa_decode_backend
self.enable_auto_select_prefill_impl = NSA_PREFILL_IMPL == "flashmla_auto"
self._arange_buf = torch.arange(16384, device=self.device, dtype=torch.int32)
@@ -231,10 +263,16 @@ class NativeSparseAttnBackend(AttentionBackend):
cu_seqlens_k = compute_cu_seqlens(cache_seqlens_int32)
assert forward_batch.seq_lens_cpu is not None
max_seqlen_k = int(forward_batch.seq_lens_cpu.max().item() + draft_token_num)
# [b, max_seqlen_k]
page_table = forward_batch.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices, :max_seqlen_k
]
page_table_1_flattened = None
topk_indices_offset = None
self.set_nsa_prefill_impl(forward_batch)
topk_transform_method = self.get_topk_transform_method()
if forward_batch.forward_mode.is_decode_or_idle():
extend_seq_lens_cpu = [1] * batch_size
max_seqlen_q = 1
@@ -295,6 +333,7 @@ class NativeSparseAttnBackend(AttentionBackend):
else:
max_seqlen_q = max_seqlen_k
cu_seqlens_q = cu_seqlens_k
seqlens_expanded = torch.cat(
[
torch.arange(
@@ -310,6 +349,24 @@ class NativeSparseAttnBackend(AttentionBackend):
)
]
)
if topk_transform_method == TopkTransformMethod.RAGGED:
page_table_1_flattened = torch.cat(
[
page_table[i, :kv_len]
for i, kv_len in enumerate(
forward_batch.seq_lens_cpu.tolist(),
)
]
)
assert (
page_table_1_flattened.shape[0] == forward_batch.seq_lens_sum
), f"{page_table_1_flattened.shape[0] = } must be the same as {forward_batch.seq_lens_sum = }"
topk_indices_offset = torch.repeat_interleave(
cu_seqlens_k[:-1],
forward_batch.extend_seq_lens,
)
else:
assert False, f"Unsupported {forward_batch.forward_mode = }"
@@ -328,7 +385,9 @@ class NativeSparseAttnBackend(AttentionBackend):
max_seq_len_k=max_seqlen_k,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
seq_lens_sum=forward_batch.seq_lens_sum,
page_table_1=page_table,
page_table_1_flattened=page_table_1_flattened,
flashmla_metadata=(
self._compute_flashmla_metadata(
cache_seqlens=nsa_cache_seqlens_int32,
@@ -344,6 +403,7 @@ class NativeSparseAttnBackend(AttentionBackend):
nsa_extend_seq_lens_list=extend_seq_lens_cpu,
real_page_table=self._transform_table_1_to_real(page_table),
nsa_max_seqlen_q=1,
topk_indices_offset=topk_indices_offset,
)
self.forward_metadata = metadata
@@ -396,6 +456,8 @@ class NativeSparseAttnBackend(AttentionBackend):
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
self.set_nsa_prefill_impl(forward_batch=None)
"""Initialize forward metadata for capturing CUDA graph."""
if forward_mode.is_decode_or_idle():
# Normal Decode
@@ -586,6 +648,8 @@ class NativeSparseAttnBackend(AttentionBackend):
"""Initialize forward metadata for replaying CUDA graph."""
assert seq_lens_cpu is not None
self.set_nsa_prefill_impl(forward_batch=None)
seq_lens = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs]
req_pool_indices = req_pool_indices[:bs]
@@ -780,17 +844,31 @@ class NativeSparseAttnBackend(AttentionBackend):
q_rope = q_all[:, :, layer.v_head_dim :]
# NOTE(dark): here, we use page size = 1
topk_transform_method = self.get_topk_transform_method()
if NSA_FUSE_TOPK:
page_table_1 = topk_indices
else:
assert metadata.nsa_extend_seq_lens_list is not None
page_table_1 = transform_index_page_table_prefill(
page_table=metadata.page_table_1,
topk_indices=topk_indices,
extend_lens_cpu=metadata.nsa_extend_seq_lens_list,
page_size=1,
)
if topk_transform_method == TopkTransformMethod.RAGGED:
topk_indices_offset = metadata.topk_indices_offset
assert topk_indices_offset is not None
mask = topk_indices != -1
topk_indices_offset = (
topk_indices_offset.unsqueeze(1)
if topk_indices_offset.ndim == 1
else topk_indices_offset
)
topk_indices = torch.where(
mask, topk_indices + topk_indices_offset, topk_indices
)
elif topk_transform_method == TopkTransformMethod.PAGED:
assert metadata.nsa_extend_seq_lens_list is not None
page_table_1 = transform_index_page_table_prefill(
page_table=metadata.page_table_1,
topk_indices=topk_indices,
extend_lens_cpu=metadata.nsa_extend_seq_lens_list,
page_size=1,
)
if NSA_PREFILL_IMPL == "tilelang":
if q_rope is not None:
q_all = torch.cat([q_nope, q_rope], dim=-1)
@@ -804,6 +882,22 @@ class NativeSparseAttnBackend(AttentionBackend):
elif NSA_PREFILL_IMPL == "flashmla_sparse":
if q_rope is not None:
q_all = torch.cat([q_nope, q_rope], dim=-1)
# NSA_FLASHMLA_BACKEND_DECODE_COMPUTE_FP8 has no effect here,
# because the flashmla_sparse kernel doesn't support fp8 compute
if topk_transform_method == TopkTransformMethod.RAGGED:
if any(forward_batch.extend_prefix_lens_cpu):
page_table_1_flattened = (
self.forward_metadata.page_table_1_flattened
)
assert page_table_1_flattened is not None
kv_cache = dequantize_k_cache_paged(
kv_cache, page_table_1_flattened
)
else:
kv_cache = torch.cat([k, k_rope], dim=-1)
page_table_1 = topk_indices
return self._forward_flashmla_sparse(
q_all=q_all,
kv_cache=kv_cache,
@@ -1121,10 +1215,52 @@ class NativeSparseAttnBackend(AttentionBackend):
"""Get the fill value for sequence length in CUDA graph."""
return 1
def set_nsa_prefill_impl(self, forward_batch: Optional[ForwardBatch] = None) -> str:
from sglang.srt.utils import is_blackwell
global NSA_PREFILL_IMPL
if self.enable_auto_select_prefill_impl:
if self.nsa_kv_cache_store_fp8:
if (
# TODO(hlu1): enable MTP
is_blackwell()
and forward_batch is not None
and forward_batch.forward_mode.is_extend()
and forward_batch.spec_algorithm.is_none()
):
total_kv_tokens = forward_batch.seq_lens_sum
total_q_tokens = forward_batch.extend_num_tokens
# Heuristic based on benchmarking flashmla_kv vs flashmla_sparse + dequantize_k_cache_paged
if total_kv_tokens < total_q_tokens * 512:
NSA_PREFILL_IMPL = "flashmla_sparse"
return
NSA_PREFILL_IMPL = "flashmla_kv"
else:
# bf16 kv cache
NSA_PREFILL_IMPL = "flashmla_sparse"
def get_topk_transform_method(self) -> TopkTransformMethod:
"""
NSA_FUSE_TOPK controls whether to fuse the topk transform into the topk kernel.
This method is used to select the topk transform method which can be fused or unfused.
"""
if (
# disable for MTP
self.nsa_kv_cache_store_fp8
and NSA_PREFILL_IMPL == "flashmla_sparse"
):
topk_transform_method = TopkTransformMethod.RAGGED
else:
topk_transform_method = TopkTransformMethod.PAGED
return topk_transform_method
def get_indexer_metadata(
self, layer_id: int, forward_batch: ForwardBatch
) -> NSAIndexerMetadata:
return NSAIndexerMetadata(attn_metadata=self.forward_metadata)
return NSAIndexerMetadata(
attn_metadata=self.forward_metadata,
topk_transform_method=self.get_topk_transform_method(),
)
def _compute_flashmla_metadata(self, cache_seqlens: torch.Tensor, seq_len_q: int):
from flash_mla import get_mla_metadata