[DeepSeek-V3.2][NSA] Enable MHA Pathway for Short Sequence Prefill on B200 (SM100) (#12788)
This commit is contained in:
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Dict, List, Literal, Optional, TypeAlias
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import get_nsa_index_topk, is_deepseek_nsa
|
from sglang.srt.configs.model_config import get_nsa_index_topk, is_deepseek_nsa
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
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.dequant_k_cache import dequantize_k_cache_paged
|
||||||
from sglang.srt.layers.attention.nsa.nsa_indexer import BaseIndexerMetadata
|
from sglang.srt.layers.attention.nsa.nsa_indexer import BaseIndexerMetadata
|
||||||
@@ -50,6 +51,10 @@ else:
|
|||||||
from sgl_kernel.flash_attn import flash_attn_varlen_func, flash_attn_with_kvcache
|
from sgl_kernel.flash_attn import flash_attn_varlen_func, flash_attn_with_kvcache
|
||||||
|
|
||||||
|
|
||||||
|
# Reuse this workspace buffer across all NSA backend instances
|
||||||
|
global_workspace_buffer = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class NSAFlashMLAMetadata:
|
class NSAFlashMLAMetadata:
|
||||||
"""Metadata only needed by FlashMLA"""
|
"""Metadata only needed by FlashMLA"""
|
||||||
@@ -231,6 +236,20 @@ class NativeSparseAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
self.speculative_step_id = speculative_step_id
|
self.speculative_step_id = speculative_step_id
|
||||||
|
|
||||||
|
# Allocate global workspace buffer for TRTLLm ragged attention kernel (SM100/B200)
|
||||||
|
device_sm_major = torch.cuda.get_device_capability()[0]
|
||||||
|
if device_sm_major >= 10:
|
||||||
|
global global_workspace_buffer
|
||||||
|
if global_workspace_buffer is None:
|
||||||
|
global_workspace_buffer = torch.empty(
|
||||||
|
envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.get(),
|
||||||
|
dtype=torch.uint8,
|
||||||
|
device=model_runner.device,
|
||||||
|
)
|
||||||
|
self.workspace_buffer = global_workspace_buffer
|
||||||
|
else:
|
||||||
|
self.workspace_buffer = None
|
||||||
|
|
||||||
def get_device_int32_arange(self, l: int) -> torch.Tensor:
|
def get_device_int32_arange(self, l: int) -> torch.Tensor:
|
||||||
if l > len(self._arange_buf):
|
if l > len(self._arange_buf):
|
||||||
next_pow_of_2 = 1 << (l - 1).bit_length()
|
next_pow_of_2 = 1 << (l - 1).bit_length()
|
||||||
@@ -1196,9 +1215,34 @@ class NativeSparseAttnBackend(AttentionBackend):
|
|||||||
f"cu_seqlens_k has {len(cu_seqlens_k)-1} requests"
|
f"cu_seqlens_k has {len(cu_seqlens_k)-1} requests"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Determine FA version: FA3 for SM90 (Hopper), FA4 for SM100+ (Blackwell and beyond)
|
# Use TRTLLm ragged attention for SM100 (Blackwell/B200) to avoid FA4 accuracy issues
|
||||||
device_sm_major = torch.cuda.get_device_capability()[0]
|
device_sm_major = torch.cuda.get_device_capability()[0]
|
||||||
fa_version = 4 if device_sm_major >= 10 else 3
|
if device_sm_major >= 10:
|
||||||
|
import flashinfer
|
||||||
|
|
||||||
|
seq_lens = metadata.cache_seqlens_int32
|
||||||
|
return flashinfer.prefill.trtllm_ragged_attention_deepseek(
|
||||||
|
query=q,
|
||||||
|
key=k,
|
||||||
|
value=v,
|
||||||
|
workspace_buffer=self.workspace_buffer,
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
max_q_len=metadata.max_seq_len_q,
|
||||||
|
max_kv_len=max_seqlen_k,
|
||||||
|
bmm1_scale=layer.scaling,
|
||||||
|
bmm2_scale=1.0,
|
||||||
|
o_sf_scale=1.0,
|
||||||
|
batch_size=forward_batch.batch_size,
|
||||||
|
window_left=-1,
|
||||||
|
cum_seq_lens_q=cu_seqlens_q,
|
||||||
|
cum_seq_lens_kv=cu_seqlens_k,
|
||||||
|
enable_pdl=False,
|
||||||
|
is_causal=causal,
|
||||||
|
return_lse=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Use FA3 for SM90 (Hopper/H200)
|
||||||
|
fa_version = 3
|
||||||
|
|
||||||
return flash_attn_varlen_func(
|
return flash_attn_varlen_func(
|
||||||
q=q,
|
q=q,
|
||||||
|
|||||||
@@ -415,11 +415,14 @@ def handle_attention_nsa(attn, forward_batch):
|
|||||||
assert forward_batch.seq_lens_cpu is not None
|
assert forward_batch.seq_lens_cpu is not None
|
||||||
max_kv_len = forward_batch.seq_lens_cpu.max().item()
|
max_kv_len = forward_batch.seq_lens_cpu.max().item()
|
||||||
|
|
||||||
# B200 (SM100) is temporarily disabled for MHA due to FA4 accuracy issues
|
# MHA path enabled for both H200 (SM90, FA3) and B200 (SM100, TRTLLm ragged)
|
||||||
# Currently only H200 (SM90) with FA3 is allowed to use MHA path
|
# B200 uses trtllm_ragged_attention_deepseek kernel instead of FA4
|
||||||
is_hopper = _device_sm == 90
|
supports_mha = _device_sm in [90, 100]
|
||||||
|
|
||||||
if max_kv_len <= attn.indexer.index_topk and is_hopper:
|
# Check if kvcache dtype is bfloat16
|
||||||
|
kv_dtype_is_bf16 = forward_batch.token_to_kv_pool.dtype == torch.bfloat16
|
||||||
|
|
||||||
|
if max_kv_len <= attn.indexer.index_topk and supports_mha and kv_dtype_is_bf16:
|
||||||
# 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)
|
||||||
|
|||||||
Reference in New Issue
Block a user