Enable XQA for SM90 and SM120 (#17115)
Co-authored-by: Xiaowei Wang <100599594+xiaoweiw-nv@users.noreply.github.com>
This commit is contained in:
co-authored by
Xiaowei Wang
parent
dc92f88a21
commit
5b2e2750b5
@@ -228,6 +228,25 @@ python3 -m sglang.launch_server \
|
|||||||
--trust-remote-code
|
--trust-remote-code
|
||||||
```
|
```
|
||||||
|
|
||||||
|
- TRTLLM MHA (Optimized for Blackwell Architecture, e.g., B200)
|
||||||
|
```bash
|
||||||
|
python3 -m sglang.launch_server \
|
||||||
|
--tp 4 \
|
||||||
|
--model Qwen/Qwen3.5-35B-A3B-FP8 \
|
||||||
|
--attention-backend trtllm_mha \
|
||||||
|
--trust-remote-code
|
||||||
|
```
|
||||||
|
|
||||||
|
- TRTLLM MHA (XQA backend) (Optimized for SM90 and SM120, e.g., H20, H200, 5090)
|
||||||
|
Note that TRTLLM XQA backend only works well for pagesize 64.
|
||||||
|
```bash
|
||||||
|
python3 -m sglang.launch_server \
|
||||||
|
--tp 4 \
|
||||||
|
--model Qwen/Qwen3.5-35B-A3B-FP8 \
|
||||||
|
--decode-attention-backend trtllm_mha \
|
||||||
|
--trust-remote-code
|
||||||
|
```
|
||||||
|
|
||||||
- FlashAttention 4 (MHA & MLA)
|
- FlashAttention 4 (MHA & MLA)
|
||||||
```bash
|
```bash
|
||||||
# FA4 for both prefill and decode on SM90/SM100
|
# FA4 for both prefill and decode on SM90/SM100
|
||||||
|
|||||||
@@ -145,3 +145,25 @@ class HybridAttnBackend(AttentionBackend):
|
|||||||
) -> Optional[BaseIndexerMetadata]:
|
) -> Optional[BaseIndexerMetadata]:
|
||||||
backend = self._select_backend(forward_batch.forward_mode)
|
backend = self._select_backend(forward_batch.forward_mode)
|
||||||
return backend.get_indexer_metadata(layer_id, forward_batch)
|
return backend.get_indexer_metadata(layer_id, forward_batch)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
q: torch.Tensor = None,
|
||||||
|
k: torch.Tensor = None,
|
||||||
|
v: torch.Tensor = None,
|
||||||
|
layer: RadixAttention = None,
|
||||||
|
forward_batch: ForwardBatch = None,
|
||||||
|
save_kv_cache: bool = True,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
"""Delegate forward to the appropriate backend based on forward mode."""
|
||||||
|
backend = self._select_backend(forward_batch.forward_mode)
|
||||||
|
return backend.forward(
|
||||||
|
q=q,
|
||||||
|
k=k,
|
||||||
|
v=v,
|
||||||
|
layer=layer,
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
save_kv_cache=save_kv_cache,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from sglang.srt.layers.attention.utils import canonicalize_stride
|
|||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool, SWATokenToKVPoolAllocator
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool, SWATokenToKVPoolAllocator
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.utils import is_flashinfer_available
|
from sglang.srt.utils import is_flashinfer_available
|
||||||
|
from sglang.srt.utils.common import is_sm90_supported, is_sm120_supported
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -138,6 +139,16 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
# Forward metadata
|
# Forward metadata
|
||||||
self.forward_metadata: Optional[TRTLLMMHAMetadata] = None
|
self.forward_metadata: Optional[TRTLLMMHAMetadata] = None
|
||||||
|
|
||||||
|
# Init backend (XQA or TRTLLM-GEN)
|
||||||
|
# We need to specify q_type and out_type for different backend
|
||||||
|
# XQA: (q_type must be bf16)
|
||||||
|
# KV bf16: q_type = bf16, out_type=model_runner.dtype
|
||||||
|
# KV fp8: q_type = bf16, out_type=model_runner.dtype
|
||||||
|
# TRTLLM-GEN:
|
||||||
|
# KV bf16: q_type = bf16, out_type=model_runner.dtype
|
||||||
|
# KV fp8: q_type = fp8, out_type=model_runner.dtype
|
||||||
|
self.is_xqa_impl = is_sm90_supported() or is_sm120_supported()
|
||||||
|
|
||||||
def _maybe_translate_swa(
|
def _maybe_translate_swa(
|
||||||
self, token_indices: torch.Tensor
|
self, token_indices: torch.Tensor
|
||||||
) -> Optional[torch.Tensor]:
|
) -> Optional[torch.Tensor]:
|
||||||
@@ -714,7 +725,8 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
layer, cache_loc, k, v, layer.k_scale, layer.v_scale
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.data_type == torch.float8_e4m3fn:
|
# For XQA, q_dtype should be bf16
|
||||||
|
if self.data_type == torch.float8_e4m3fn and (not self.is_xqa_impl):
|
||||||
q = q.to(torch.float8_e4m3fn)
|
q = q.to(torch.float8_e4m3fn)
|
||||||
q = q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim)
|
q = q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim)
|
||||||
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
k_cache, v_cache = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||||
|
|||||||
@@ -2035,9 +2035,28 @@ class ServerArgs:
|
|||||||
or self.decode_attention_backend == "trtllm_mha"
|
or self.decode_attention_backend == "trtllm_mha"
|
||||||
or self.prefill_attention_backend == "trtllm_mha"
|
or self.prefill_attention_backend == "trtllm_mha"
|
||||||
):
|
):
|
||||||
if not is_sm100_supported():
|
# Check prefill backend
|
||||||
|
prefill_backend = (
|
||||||
|
self.prefill_attention_backend
|
||||||
|
if self.prefill_attention_backend is not None
|
||||||
|
else self.attention_backend
|
||||||
|
)
|
||||||
|
if prefill_backend == "trtllm_mha" and not is_sm100_supported():
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"TRTLLM MHA backend is only supported on Blackwell GPUs (SM100). Please use a different backend."
|
"TRTLLM MHA backend for prefill is only supported on Blackwell GPUs (SM100). Please use a different prefill backend."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Check decode backend
|
||||||
|
decode_backend = (
|
||||||
|
self.decode_attention_backend
|
||||||
|
if self.decode_attention_backend is not None
|
||||||
|
else self.attention_backend
|
||||||
|
)
|
||||||
|
if decode_backend == "trtllm_mha" and not (
|
||||||
|
is_sm90_supported() or is_sm100_supported() or is_sm120_supported()
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"TRTLLM MHA backend for decode is only supported on Hopper (SM90), Blackwell (SM100) and (SM120) GPUs. Please use a different decode backend."
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.page_size not in [16, 32, 64]:
|
if self.page_size not in [16, 32, 64]:
|
||||||
|
|||||||
Reference in New Issue
Block a user