Support GPT-OSS zigzag CP with TRTLLM-MHA (#31732)
This commit is contained in:
@@ -25,6 +25,8 @@ from sglang.srt.layers.attention.flashinfer_backend import (
|
|||||||
FlashInferAttnBackend,
|
FlashInferAttnBackend,
|
||||||
FlashInferMultiStepDraftBackend,
|
FlashInferMultiStepDraftBackend,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy
|
||||||
|
from sglang.srt.layers.cp.utils import is_cp_v2_active
|
||||||
from sglang.srt.layers.quantization.fp4_kv_cache_quant_method import (
|
from sglang.srt.layers.quantization.fp4_kv_cache_quant_method import (
|
||||||
KVCacheAttentionAccessKind,
|
KVCacheAttentionAccessKind,
|
||||||
)
|
)
|
||||||
@@ -643,9 +645,16 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
"""Get the fill value for sequence lengths in CUDA graph."""
|
"""Get the fill value for sequence lengths in CUDA graph."""
|
||||||
return 1
|
return 1
|
||||||
|
|
||||||
def _should_use_fused_fp8_path(self, save_kv_cache: bool, k: torch.Tensor) -> bool:
|
def _should_use_fused_fp8_path(
|
||||||
|
self, save_kv_cache: bool, k: torch.Tensor, forward_batch: ForwardBatch
|
||||||
|
) -> bool:
|
||||||
"""Check if we should use the fused FP8 KV cache write path."""
|
"""Check if we should use the fused FP8 KV cache write path."""
|
||||||
return save_kv_cache and k is not None and self.data_type == torch.float8_e4m3fn
|
return (
|
||||||
|
not is_cp_v2_active(forward_batch)
|
||||||
|
and save_kv_cache
|
||||||
|
and k is not None
|
||||||
|
and self.data_type == torch.float8_e4m3fn
|
||||||
|
)
|
||||||
|
|
||||||
def _fused_fp8_qkv_kv_cache(
|
def _fused_fp8_qkv_kv_cache(
|
||||||
self,
|
self,
|
||||||
@@ -931,7 +940,9 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
"""Run forward for decode using TRTLLM MHA kernel."""
|
"""Run forward for decode using TRTLLM MHA kernel."""
|
||||||
cache_loc = forward_batch.out_cache_loc
|
cache_loc = forward_batch.out_cache_loc
|
||||||
|
|
||||||
use_fused_fp8_path = self._should_use_fused_fp8_path(save_kv_cache, k)
|
use_fused_fp8_path = self._should_use_fused_fp8_path(
|
||||||
|
save_kv_cache, k, forward_batch
|
||||||
|
)
|
||||||
use_fused_qkv = use_fused_fp8_path and not self.is_xqa_impl
|
use_fused_qkv = use_fused_fp8_path and not self.is_xqa_impl
|
||||||
pool = self.token_to_kv_pool
|
pool = self.token_to_kv_pool
|
||||||
|
|
||||||
@@ -1020,8 +1031,13 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
)
|
)
|
||||||
|
|
||||||
cache_loc = forward_batch.out_cache_loc
|
cache_loc = forward_batch.out_cache_loc
|
||||||
|
cp_v2_active = is_cp_v2_active(forward_batch)
|
||||||
|
|
||||||
use_fused_fp8_path = self._should_use_fused_fp8_path(save_kv_cache, k)
|
# The fused path writes rank-local K/V directly to cache. CP-v2 needs
|
||||||
|
# the strategy to gather K/V into full logical token order first.
|
||||||
|
use_fused_fp8_path = self._should_use_fused_fp8_path(
|
||||||
|
save_kv_cache, k, forward_batch
|
||||||
|
)
|
||||||
use_fused_qkv = use_fused_fp8_path and not self.is_xqa_impl
|
use_fused_qkv = use_fused_fp8_path and not self.is_xqa_impl
|
||||||
|
|
||||||
if use_fused_fp8_path:
|
if use_fused_fp8_path:
|
||||||
@@ -1033,8 +1049,18 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
k = None
|
k = None
|
||||||
v = None
|
v = None
|
||||||
else:
|
else:
|
||||||
# Use original set_kv_buffer path
|
|
||||||
if save_kv_cache and k is not None:
|
if save_kv_cache and k is not None:
|
||||||
|
if cp_v2_active:
|
||||||
|
cp_strategy = get_cp_strategy()
|
||||||
|
assert cp_strategy is not None
|
||||||
|
cp_strategy.materialize_full_kv(
|
||||||
|
forward_batch,
|
||||||
|
layer,
|
||||||
|
k,
|
||||||
|
v,
|
||||||
|
swa_loc=self.forward_metadata.swa_out_cache_loc,
|
||||||
|
)
|
||||||
|
else:
|
||||||
self.token_to_kv_pool.set_kv_buffer(
|
self.token_to_kv_pool.set_kv_buffer(
|
||||||
layer,
|
layer,
|
||||||
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc),
|
||||||
@@ -1116,23 +1142,50 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
q_len_per_req=self.forward_metadata.max_seq_len_q,
|
q_len_per_req=self.forward_metadata.max_seq_len_q,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
o = flashinfer.prefill.trtllm_batch_context_with_kv_cache(
|
|
||||||
query=q,
|
def _trtllm_context_attn(
|
||||||
|
q_chunk,
|
||||||
|
cu_seqlens_q,
|
||||||
|
cache_seqlens,
|
||||||
|
max_seqlen_q,
|
||||||
|
cu_seqlens_kv,
|
||||||
|
):
|
||||||
|
return flashinfer.prefill.trtllm_batch_context_with_kv_cache(
|
||||||
|
query=q_chunk,
|
||||||
kv_cache=kv_cache,
|
kv_cache=kv_cache,
|
||||||
workspace_buffer=self.workspace_buffer,
|
workspace_buffer=self.workspace_buffer,
|
||||||
block_tables=page_table,
|
block_tables=page_table,
|
||||||
seq_lens=self.forward_metadata.cache_seqlens_int32,
|
seq_lens=cache_seqlens,
|
||||||
max_q_len=self.forward_metadata.max_seq_len_q,
|
max_q_len=max_seqlen_q,
|
||||||
max_kv_len=self.max_context_len,
|
max_kv_len=self.max_context_len,
|
||||||
bmm1_scale=bmm1_scale,
|
bmm1_scale=bmm1_scale,
|
||||||
bmm2_scale=bmm2_scale,
|
bmm2_scale=bmm2_scale,
|
||||||
batch_size=self.forward_metadata.cu_seqlens_q.shape[0] - 1,
|
batch_size=cu_seqlens_q.shape[0] - 1,
|
||||||
cum_seq_lens_q=self.forward_metadata.cu_seqlens_q,
|
cum_seq_lens_q=cu_seqlens_q,
|
||||||
cum_seq_lens_kv=self.forward_metadata.cu_seqlens_k,
|
cum_seq_lens_kv=cu_seqlens_kv,
|
||||||
window_left=layer.sliding_window_size,
|
window_left=layer.sliding_window_size,
|
||||||
sinks=attention_sink,
|
sinks=attention_sink,
|
||||||
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get(),
|
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get(),
|
||||||
out_dtype=self.q_data_type, # model_runner.dtype
|
out_dtype=self.q_data_type,
|
||||||
|
)
|
||||||
|
|
||||||
|
if cp_v2_active:
|
||||||
|
cp_strategy = get_cp_strategy()
|
||||||
|
assert cp_strategy is not None
|
||||||
|
o = cp_strategy.run_attention(
|
||||||
|
q,
|
||||||
|
forward_batch,
|
||||||
|
self.device,
|
||||||
|
_trtllm_context_attn,
|
||||||
|
attention_backend=CPAttentionBackendKind.TRTLLM_MHA,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
o = _trtllm_context_attn(
|
||||||
|
q,
|
||||||
|
self.forward_metadata.cu_seqlens_q,
|
||||||
|
self.forward_metadata.cache_seqlens_int32,
|
||||||
|
self.forward_metadata.max_seq_len_q,
|
||||||
|
cu_seqlens_kv=self.forward_metadata.cu_seqlens_k,
|
||||||
)
|
)
|
||||||
|
|
||||||
return o.view(-1, layer.tp_q_head_num * layer.head_dim)
|
return o.view(-1, layer.tp_q_head_num * layer.head_dim)
|
||||||
|
|||||||
@@ -67,14 +67,17 @@ class CPAttentionBackendKind(IntEnum):
|
|||||||
"""Attention backend calling convention used by CP strategy dispatch."""
|
"""Attention backend calling convention used by CP strategy dispatch."""
|
||||||
|
|
||||||
FLASH_ATTENTION = 0
|
FLASH_ATTENTION = 0
|
||||||
|
TRTLLM_MHA = 1
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_string(cls, value: str) -> CPAttentionBackendKind:
|
def from_string(cls, value: str) -> CPAttentionBackendKind:
|
||||||
if value in ("fa3", "fa4", "flashinfer"):
|
if value in ("fa3", "fa4", "flashinfer"):
|
||||||
return cls.FLASH_ATTENTION
|
return cls.FLASH_ATTENTION
|
||||||
|
if value == "trtllm_mha":
|
||||||
|
return cls.TRTLLM_MHA
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported attention_backend={value!r} for CP strategy; expected one "
|
f"Unsupported attention_backend={value!r} for CP strategy; expected one "
|
||||||
"of {'fa3', 'fa4', 'flashinfer'}"
|
"of {'fa3', 'fa4', 'flashinfer', 'trtllm_mha'}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -40,6 +40,7 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
CP_V2_DEFAULT_MODEL_CLASSES = frozenset(
|
CP_V2_DEFAULT_MODEL_CLASSES = frozenset(
|
||||||
{
|
{
|
||||||
|
"GptOssForCausalLM",
|
||||||
"MiMoV2FlashForCausalLM",
|
"MiMoV2FlashForCausalLM",
|
||||||
"MiMoV2ForCausalLM",
|
"MiMoV2ForCausalLM",
|
||||||
"Qwen3MoeForCausalLM",
|
"Qwen3MoeForCausalLM",
|
||||||
|
|||||||
@@ -72,6 +72,8 @@ class ZigzagContextParallelMetadata(BaseContextParallelMetadata):
|
|||||||
# Per-sequence FlashAttention tensors (shape [bs] or [bs + 1]).
|
# Per-sequence FlashAttention tensors (shape [bs] or [bs + 1]).
|
||||||
kv_len_prev_tensor: Optional[Any] = None
|
kv_len_prev_tensor: Optional[Any] = None
|
||||||
kv_len_next_tensor: Optional[Any] = None
|
kv_len_next_tensor: Optional[Any] = None
|
||||||
|
cu_seqlens_kv_prev_tensor: Optional[Any] = None
|
||||||
|
cu_seqlens_kv_next_tensor: Optional[Any] = None
|
||||||
actual_seq_q_prev_tensor: Optional[Any] = None
|
actual_seq_q_prev_tensor: Optional[Any] = None
|
||||||
actual_seq_q_next_tensor: Optional[Any] = None
|
actual_seq_q_next_tensor: Optional[Any] = None
|
||||||
cu_seqlens_q_prev_tensor: Optional[Any] = None
|
cu_seqlens_q_prev_tensor: Optional[Any] = None
|
||||||
@@ -214,6 +216,8 @@ class ZigzagCPStrategy(ContextParallelStrategy):
|
|||||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
cu_prev = [0] + list(accumulate(actual_seq_q_prev_list))
|
cu_prev = [0] + list(accumulate(actual_seq_q_prev_list))
|
||||||
cu_next = [0] + list(accumulate(actual_seq_q_next_list))
|
cu_next = [0] + list(accumulate(actual_seq_q_next_list))
|
||||||
|
cu_kv_prev = [0] + list(accumulate(kv_len_prev_list))
|
||||||
|
cu_kv_next = [0] + list(accumulate(kv_len_next_list))
|
||||||
|
|
||||||
total_seq_lens = sum(extend_seqs_len)
|
total_seq_lens = sum(extend_seqs_len)
|
||||||
assert len(split_list) == bs * cp_segment_num
|
assert len(split_list) == bs * cp_segment_num
|
||||||
@@ -236,6 +240,12 @@ class ZigzagCPStrategy(ContextParallelStrategy):
|
|||||||
kv_len_next_tensor=torch.tensor(
|
kv_len_next_tensor=torch.tensor(
|
||||||
kv_len_next_list, device=device, dtype=torch.int32
|
kv_len_next_list, device=device, dtype=torch.int32
|
||||||
),
|
),
|
||||||
|
cu_seqlens_kv_prev_tensor=torch.tensor(
|
||||||
|
cu_kv_prev, device=device, dtype=torch.int32
|
||||||
|
),
|
||||||
|
cu_seqlens_kv_next_tensor=torch.tensor(
|
||||||
|
cu_kv_next, device=device, dtype=torch.int32
|
||||||
|
),
|
||||||
actual_seq_q_prev_tensor=torch.tensor(
|
actual_seq_q_prev_tensor=torch.tensor(
|
||||||
actual_seq_q_prev_list, device=device, dtype=torch.int32
|
actual_seq_q_prev_list, device=device, dtype=torch.int32
|
||||||
),
|
),
|
||||||
@@ -301,7 +311,10 @@ class ZigzagCPStrategy(ContextParallelStrategy):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_supported_attention_backend(self):
|
def get_supported_attention_backend(self):
|
||||||
return [CPAttentionBackendKind.FLASH_ATTENTION]
|
return [
|
||||||
|
CPAttentionBackendKind.FLASH_ATTENTION,
|
||||||
|
CPAttentionBackendKind.TRTLLM_MHA,
|
||||||
|
]
|
||||||
|
|
||||||
def run_attention(
|
def run_attention(
|
||||||
self,
|
self,
|
||||||
@@ -320,17 +333,25 @@ class ZigzagCPStrategy(ContextParallelStrategy):
|
|||||||
logical_tokens = meta.total_q_prev_tokens + meta.total_q_next_tokens
|
logical_tokens = meta.total_q_prev_tokens + meta.total_q_next_tokens
|
||||||
q_next = q[meta.total_q_prev_tokens : logical_tokens]
|
q_next = q[meta.total_q_prev_tokens : logical_tokens]
|
||||||
|
|
||||||
|
prev_kwargs = {}
|
||||||
|
next_kwargs = {}
|
||||||
|
if attention_backend == CPAttentionBackendKind.TRTLLM_MHA:
|
||||||
|
prev_kwargs["cu_seqlens_kv"] = meta.cu_seqlens_kv_prev_tensor
|
||||||
|
next_kwargs["cu_seqlens_kv"] = meta.cu_seqlens_kv_next_tensor
|
||||||
|
|
||||||
result_prev = attn_fn(
|
result_prev = attn_fn(
|
||||||
q_prev,
|
q_prev,
|
||||||
meta.cu_seqlens_q_prev_tensor,
|
meta.cu_seqlens_q_prev_tensor,
|
||||||
meta.kv_len_prev_tensor,
|
meta.kv_len_prev_tensor,
|
||||||
meta.max_seqlen_q_prev,
|
meta.max_seqlen_q_prev,
|
||||||
|
**prev_kwargs,
|
||||||
)
|
)
|
||||||
result_next = attn_fn(
|
result_next = attn_fn(
|
||||||
q_next,
|
q_next,
|
||||||
meta.cu_seqlens_q_next_tensor,
|
meta.cu_seqlens_q_next_tensor,
|
||||||
meta.kv_len_next_tensor,
|
meta.kv_len_next_tensor,
|
||||||
meta.max_seqlen_q_next,
|
meta.max_seqlen_q_next,
|
||||||
|
**next_kwargs,
|
||||||
)
|
)
|
||||||
result = torch.cat([result_prev, result_next], dim=0)
|
result = torch.cat([result_prev, result_next], dim=0)
|
||||||
pad_size = q.shape[0] - logical_tokens
|
pad_size = q.shape[0] - logical_tokens
|
||||||
|
|||||||
@@ -0,0 +1,32 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.gpt_oss_common import BaseTestGptOss
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=220, stage="extra-b", runner_config="4-gpu-b200")
|
||||||
|
|
||||||
|
|
||||||
|
class TestGptOss4GpuMxfp4CP(BaseTestGptOss):
|
||||||
|
def test_mxfp4_120b(self):
|
||||||
|
self.run_test(
|
||||||
|
model_variant="120b",
|
||||||
|
quantization="mxfp4",
|
||||||
|
expected_score_of_reasoning_effort={
|
||||||
|
"low": 0.58,
|
||||||
|
},
|
||||||
|
other_args=[
|
||||||
|
"--tp",
|
||||||
|
"4",
|
||||||
|
"--enable-prefill-cp",
|
||||||
|
"--attn-cp-size",
|
||||||
|
"4",
|
||||||
|
"--cp-strategy",
|
||||||
|
"zigzag",
|
||||||
|
"--cuda-graph-max-bs-decode",
|
||||||
|
"200",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user