Support GPT-OSS zigzag CP with TRTLLM-MHA (#31732)

This commit is contained in:
Baizhou Zhang
2026-07-20 00:48:04 -07:00
committed by GitHub
parent 2f14d6c6f2
commit 9668d9ea72
5 changed files with 143 additions and 33 deletions
@@ -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)
+4 -1
View File
@@ -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'}"
)
+1
View File
@@ -40,6 +40,7 @@ if TYPE_CHECKING:
CP_V2_DEFAULT_MODEL_CLASSES = frozenset(
{
"GptOssForCausalLM",
"MiMoV2FlashForCausalLM",
"MiMoV2ForCausalLM",
"Qwen3MoeForCausalLM",
+22 -1
View File
@@ -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
@@ -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()