[DeepseekV3.2] Deepseek fp8 support for MHA path (#12964)
This commit is contained in:
@@ -455,7 +455,14 @@ class NativeSparseAttnBackend(AttentionBackend):
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
if topk_transform_method == TopkTransformMethod.RAGGED:
|
# Generate page_table_1_flattened when needed:
|
||||||
|
mha_dequantize_needed = (
|
||||||
|
self.nsa_kv_cache_store_fp8 and max_seqlen_k <= self.nsa_index_topk
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
topk_transform_method == TopkTransformMethod.RAGGED
|
||||||
|
or mha_dequantize_needed
|
||||||
|
):
|
||||||
page_table_1_flattened = torch.cat(
|
page_table_1_flattened = torch.cat(
|
||||||
[
|
[
|
||||||
page_table[i, :kv_len]
|
page_table[i, :kv_len]
|
||||||
@@ -468,6 +475,7 @@ class NativeSparseAttnBackend(AttentionBackend):
|
|||||||
page_table_1_flattened.shape[0] == forward_batch.seq_lens_sum
|
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 = }"
|
), f"{page_table_1_flattened.shape[0] = } must be the same as {forward_batch.seq_lens_sum = }"
|
||||||
|
|
||||||
|
if topk_transform_method == TopkTransformMethod.RAGGED:
|
||||||
topk_indices_offset = torch.repeat_interleave(
|
topk_indices_offset = torch.repeat_interleave(
|
||||||
cu_seqlens_k[:-1],
|
cu_seqlens_k[:-1],
|
||||||
forward_batch.extend_seq_lens,
|
forward_batch.extend_seq_lens,
|
||||||
|
|||||||
@@ -53,6 +53,7 @@ from sglang.srt.layers.attention.npu_ops.mla_preprocess import (
|
|||||||
NPUFusedMLAPreprocess,
|
NPUFusedMLAPreprocess,
|
||||||
is_mla_preprocess_enabled,
|
is_mla_preprocess_enabled,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.attention.nsa.dequant_k_cache import dequantize_k_cache_paged
|
||||||
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
|
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
|
||||||
from sglang.srt.layers.attention.nsa.utils import (
|
from sglang.srt.layers.attention.nsa.utils import (
|
||||||
can_cp_split,
|
can_cp_split,
|
||||||
@@ -421,8 +422,6 @@ def handle_attention_nsa(attn, forward_batch):
|
|||||||
- Decode: MLA (avoids per-token decompression)
|
- Decode: MLA (avoids per-token decompression)
|
||||||
- Prefill <= 2048: MHA (topk ineffective, MHA has lower FLOPs)
|
- Prefill <= 2048: MHA (topk ineffective, MHA has lower FLOPs)
|
||||||
- Prefill > 2048: MLA (topk filtering reduces computation significantly)
|
- Prefill > 2048: MLA (topk filtering reduces computation significantly)
|
||||||
|
|
||||||
TODO: B200 (SM100) MHA path is temporarily disabled due to FA4 gpqa accuracy issues.
|
|
||||||
"""
|
"""
|
||||||
if forward_batch.forward_mode.is_decode_or_idle():
|
if forward_batch.forward_mode.is_decode_or_idle():
|
||||||
return AttnForwardMethod.MLA
|
return AttnForwardMethod.MLA
|
||||||
@@ -437,10 +436,17 @@ def handle_attention_nsa(attn, forward_batch):
|
|||||||
# B200 uses trtllm_ragged_attention_deepseek kernel instead of FA4
|
# B200 uses trtllm_ragged_attention_deepseek kernel instead of FA4
|
||||||
supports_mha = _device_sm in [90, 100]
|
supports_mha = _device_sm in [90, 100]
|
||||||
|
|
||||||
# Check if kvcache dtype is bfloat16
|
# MHA supports both BF16 and FP8 KV cache (FP8 will be dequantized on-demand)
|
||||||
kv_dtype_is_bf16 = forward_batch.token_to_kv_pool.dtype == torch.bfloat16
|
kv_dtype_supported = forward_batch.token_to_kv_pool.dtype in [
|
||||||
|
torch.bfloat16,
|
||||||
|
torch.float8_e4m3fn,
|
||||||
|
]
|
||||||
|
|
||||||
if max_kv_len <= attn.indexer.index_topk and supports_mha and kv_dtype_is_bf16:
|
if (
|
||||||
|
max_kv_len <= attn.indexer.index_topk
|
||||||
|
and supports_mha
|
||||||
|
and kv_dtype_supported
|
||||||
|
):
|
||||||
# NSA backend uses varlen kernel which supports MHA_ONE_SHOT
|
# NSA backend uses varlen kernel which supports MHA_ONE_SHOT
|
||||||
# Check if total sequence length fits in chunk capacity
|
# Check if total sequence length fits in chunk capacity
|
||||||
sum_seq_lens = sum(forward_batch.seq_lens_cpu)
|
sum_seq_lens = sum(forward_batch.seq_lens_cpu)
|
||||||
@@ -1659,9 +1665,16 @@ class DeepseekV2AttentionMLA(nn.Module):
|
|||||||
forward_batch.mha_one_shot
|
forward_batch.mha_one_shot
|
||||||
and sum(forward_batch.extend_prefix_lens_cpu) != 0
|
and sum(forward_batch.extend_prefix_lens_cpu) != 0
|
||||||
):
|
):
|
||||||
kv_a, k_pe = self._get_mla_kv_buffer(
|
if self.use_nsa and self.kv_cache_dtype == "fp8_e4m3":
|
||||||
forward_batch.fetch_mha_one_shot_kv_indices(), q.dtype, forward_batch
|
# FP8 path: dequantize NSA-specific FP8 format to BF16
|
||||||
)
|
kv_a, k_pe = self._get_mla_kv_buffer_from_fp8(forward_batch)
|
||||||
|
else:
|
||||||
|
# BF16/FP16 path: directly fetch from cache
|
||||||
|
kv_a, k_pe = self._get_mla_kv_buffer(
|
||||||
|
forward_batch.fetch_mha_one_shot_kv_indices(),
|
||||||
|
q.dtype,
|
||||||
|
forward_batch,
|
||||||
|
)
|
||||||
kv = self.kv_b_proj(kv_a)[0]
|
kv = self.kv_b_proj(kv_a)[0]
|
||||||
kv = kv.view(-1, self.num_local_heads, self.qk_nope_head_dim + self.v_head_dim)
|
kv = kv.view(-1, self.num_local_heads, self.qk_nope_head_dim + self.v_head_dim)
|
||||||
k_nope = kv[..., : self.qk_nope_head_dim]
|
k_nope = kv[..., : self.qk_nope_head_dim]
|
||||||
@@ -2733,6 +2746,31 @@ class DeepseekV2AttentionMLA(nn.Module):
|
|||||||
kv_a = kv_a.squeeze(1).contiguous()
|
kv_a = kv_a.squeeze(1).contiguous()
|
||||||
return kv_a, k_pe
|
return kv_a, k_pe
|
||||||
|
|
||||||
|
def _get_mla_kv_buffer_from_fp8(
|
||||||
|
self,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Dequantize FP8 KV cache to BF16 for MLA attention (NSA-specific format).
|
||||||
|
|
||||||
|
Returns: (kv_a, k_pe) both in BF16
|
||||||
|
"""
|
||||||
|
kv_indices = forward_batch.attn_backend.forward_metadata.page_table_1_flattened
|
||||||
|
assert (
|
||||||
|
kv_indices is not None
|
||||||
|
), "page_table_1_flattened should have been generated for FP8 MHA path"
|
||||||
|
|
||||||
|
kv_cache_fp8 = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||||
|
self.attn_mha.layer_id
|
||||||
|
)
|
||||||
|
|
||||||
|
kv_latent_bf16 = dequantize_k_cache_paged(kv_cache_fp8, kv_indices)
|
||||||
|
|
||||||
|
kv_a = kv_latent_bf16[:, :, : self.kv_lora_rank].squeeze(1).contiguous()
|
||||||
|
k_pe = kv_latent_bf16[:, :, self.kv_lora_rank :]
|
||||||
|
|
||||||
|
return kv_a, k_pe
|
||||||
|
|
||||||
def _concat_and_cast_mha_k(self, k_nope, k_pe, forward_batch):
|
def _concat_and_cast_mha_k(self, k_nope, k_pe, forward_batch):
|
||||||
# Temporary for DeepSeek V3/R1 only, but can generalize if needed
|
# Temporary for DeepSeek V3/R1 only, but can generalize if needed
|
||||||
k_shape = (k_nope.shape[0], self.num_local_heads, self.qk_head_dim)
|
k_shape = (k_nope.shape[0], self.num_local_heads, self.qk_head_dim)
|
||||||
|
|||||||
Reference in New Issue
Block a user