From 9668d9ea72ae9385c96f8dc1202df38222b98208 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Mon, 20 Jul 2026 00:48:04 -0700 Subject: [PATCH] Support GPT-OSS zigzag CP with TRTLLM-MHA (#31732) --- .../layers/attention/trtllm_mha_backend.py | 115 +++++++++++++----- python/sglang/srt/layers/cp/base.py | 5 +- python/sglang/srt/layers/cp/utils.py | 1 + python/sglang/srt/layers/cp/zigzag.py | 23 +++- .../cp/test_gpt_oss_4gpu_mxfp4_cp.py | 32 +++++ 5 files changed, 143 insertions(+), 33 deletions(-) create mode 100644 test/registered/cp/test_gpt_oss_4gpu_mxfp4_cp.py diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index cbd1a31a3..548c1c389 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -25,6 +25,8 @@ from sglang.srt.layers.attention.flashinfer_backend import ( FlashInferAttnBackend, 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 ( KVCacheAttentionAccessKind, ) @@ -643,9 +645,16 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): """Get the fill value for sequence lengths in CUDA graph.""" 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.""" - 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( self, @@ -931,7 +940,9 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): """Run forward for decode using TRTLLM MHA kernel.""" 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 pool = self.token_to_kv_pool @@ -1020,8 +1031,13 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): ) 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 if use_fused_fp8_path: @@ -1033,16 +1049,26 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): k = None v = None else: - # Use original set_kv_buffer path if save_kv_cache and k is not None: - self.token_to_kv_pool.set_kv_buffer( - layer, - KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc), - k, - v, - layer.k_scale, - layer.v_scale, - ) + 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( + layer, + KVWriteLoc(cache_loc, self.forward_metadata.swa_out_cache_loc), + k, + v, + layer.k_scale, + layer.v_scale, + ) q_scale = 1.0 if ( @@ -1116,24 +1142,51 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): q_len_per_req=self.forward_metadata.max_seq_len_q, ) else: - o = flashinfer.prefill.trtllm_batch_context_with_kv_cache( - query=q, - kv_cache=kv_cache, - workspace_buffer=self.workspace_buffer, - block_tables=page_table, - seq_lens=self.forward_metadata.cache_seqlens_int32, - max_q_len=self.forward_metadata.max_seq_len_q, - max_kv_len=self.max_context_len, - bmm1_scale=bmm1_scale, - bmm2_scale=bmm2_scale, - batch_size=self.forward_metadata.cu_seqlens_q.shape[0] - 1, - cum_seq_lens_q=self.forward_metadata.cu_seqlens_q, - cum_seq_lens_kv=self.forward_metadata.cu_seqlens_k, - window_left=layer.sliding_window_size, - sinks=attention_sink, - skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get(), - out_dtype=self.q_data_type, # model_runner.dtype - ) + + 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, + workspace_buffer=self.workspace_buffer, + block_tables=page_table, + seq_lens=cache_seqlens, + max_q_len=max_seqlen_q, + max_kv_len=self.max_context_len, + bmm1_scale=bmm1_scale, + bmm2_scale=bmm2_scale, + batch_size=cu_seqlens_q.shape[0] - 1, + cum_seq_lens_q=cu_seqlens_q, + cum_seq_lens_kv=cu_seqlens_kv, + window_left=layer.sliding_window_size, + sinks=attention_sink, + skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get(), + 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) diff --git a/python/sglang/srt/layers/cp/base.py b/python/sglang/srt/layers/cp/base.py index 4fb2a9da2..6b66a21f4 100644 --- a/python/sglang/srt/layers/cp/base.py +++ b/python/sglang/srt/layers/cp/base.py @@ -67,14 +67,17 @@ class CPAttentionBackendKind(IntEnum): """Attention backend calling convention used by CP strategy dispatch.""" FLASH_ATTENTION = 0 + TRTLLM_MHA = 1 @classmethod def from_string(cls, value: str) -> CPAttentionBackendKind: if value in ("fa3", "fa4", "flashinfer"): return cls.FLASH_ATTENTION + if value == "trtllm_mha": + return cls.TRTLLM_MHA raise ValueError( f"Unsupported attention_backend={value!r} for CP strategy; expected one " - "of {'fa3', 'fa4', 'flashinfer'}" + "of {'fa3', 'fa4', 'flashinfer', 'trtllm_mha'}" ) diff --git a/python/sglang/srt/layers/cp/utils.py b/python/sglang/srt/layers/cp/utils.py index 24de66560..1ec8887f8 100644 --- a/python/sglang/srt/layers/cp/utils.py +++ b/python/sglang/srt/layers/cp/utils.py @@ -40,6 +40,7 @@ if TYPE_CHECKING: CP_V2_DEFAULT_MODEL_CLASSES = frozenset( { + "GptOssForCausalLM", "MiMoV2FlashForCausalLM", "MiMoV2ForCausalLM", "Qwen3MoeForCausalLM", diff --git a/python/sglang/srt/layers/cp/zigzag.py b/python/sglang/srt/layers/cp/zigzag.py index 0f279fcb5..b7fe868d6 100644 --- a/python/sglang/srt/layers/cp/zigzag.py +++ b/python/sglang/srt/layers/cp/zigzag.py @@ -72,6 +72,8 @@ class ZigzagContextParallelMetadata(BaseContextParallelMetadata): # Per-sequence FlashAttention tensors (shape [bs] or [bs + 1]). kv_len_prev_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_next_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") cu_prev = [0] + list(accumulate(actual_seq_q_prev_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) assert len(split_list) == bs * cp_segment_num @@ -236,6 +240,12 @@ class ZigzagCPStrategy(ContextParallelStrategy): kv_len_next_tensor=torch.tensor( 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_list, device=device, dtype=torch.int32 ), @@ -301,7 +311,10 @@ class ZigzagCPStrategy(ContextParallelStrategy): ) def get_supported_attention_backend(self): - return [CPAttentionBackendKind.FLASH_ATTENTION] + return [ + CPAttentionBackendKind.FLASH_ATTENTION, + CPAttentionBackendKind.TRTLLM_MHA, + ] def run_attention( self, @@ -320,17 +333,25 @@ class ZigzagCPStrategy(ContextParallelStrategy): logical_tokens = meta.total_q_prev_tokens + meta.total_q_next_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( q_prev, meta.cu_seqlens_q_prev_tensor, meta.kv_len_prev_tensor, meta.max_seqlen_q_prev, + **prev_kwargs, ) result_next = attn_fn( q_next, meta.cu_seqlens_q_next_tensor, meta.kv_len_next_tensor, meta.max_seqlen_q_next, + **next_kwargs, ) result = torch.cat([result_prev, result_next], dim=0) pad_size = q.shape[0] - logical_tokens diff --git a/test/registered/cp/test_gpt_oss_4gpu_mxfp4_cp.py b/test/registered/cp/test_gpt_oss_4gpu_mxfp4_cp.py new file mode 100644 index 000000000..6ad3aaaa9 --- /dev/null +++ b/test/registered/cp/test_gpt_oss_4gpu_mxfp4_cp.py @@ -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()