[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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user