[3/N][CP] Implement zigzag CP strategy (#28421)

This commit is contained in:
Baizhou Zhang
2026-06-18 15:10:30 -07:00
committed by GitHub
parent 9fc9d37f6d
commit e3026ef016
13 changed files with 1091 additions and 56 deletions
@@ -13,6 +13,8 @@ from sglang.srt.layers.attention.triton_ops.metadata import (
prepare_swa_spec_page_table_triton, prepare_swa_spec_page_table_triton,
) )
from sglang.srt.layers.attention.utils import assert_buffer_fits from sglang.srt.layers.attention.utils import assert_buffer_fits
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.radix_attention import AttentionType from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.layers.utils.cp_utils import ( from sglang.srt.layers.utils.cp_utils import (
cp_allgather_and_save_kv_cache, cp_allgather_and_save_kv_cache,
@@ -811,17 +813,25 @@ class FlashAttentionBackend(AttentionBackend):
elif is_cp_mode: elif is_cp_mode:
# Dense-MHA CP: k, v are still rank-local; backend # Dense-MHA CP: k, v are still rank-local; backend
# all-gathers and writes to the per-rank pool. # all-gathers and writes to the per-rank pool.
swa_loc = (
self.forward_metadata.swa_out_cache_loc
if self.use_sliding_window_kv_pool
else None
)
if is_cp_v2_active(forward_batch):
cp_strategy = get_cp_strategy()
assert cp_strategy is not None
cp_strategy.materialize_full_kv(
forward_batch, layer, k, v, swa_loc=swa_loc
)
else:
cp_allgather_and_save_kv_cache( cp_allgather_and_save_kv_cache(
forward_batch, forward_batch,
layer, layer,
k, k,
v, v,
self.attn_cp_size, self.attn_cp_size,
swa_loc=( swa_loc=swa_loc,
self.forward_metadata.swa_out_cache_loc
if self.use_sliding_window_kv_pool
else None
),
) )
else: else:
self.token_to_kv_pool.set_kv_buffer( self.token_to_kv_pool.set_kv_buffer(
@@ -967,9 +977,21 @@ class FlashAttentionBackend(AttentionBackend):
**kwargs, **kwargs,
) )
q_cp = q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim)
if is_cp_v2_active(forward_batch):
cp_strategy = get_cp_strategy()
assert cp_strategy is not None
result = cp_strategy.run_attention(
q_cp,
forward_batch,
self.device,
_fa_cp_attn,
attention_backend=CPAttentionBackendKind.FLASH_ATTENTION,
)
else:
result = cp_attn_forward_extend( result = cp_attn_forward_extend(
forward_batch, forward_batch,
q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim), q_cp,
self.device, self.device,
_fa_cp_attn, _fa_cp_attn,
) )
+4 -8
View File
@@ -185,6 +185,7 @@ class ContextParallelStrategy(ABC):
layer: Any, layer: Any,
k: Any, k: Any,
v: Any, v: Any,
swa_loc: Optional[Any] = None,
) -> None: ) -> None:
"""Write full-layout K/V to the backend cache if needed.""" """Write full-layout K/V to the backend cache if needed."""
@@ -235,7 +236,7 @@ def init_cp_strategy(server_args: ServerArgs) -> None:
) )
def _get_cp_strategy() -> Optional[ContextParallelStrategy]: def get_cp_strategy() -> Optional[ContextParallelStrategy]:
"""Return the configured strategy, initializing lazily on first call. """Return the configured strategy, initializing lazily on first call.
Subprocesses re-import this module with ``_STRATEGY = None`` and never Subprocesses re-import this module with ``_STRATEGY = None`` and never
@@ -257,20 +258,15 @@ def _get_cp_strategy() -> Optional[ContextParallelStrategy]:
return _STRATEGY return _STRATEGY
def get_cp_strategy() -> Optional[ContextParallelStrategy]:
"""Return the configured CP strategy for runtime dispatch."""
return _get_cp_strategy()
def get_cp_strategy_kind() -> ContextParallelStrategyKind: def get_cp_strategy_kind() -> ContextParallelStrategyKind:
strategy = _get_cp_strategy() strategy = get_cp_strategy()
if strategy is None: if strategy is None:
return ContextParallelStrategyKind.NONE return ContextParallelStrategyKind.NONE
return strategy.kind return strategy.kind
def is_cp_enabled() -> bool: def is_cp_enabled() -> bool:
return _get_cp_strategy() is not None return get_cp_strategy() is not None
def is_zigzag() -> bool: def is_zigzag() -> bool:
+3 -1
View File
@@ -99,7 +99,9 @@ class InterleaveCPStrategy(ContextParallelStrategy):
"Interleave attention dispatch will land in a follow-up PR" "Interleave attention dispatch will land in a follow-up PR"
) )
def materialize_full_kv(self, forward_batch, layer: Any, k: Any, v: Any) -> None: def materialize_full_kv(
self, forward_batch, layer: Any, k: Any, v: Any, swa_loc: Optional[Any] = None
) -> None:
raise NotImplementedError( raise NotImplementedError(
"Interleave KV materialization will land in a follow-up PR" "Interleave KV materialization will land in a follow-up PR"
) )
+101 -1
View File
@@ -12,13 +12,16 @@
# limitations under the License. # limitations under the License.
# ============================================================================== # ==============================================================================
"""Public import facade for context parallel strategy helpers.""" """Public import facade and runtime helpers for context parallel strategies."""
from typing import Any, Optional, Tuple
from sglang.srt.layers.cp.base import ( from sglang.srt.layers.cp.base import (
BaseContextParallelMetadata, BaseContextParallelMetadata,
ContextParallelStrategy, ContextParallelStrategy,
ContextParallelStrategyKind, ContextParallelStrategyKind,
CPAttentionBackendKind, CPAttentionBackendKind,
get_cp_strategy,
) )
from sglang.srt.layers.cp.interleave import ( from sglang.srt.layers.cp.interleave import (
InterleaveContextParallelMetadata, InterleaveContextParallelMetadata,
@@ -30,6 +33,96 @@ from sglang.srt.layers.cp.zigzag import (
ZigzagCPStrategy, ZigzagCPStrategy,
) )
CP_V2_DEFAULT_MODEL_CLASSES = frozenset(
{
"Qwen3MoeForCausalLM",
}
)
def enable_cp_v2() -> bool:
"""Return whether the CP-v2 path is enabled for this process."""
from sglang.srt.environ import envs
return bool(envs.SGLANG_ENABLE_CP_V2.get())
def is_cp_v2_active(forward_batch) -> bool:
"""Return whether the current forward batch is running through CP-v2."""
if not enable_cp_v2():
return False
forward_mode = getattr(forward_batch, "forward_mode", None)
if forward_mode is None or not forward_mode.is_context_parallel_extend():
return False
strategy = get_cp_strategy()
if strategy is None:
return False
input_ids = getattr(forward_batch, "input_ids", None)
if input_ids is None:
return False
return strategy.can_apply(len(input_ids), forward_batch)
def prepare_cp_forward(forward_batch) -> None:
"""Build CP-v2 metadata for an active context-parallel prefill batch."""
assert is_cp_v2_active(forward_batch)
strategy = get_cp_strategy()
assert strategy is not None
num_tokens = len(forward_batch.input_ids)
seq_lens_cpu = _to_int_list(getattr(forward_batch, "seq_lens_cpu", None))
extend_lens_cpu = _to_int_list(getattr(forward_batch, "extend_seq_lens_cpu", None))
forward_batch.attn_cp_metadata = strategy.build_metadata(
num_tokens=num_tokens,
seqs_len=seq_lens_cpu,
extend_seqs_len=extend_lens_cpu,
)
def cp_split_before_forward(
complete_hidden_states: Any,
complete_position_ids: Any,
forward_batch,
) -> Tuple[Optional[Any], Optional[Any]]:
"""Shard embeddings and positions for CP-v2 model-runner forwarding."""
assert is_cp_v2_active(forward_batch)
strategy = get_cp_strategy()
assert strategy is not None
assert complete_hidden_states is not None
assert getattr(forward_batch, "attn_cp_metadata", None) is not None
return (
strategy.shard_hidden_states(complete_hidden_states, forward_batch),
strategy.shard_position_ids(complete_position_ids, forward_batch),
)
def cp_gather_after_forward(x: Any, forward_batch, stream: Optional[Any] = None):
"""Gather CP-v2 hidden states at the model boundary when this batch is active."""
assert is_cp_v2_active(forward_batch)
strategy = get_cp_strategy()
assert strategy is not None
if isinstance(x, tuple):
hidden_states, *rest = x
hidden_states = strategy.gather_hidden_states(
hidden_states, forward_batch, stream
)
return (hidden_states, *rest)
return strategy.gather_hidden_states(x, forward_batch, stream)
def _to_int_list(values) -> Optional[list[int]]:
if values is None:
return None
if hasattr(values, "tolist"):
values = values.tolist()
return [int(x) for x in values]
__all__ = [ __all__ = [
"BaseContextParallelMetadata", "BaseContextParallelMetadata",
"CPAttentionBackendKind", "CPAttentionBackendKind",
@@ -40,4 +133,11 @@ __all__ = [
"InterleaveContextParallelMetadata", "InterleaveContextParallelMetadata",
"ZigzagCPStrategy", "ZigzagCPStrategy",
"ZigzagContextParallelMetadata", "ZigzagContextParallelMetadata",
"CP_V2_DEFAULT_MODEL_CLASSES",
"enable_cp_v2",
"get_cp_strategy",
"is_cp_v2_active",
"cp_gather_after_forward",
"cp_split_before_forward",
"prepare_cp_forward",
] ]
+264 -15
View File
@@ -30,15 +30,29 @@ After all-gather, the blocks are reranged back to their original order:
from __future__ import annotations from __future__ import annotations
from contextlib import nullcontext
from dataclasses import dataclass from dataclasses import dataclass
from itertools import accumulate
from typing import Any, List, Optional from typing import Any, List, Optional
import torch
import torch.nn.functional as F
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory,
)
from sglang.srt.layers.cp.base import ( from sglang.srt.layers.cp.base import (
BaseContextParallelMetadata, BaseContextParallelMetadata,
ContextParallelStrategy, ContextParallelStrategy,
ContextParallelStrategyKind, ContextParallelStrategyKind,
CPAttentionBackendKind, CPAttentionBackendKind,
) )
from sglang.srt.layers.dp_attention import (
get_attention_cp_group,
is_allocation_symmetric,
)
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
@dataclass @dataclass
@@ -85,7 +99,13 @@ class ZigzagCPStrategy(ContextParallelStrategy):
if self.cp_size <= 1 or num_tokens < self.cp_size * 2: if self.cp_size <= 1 or num_tokens < self.cp_size * 2:
return False return False
forward_mode = getattr(forward_batch, "forward_mode", None) forward_mode = getattr(forward_batch, "forward_mode", None)
return forward_mode is None or forward_mode.is_context_parallel_extend() if forward_mode is not None and not forward_mode.is_context_parallel_extend():
return False
extend_lens = getattr(forward_batch, "extend_seq_lens_cpu", None)
if extend_lens is None:
return True
return all(int(length) >= self.cp_size * 2 for length in extend_lens)
def build_metadata( def build_metadata(
self, self,
@@ -93,32 +113,193 @@ class ZigzagCPStrategy(ContextParallelStrategy):
seqs_len: Optional[List[int]], seqs_len: Optional[List[int]],
extend_seqs_len: Optional[List[int]] = None, extend_seqs_len: Optional[List[int]] = None,
) -> ZigzagContextParallelMetadata: ) -> ZigzagContextParallelMetadata:
if extend_seqs_len is None:
extend_seqs_len = seqs_len or [num_tokens]
extend_seqs_len = [int(x) for x in extend_seqs_len]
pad_len = int(num_tokens) - sum(extend_seqs_len)
if pad_len > 0:
extend_seqs_len[-1] += pad_len
if seqs_len is not None and len(seqs_len) == len(extend_seqs_len):
seqs_len = list(seqs_len)
seqs_len[-1] += pad_len
bs = len(extend_seqs_len)
cp_segment_num = self.cp_size * 2
if seqs_len is not None and len(seqs_len) == bs:
prefix_offsets = [
max(int(seqs_len[i]) - extend_seqs_len[i], 0) for i in range(bs)
]
else:
prefix_offsets = [0] * bs
# TODO: move these per-request layout/index computations to a Triton
# kernel if Python-side metadata construction becomes a bottleneck.
per_seq_block_sizes: List[List[int]] = []
split_list: List[int] = []
for length in extend_seqs_len:
base = length // cp_segment_num
rem = length % cp_segment_num
block_sizes = [
base + 1 if block_id < rem else base
for block_id in range(cp_segment_num)
]
per_seq_block_sizes.append(block_sizes)
split_list.extend(block_sizes)
per_rank_actual_token = []
for rank in range(self.cp_size):
per_rank_actual_token.append(
sum(
block_sizes[rank] + block_sizes[cp_segment_num - 1 - rank]
for block_sizes in per_seq_block_sizes
)
)
max_rank_len = [max(per_rank_actual_token)] * self.cp_size
cp_rank = self.cp_rank
zigzag_index = list(
range(cp_rank, cp_rank + bs * cp_segment_num, cp_segment_num)
) + list(
range(
cp_segment_num - cp_rank - 1,
bs * cp_segment_num,
cp_segment_num,
)
)
cp_reverse_index: List[int] = []
for batch_id in range(bs):
cp_reverse_index.extend(
list(range(batch_id, cp_segment_num * bs, 2 * bs))
+ list(
range(
(cp_segment_num - 1) * bs + batch_id,
0,
-2 * bs,
)
)
)
reverse_split_len: List[int] = []
for rank in range(self.cp_size):
for batch_id in range(bs):
reverse_split_len.append(per_seq_block_sizes[batch_id][rank])
for batch_id in range(bs):
reverse_split_len.append(
per_seq_block_sizes[batch_id][cp_segment_num - 1 - rank]
)
kv_len_prev_list: List[int] = []
kv_len_next_list: List[int] = []
actual_seq_q_prev_list: List[int] = []
actual_seq_q_next_list: List[int] = []
for batch_id, block_sizes in enumerate(per_seq_block_sizes):
kv_len_prev_list.append(
prefix_offsets[batch_id] + sum(block_sizes[: cp_rank + 1])
)
kv_len_next_list.append(
prefix_offsets[batch_id] + sum(block_sizes[: cp_segment_num - cp_rank])
)
actual_seq_q_prev_list.append(block_sizes[cp_rank])
actual_seq_q_next_list.append(block_sizes[cp_segment_num - cp_rank - 1])
from sglang.srt.server_args import get_global_server_args
try:
device = torch.device(get_global_server_args().device)
except Exception:
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))
total_seq_lens = sum(extend_seqs_len)
assert len(split_list) == bs * cp_segment_num
assert sum(split_list) == total_seq_lens
assert len(zigzag_index) == 2 * bs
assert len(cp_reverse_index) == bs * cp_segment_num
assert sorted(cp_reverse_index) == list(range(bs * cp_segment_num))
assert sum(per_rank_actual_token) == total_seq_lens
return ZigzagContextParallelMetadata( return ZigzagContextParallelMetadata(
total_seq_lens=sum(extend_seqs_len or seqs_len or [num_tokens]), split_list=split_list,
bs=len(extend_seqs_len or seqs_len or [num_tokens]), zigzag_index=zigzag_index,
cp_reverse_index=cp_reverse_index,
reverse_split_len=reverse_split_len,
per_rank_actual_token=per_rank_actual_token,
max_rank_len=max_rank_len,
kv_len_prev_tensor=torch.tensor(
kv_len_prev_list, device=device, dtype=torch.int32
),
kv_len_next_tensor=torch.tensor(
kv_len_next_list, device=device, dtype=torch.int32
),
actual_seq_q_prev_tensor=torch.tensor(
actual_seq_q_prev_list, device=device, dtype=torch.int32
),
actual_seq_q_next_tensor=torch.tensor(
actual_seq_q_next_list, device=device, dtype=torch.int32
),
cu_seqlens_q_prev_tensor=torch.tensor(
cu_prev, device=device, dtype=torch.int32
),
cu_seqlens_q_next_tensor=torch.tensor(
cu_next, device=device, dtype=torch.int32
),
total_q_prev_tokens=cu_prev[-1],
total_q_next_tokens=cu_next[-1],
max_seqlen_q_prev=(
max(actual_seq_q_prev_list) if actual_seq_q_prev_list else 0
),
max_seqlen_q_next=(
max(actual_seq_q_next_list) if actual_seq_q_next_list else 0
),
kv_len_prev_list=kv_len_prev_list,
kv_len_next_list=kv_len_next_list,
actual_seq_q_prev_list=actual_seq_q_prev_list,
actual_seq_q_next_list=actual_seq_q_next_list,
total_seq_lens=total_seq_lens,
bs=bs,
) )
def shard_hidden_states(self, x: Any, forward_batch) -> Any: def shard_hidden_states(self, x: Any, forward_batch) -> Any:
raise NotImplementedError( chunks = torch.split(x, forward_batch.attn_cp_metadata.split_list, dim=0)
"Zigzag hidden-state sharding will land in a follow-up PR" return torch.cat(
[chunks[i] for i in forward_batch.attn_cp_metadata.zigzag_index], dim=0
) )
def shard_position_ids(self, positions: Any, forward_batch) -> Any: def shard_position_ids(self, positions: Any, forward_batch) -> Any:
raise NotImplementedError( chunks = torch.split(
"Zigzag position-id sharding will land in a follow-up PR" positions, forward_batch.attn_cp_metadata.split_list, dim=-1
)
return torch.cat(
[chunks[i] for i in forward_batch.attn_cp_metadata.zigzag_index], dim=-1
) )
def gather_hidden_states( def gather_hidden_states(
self, x: Any, forward_batch, stream: Optional[Any] = None self, x: Any, forward_batch, stream: Optional[Any] = None
) -> Any: ) -> Any:
raise NotImplementedError( gathered = self._all_gather_reorganized(x, forward_batch, stream)
"Zigzag hidden-state gather will land in a follow-up PR" chunks = torch.split(
gathered, forward_batch.attn_cp_metadata.reverse_split_len, dim=0
)
return torch.cat(
[chunks[i] for i in forward_batch.attn_cp_metadata.cp_reverse_index], dim=0
) )
def gather_kv_cache( def gather_kv_cache(
self, x: Any, forward_batch, stream: Optional[Any] = None self, x: Any, forward_batch, stream: Optional[Any] = None
) -> Any: ) -> Any:
raise NotImplementedError("Zigzag KV gather will land in a follow-up PR") gathered = self._all_gather_reorganized(x, forward_batch, stream)
chunks = torch.split(
gathered, forward_batch.attn_cp_metadata.reverse_split_len, dim=0
)
return torch.cat(
[chunks[i] for i in forward_batch.attn_cp_metadata.cp_reverse_index], dim=0
)
def get_supported_attention_backend(self):
return [CPAttentionBackendKind.FLASH_ATTENTION]
def run_attention( def run_attention(
self, self,
@@ -128,11 +309,79 @@ class ZigzagCPStrategy(ContextParallelStrategy):
attn_fn, attn_fn,
attention_backend: CPAttentionBackendKind = CPAttentionBackendKind.FLASH_ATTENTION, attention_backend: CPAttentionBackendKind = CPAttentionBackendKind.FLASH_ATTENTION,
) -> Any: ) -> Any:
raise NotImplementedError( assert (
"Zigzag attention dispatch will land in a follow-up PR" attention_backend in self.get_supported_attention_backend()
), f"{self.name} CP does not support {attention_backend=}"
meta = forward_batch.attn_cp_metadata
q_prev = q[: meta.total_q_prev_tokens]
q_next = q[meta.total_q_prev_tokens :]
result_prev = attn_fn(
q_prev,
meta.cu_seqlens_q_prev_tensor,
meta.kv_len_prev_tensor,
meta.max_seqlen_q_prev,
)
result_next = attn_fn(
q_next,
meta.cu_seqlens_q_next_tensor,
meta.kv_len_next_tensor,
meta.max_seqlen_q_next,
)
return torch.cat([result_prev, result_next], dim=0)
def materialize_full_kv(
self, forward_batch, layer: Any, k: Any, v: Any, swa_loc: Optional[Any] = None
) -> None:
cache_loc = (
forward_batch.out_cache_loc
if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc
)
key_cache_full = self.gather_kv_cache(
k.contiguous(), forward_batch, torch.cuda.current_stream()
)
value_cache_full = self.gather_kv_cache(
v.contiguous(), forward_batch, torch.cuda.current_stream()
)
get_token_to_kv_pool().set_kv_buffer(
layer,
KVWriteLoc(cache_loc, swa_loc),
key_cache_full,
value_cache_full,
layer.k_scale,
layer.v_scale,
) )
def materialize_full_kv(self, forward_batch, layer: Any, k: Any, v: Any) -> None: def _all_gather_reorganized(self, x: torch.Tensor, forward_batch, stream):
raise NotImplementedError( meta = forward_batch.attn_cp_metadata
"Zigzag KV materialization will land in a follow-up PR" max_len = meta.max_rank_len[0]
pad_size = max_len - x.shape[0]
if pad_size > 0:
padding = [0, 0] * (x.ndim - 1) + [0, pad_size]
x = F.pad(x, padding, mode="constant", value=0)
group = get_attention_cp_group()
ctx = (
use_symmetric_memory(group, disabled=not is_allocation_symmetric())
if x.is_cuda
else nullcontext()
)
with ctx:
gathered = torch.empty(
max_len * self.cp_size,
*x.shape[1:],
device=x.device,
dtype=x.dtype,
)
group.cp_all_gather_into_tensor_async(gathered, x, stream)
chunks = torch.split(gathered, meta.max_rank_len, dim=0)
return torch.cat(
[
chunks[rank][:per_rank_len]
for rank, per_rank_len in enumerate(meta.per_rank_actual_token)
],
dim=0,
) )
@@ -123,6 +123,13 @@ from sglang.srt.layers.attention.attention_registry import (
) )
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
from sglang.srt.layers.cp.utils import (
cp_gather_after_forward,
cp_split_before_forward,
get_cp_strategy,
is_cp_v2_active,
prepare_cp_forward,
)
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
DpPaddingMode, DpPaddingMode,
get_attention_tp_group, get_attention_tp_group,
@@ -3411,6 +3418,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.prefill_cuda_graph_runner is not None self.prefill_cuda_graph_runner is not None
and self.prefill_cuda_graph_runner.can_run(forward_batch) and self.prefill_cuda_graph_runner.can_run(forward_batch)
) )
if get_cp_strategy() is not None:
can_run_graph = False
if can_run_graph: if can_run_graph:
# TODO: device_timer.wrap is too broad here — it also includes # TODO: device_timer.wrap is too broad here — it also includes
# replay_prepare time. Move timing into the prefill cuda graph # replay_prepare time. Move timing into the prefill cuda graph
@@ -3434,6 +3443,21 @@ class ModelRunner(ModelRunnerKVCacheMixin):
# e.g. Moss-VL's prefill cross-attention custom mask. # e.g. Moss-VL's prefill cross-attention custom mask.
self.model.prepare_forward_batch(forward_batch) self.model.prepare_forward_batch(forward_batch)
self.attn_backend.init_forward_metadata(forward_batch) self.attn_backend.init_forward_metadata(forward_batch)
cp_v2_active = is_cp_v2_active(forward_batch)
forward_positions = forward_batch.positions
if cp_v2_active:
prepare_cp_forward(forward_batch)
complete_hidden_states = kwargs.get("input_embeds")
if complete_hidden_states is None:
embed_layer = self.model.get_input_embeddings()
complete_hidden_states = embed_layer(forward_batch.input_ids)
sharded_hidden_states, sharded_positions = cp_split_before_forward(
complete_hidden_states,
forward_batch.positions,
forward_batch,
)
kwargs["input_embeds"] = sharded_hidden_states
forward_positions = sharded_positions
ctx = ( ctx = (
self.device_timer.wrap(metadata={"category": "extend"}) self.device_timer.wrap(metadata={"category": "extend"})
@@ -3441,7 +3465,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
else contextlib.nullcontext() else contextlib.nullcontext()
) )
with ctx: with ctx:
if _is_hip and self.prefill_cuda_graph_runner is not None: if (
_is_hip
and self.prefill_cuda_graph_runner is not None
and not cp_v2_active
):
# AMD/HIP: when PCG is enabled but the batch exceeds max captured # AMD/HIP: when PCG is enabled but the batch exceeds max captured
# size, run eagerly under enable_tc_piecewise_cuda_graph() and # size, run eagerly under enable_tc_piecewise_cuda_graph() and
# set_tc_piecewise_forward_context() so that (a) Dynamo guards on # set_tc_piecewise_forward_context() so that (a) Dynamo guards on
@@ -3461,14 +3489,47 @@ class ModelRunner(ModelRunnerKVCacheMixin):
): ):
ret = self.model.forward( ret = self.model.forward(
forward_batch.input_ids, forward_batch.input_ids,
forward_batch.positions, forward_positions,
forward_batch, forward_batch,
**kwargs, **kwargs,
) )
elif cp_v2_active:
hidden_states = self.model.model(
forward_batch.input_ids,
forward_positions,
forward_batch,
input_embeds=kwargs.get("input_embeds"),
pp_proxy_tensors=kwargs.get("pp_proxy_tensors"),
)
aux_hidden_states = None
capture_aux_hidden_states = getattr(
self.model, "capture_aux_hidden_states", False
)
if capture_aux_hidden_states:
hidden_states, aux_hidden_states = hidden_states
if self.model.pp_group.is_last_rank:
hidden_states = cp_gather_after_forward(
hidden_states,
forward_batch,
torch.cuda.current_stream(),
)
ret = self.model.logits_processor(
forward_batch.input_ids,
hidden_states,
self.model.lm_head,
forward_batch,
aux_hidden_states,
)
elif capture_aux_hidden_states:
ret = hidden_states, aux_hidden_states
else:
ret = hidden_states
else: else:
ret = self.model.forward( ret = self.model.forward(
forward_batch.input_ids, forward_batch.input_ids,
forward_batch.positions, forward_positions,
forward_batch, forward_batch,
**kwargs, **kwargs,
) )
+3
View File
@@ -42,6 +42,7 @@ from sglang.srt.layers.communicator import (
LayerScatterModes, LayerScatterModes,
ScatterMode, ScatterMode,
) )
from sglang.srt.layers.cp.utils import is_cp_v2_active
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
is_dp_attention_enabled, is_dp_attention_enabled,
) )
@@ -889,6 +890,7 @@ class Qwen2MoeModel(nn.Module):
if ( if (
is_prefill_context_parallel_enabled() is_prefill_context_parallel_enabled()
and not is_cp_v2_active(forward_batch)
and forward_batch.forward_mode.is_context_parallel_extend() and forward_batch.forward_mode.is_context_parallel_extend()
and forward_batch.attn_cp_metadata is not None and forward_batch.attn_cp_metadata is not None
): ):
@@ -944,6 +946,7 @@ class Qwen2MoeModel(nn.Module):
if ( if (
self.pp_group.is_last_rank self.pp_group.is_last_rank
and not is_cp_v2_active(forward_batch)
and is_prefill_context_parallel_enabled() and is_prefill_context_parallel_enabled()
and forward_batch.forward_mode.is_context_parallel_extend() and forward_batch.forward_mode.is_context_parallel_extend()
and forward_batch.attn_cp_metadata is not None and forward_batch.attn_cp_metadata is not None
+2 -1
View File
@@ -34,6 +34,7 @@ from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_r
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
from sglang.srt.layers.cp.utils import is_cp_v2_active
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
QKVParallelLinear, QKVParallelLinear,
@@ -988,7 +989,7 @@ class Qwen3MoeForCausalLM(nn.Module):
input_embeds: torch.Tensor = None, input_embeds: torch.Tensor = None,
pp_proxy_tensors: Optional[PPProxyTensors] = None, pp_proxy_tensors: Optional[PPProxyTensors] = None,
) -> torch.Tensor: ) -> torch.Tensor:
if is_prefill_context_parallel_enabled(): if is_prefill_context_parallel_enabled() and not is_cp_v2_active(forward_batch):
if can_cp_split(len(input_ids), self.attn_cp_size, forward_batch): if can_cp_split(len(input_ids), self.attn_cp_size, forward_batch):
forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( forward_batch.attn_cp_metadata = prepare_context_parallel_metadata(
len(input_ids), len(input_ids),
+11
View File
@@ -3436,6 +3436,17 @@ class ServerArgs:
self.prefill_cp_mode = mode self.prefill_cp_mode = mode
def _handle_context_parallelism(self): def _handle_context_parallelism(self):
if parse_connector_type(self.model_path) != ConnectorType.INSTANCE:
from sglang.srt.layers.cp.utils import CP_V2_DEFAULT_MODEL_CLASSES
model_config = self.get_model_config()
model_arch = model_config.hf_config.architectures[0]
if (
model_arch in CP_V2_DEFAULT_MODEL_CLASSES
and not envs.SGLANG_ENABLE_CP_V2.is_set()
):
envs.SGLANG_ENABLE_CP_V2.set(True)
if self.enable_prefill_cp and self.cp_strategy is None: if self.enable_prefill_cp and self.cp_strategy is None:
raise ValueError( raise ValueError(
"--cp-strategy must be set when --enable-prefill-cp is enabled." "--cp-strategy must be set when --enable-prefill-cp is enabled."
+386 -1
View File
@@ -2,6 +2,8 @@ import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import patch from unittest.mock import patch
import torch
from sglang.srt.layers.cp.base import ( from sglang.srt.layers.cp.base import (
ContextParallelStrategyKind, ContextParallelStrategyKind,
get_cp_strategy, get_cp_strategy,
@@ -11,10 +13,31 @@ from sglang.srt.layers.cp.base import (
is_interleave, is_interleave,
is_zigzag, is_zigzag,
) )
from sglang.srt.layers.cp.utils import (
cp_split_before_forward,
enable_cp_v2,
is_cp_v2_active,
)
from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy
from sglang.srt.runtime_context import get_parallel
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="base-a-test-cpu") register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class _ExtendMode:
def is_context_parallel_extend(self):
return True
class _FakeCPGroup:
def __init__(self, all_rank_tensors):
self.all_rank_tensors = all_rank_tensors
def cp_all_gather_into_tensor_async(self, output, input_tensor, stream):
del input_tensor, stream
torch.cat(self.all_rank_tensors, dim=0, out=output)
class TestCPStrategyUnit(CustomTestCase): class TestCPStrategyUnit(CustomTestCase):
@@ -70,5 +93,367 @@ class TestCPStrategyUnit(CustomTestCase):
self.assertIsNotNone(get_cp_strategy()) self.assertIsNotNone(get_cp_strategy())
class TestCPZigzagStrategy(CustomTestCase):
def setUp(self):
init_cp_strategy(
SimpleNamespace(
enable_prefill_cp=True,
cp_strategy="zigzag",
attn_cp_size=4,
attention_backend="fa3",
)
)
def tearDown(self):
init_cp_strategy(SimpleNamespace(enable_prefill_cp=False))
def _metadata_for_rank(self, rank, *, cp_size, seq_lens, extend_seq_lens):
strategy = ZigzagCPStrategy(cp_size=cp_size)
with get_parallel().override(attn_cp_rank=rank):
return strategy.build_metadata(
num_tokens=sum(extend_seq_lens),
seqs_len=seq_lens,
extend_seqs_len=extend_seq_lens,
)
def _forward_batch(self, metadata, extend_seq_lens):
return SimpleNamespace(
input_ids=torch.arange(sum(extend_seq_lens)),
forward_mode=_ExtendMode(),
extend_seq_lens_cpu=extend_seq_lens,
attn_cp_metadata=metadata,
)
def test_enable_cp_v2_and_is_cp_v2_active(self):
active_batch = SimpleNamespace(
input_ids=torch.arange(8),
forward_mode=_ExtendMode(),
extend_seq_lens_cpu=[8],
)
inactive_batch = SimpleNamespace(
input_ids=torch.arange(7),
forward_mode=_ExtendMode(),
extend_seq_lens_cpu=[7],
)
with patch(
"sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", return_value=False
):
self.assertFalse(enable_cp_v2())
self.assertFalse(is_cp_v2_active(active_batch))
with patch(
"sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", return_value=True
):
self.assertTrue(enable_cp_v2())
self.assertTrue(is_cp_v2_active(active_batch))
self.assertFalse(is_cp_v2_active(inactive_batch))
def _expected_metadata(self, *, rank, cp_size, seq_lens, extend_seq_lens):
bs = len(extend_seq_lens)
cp_segment_num = cp_size * 2
prefix_offsets = [
max(int(seq_lens[i]) - int(extend_seq_lens[i]), 0) for i in range(bs)
]
per_seq_block_sizes = []
split_list = []
for length in extend_seq_lens:
base = length // cp_segment_num
rem = length % cp_segment_num
block_sizes = [
base + 1 if block_id < rem else base
for block_id in range(cp_segment_num)
]
per_seq_block_sizes.append(block_sizes)
split_list.extend(block_sizes)
per_rank_actual_token = [
sum(
block_sizes[rank_id] + block_sizes[cp_segment_num - 1 - rank_id]
for block_sizes in per_seq_block_sizes
)
for rank_id in range(cp_size)
]
max_rank_len = [max(per_rank_actual_token)] * cp_size
zigzag_index = list(range(rank, rank + bs * cp_segment_num, cp_segment_num))
zigzag_index += list(
range(cp_segment_num - rank - 1, bs * cp_segment_num, cp_segment_num)
)
cp_reverse_index = []
for batch_id in range(bs):
cp_reverse_index.extend(
list(range(batch_id, cp_segment_num * bs, 2 * bs))
+ list(range((cp_segment_num - 1) * bs + batch_id, 0, -2 * bs))
)
reverse_split_len = []
for rank_id in range(cp_size):
for batch_id in range(bs):
reverse_split_len.append(per_seq_block_sizes[batch_id][rank_id])
for batch_id in range(bs):
reverse_split_len.append(
per_seq_block_sizes[batch_id][cp_segment_num - 1 - rank_id]
)
kv_len_prev_list = []
kv_len_next_list = []
actual_seq_q_prev_list = []
actual_seq_q_next_list = []
for batch_id, block_sizes in enumerate(per_seq_block_sizes):
kv_len_prev_list.append(
prefix_offsets[batch_id] + sum(block_sizes[: rank + 1])
)
kv_len_next_list.append(
prefix_offsets[batch_id] + sum(block_sizes[: cp_segment_num - rank])
)
actual_seq_q_prev_list.append(block_sizes[rank])
actual_seq_q_next_list.append(block_sizes[cp_segment_num - rank - 1])
return {
"bs": bs,
"total_seq_lens": sum(extend_seq_lens),
"split_list": split_list,
"zigzag_index": zigzag_index,
"per_rank_actual_token": per_rank_actual_token,
"max_rank_len": max_rank_len,
"reverse_split_len": reverse_split_len,
"cp_reverse_index": cp_reverse_index,
"kv_len_prev_list": kv_len_prev_list,
"kv_len_next_list": kv_len_next_list,
"actual_seq_q_prev_list": actual_seq_q_prev_list,
"actual_seq_q_next_list": actual_seq_q_next_list,
}
def _assert_metadata_matches(self, metadata, expected):
self.assertEqual(metadata.bs, expected["bs"])
self.assertEqual(metadata.total_seq_lens, expected["total_seq_lens"])
self.assertEqual(metadata.split_list, expected["split_list"])
self.assertEqual(metadata.zigzag_index, expected["zigzag_index"])
self.assertEqual(
metadata.per_rank_actual_token, expected["per_rank_actual_token"]
)
self.assertEqual(metadata.max_rank_len, expected["max_rank_len"])
self.assertEqual(metadata.reverse_split_len, expected["reverse_split_len"])
self.assertEqual(metadata.cp_reverse_index, expected["cp_reverse_index"])
self.assertEqual(metadata.kv_len_prev_list, expected["kv_len_prev_list"])
self.assertEqual(metadata.kv_len_next_list, expected["kv_len_next_list"])
self.assertEqual(
metadata.actual_seq_q_prev_list, expected["actual_seq_q_prev_list"]
)
self.assertEqual(
metadata.actual_seq_q_next_list, expected["actual_seq_q_next_list"]
)
self.assertEqual(
metadata.cu_seqlens_q_prev_tensor.cpu().tolist(),
[0]
+ list(
torch.tensor(expected["actual_seq_q_prev_list"]).cumsum(dim=0).tolist()
),
)
self.assertEqual(
metadata.cu_seqlens_q_next_tensor.cpu().tolist(),
[0]
+ list(
torch.tensor(expected["actual_seq_q_next_list"]).cumsum(dim=0).tolist()
),
)
def _padded_rank_tensors(self, x, *, cp_size, seq_lens, extend_seq_lens):
per_rank = []
metas = []
for rank in range(cp_size):
metadata = self._metadata_for_rank(
rank,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
metas.append(metadata)
fb = self._forward_batch(metadata, extend_seq_lens)
local = ZigzagCPStrategy(cp_size=cp_size).shard_hidden_states(x, fb)
pad = metadata.max_rank_len[0] - local.shape[0]
if pad:
local = torch.nn.functional.pad(
local,
[0, 0] * (local.ndim - 1) + [0, pad],
)
per_rank.append(local)
return metas, per_rank
def test_zigzag_metadata_for_batched_sequences(self):
cases = [
(4, [11, 13], [9, 10]),
(2, [8], [8]),
(4, [100000, 200000, 80], [100000, 200000, 64]),
(4, [100005, 200011, 25], [100000, 200000, 16]),
]
for cp_size, seq_lens, extend_seq_lens in cases:
for rank in range(cp_size):
with self.subTest(
cp_size=cp_size,
rank=rank,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
):
metadata = self._metadata_for_rank(
rank,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
expected = self._expected_metadata(
rank=rank,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
self._assert_metadata_matches(metadata, expected)
def test_zigzag_shards_hidden_states_and_position_ids(self):
cp_size = 4
seq_lens = [11, 13]
extend_seq_lens = [9, 10]
x = torch.arange(sum(extend_seq_lens) * 2).view(sum(extend_seq_lens), 2)
positions = torch.arange(sum(extend_seq_lens))
for rank in range(cp_size):
metadata = self._metadata_for_rank(
rank,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
fb = self._forward_batch(metadata, extend_seq_lens)
strategy = ZigzagCPStrategy(cp_size=cp_size)
chunks = torch.split(x, metadata.split_list, dim=0)
position_chunks = torch.split(positions, metadata.split_list, dim=-1)
expected_x = torch.cat([chunks[i] for i in metadata.zigzag_index], dim=0)
expected_positions = torch.cat(
[position_chunks[i] for i in metadata.zigzag_index], dim=-1
)
local_x = strategy.shard_hidden_states(x, fb)
local_positions = strategy.shard_position_ids(positions, fb)
with patch(
"sglang.srt.environ.envs.SGLANG_ENABLE_CP_V2.get", return_value=True
):
helper_x, helper_positions = cp_split_before_forward(
x,
positions,
fb,
)
self.assertTrue(torch.equal(local_x, expected_x))
self.assertTrue(torch.equal(local_positions, expected_positions))
self.assertTrue(torch.equal(helper_x, expected_x))
self.assertTrue(torch.equal(helper_positions, expected_positions))
def test_zigzag_gathers_hidden_states_to_original_order(self):
cp_size = 4
seq_lens = [11, 13]
extend_seq_lens = [9, 10]
x = torch.arange(sum(extend_seq_lens) * 2).view(sum(extend_seq_lens), 2)
metas, padded_rank_tensors = self._padded_rank_tensors(
x,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
for rank in range(cp_size):
local_x = padded_rank_tensors[rank][
: metas[rank].per_rank_actual_token[rank]
]
fb = self._forward_batch(metas[rank], extend_seq_lens)
with (
patch(
"sglang.srt.layers.cp.zigzag.get_attention_cp_group",
return_value=_FakeCPGroup(padded_rank_tensors),
),
patch(
"sglang.srt.distributed.device_communicators.pynccl_allocator.use_symmetric_memory",
return_value=torch.no_grad(),
),
):
gathered = ZigzagCPStrategy(cp_size=cp_size).gather_hidden_states(
local_x, fb, stream=None
)
self.assertTrue(torch.equal(gathered, x))
def test_zigzag_gathers_kv_cache_to_original_order(self):
cp_size = 4
seq_lens = [11, 13]
extend_seq_lens = [9, 10]
kv = torch.arange(sum(extend_seq_lens) * 2 * 3).view(sum(extend_seq_lens), 2, 3)
metas, padded_rank_tensors = self._padded_rank_tensors(
kv,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
for rank in range(cp_size):
local_kv = padded_rank_tensors[rank][
: metas[rank].per_rank_actual_token[rank]
]
fb = self._forward_batch(metas[rank], extend_seq_lens)
with (
patch(
"sglang.srt.layers.cp.zigzag.get_attention_cp_group",
return_value=_FakeCPGroup(padded_rank_tensors),
),
patch(
"sglang.srt.distributed.device_communicators.pynccl_allocator.use_symmetric_memory",
return_value=torch.no_grad(),
),
):
gathered = ZigzagCPStrategy(cp_size=cp_size).gather_kv_cache(
local_kv, fb, stream=None
)
self.assertTrue(torch.equal(gathered, kv))
def test_zigzag_attention_dispatch_runs_prev_then_next(self):
cp_size = 2
seq_lens = [8]
extend_seq_lens = [8]
metadata = self._metadata_for_rank(
0,
cp_size=cp_size,
seq_lens=seq_lens,
extend_seq_lens=extend_seq_lens,
)
fb = SimpleNamespace(attn_cp_metadata=metadata)
q = torch.arange(4 * 2).view(4, 2)
calls = []
def attn_fn(q_chunk, cu_seqlens_q, cache_seqlens, max_seqlen_q):
calls.append(
(
q_chunk.clone(),
cu_seqlens_q.clone(),
cache_seqlens.clone(),
max_seqlen_q,
)
)
return q_chunk + 100
out = ZigzagCPStrategy(cp_size=cp_size).run_attention(
q, fb, device=torch.device("cpu"), attn_fn=attn_fn
)
self.assertEqual(len(calls), 2)
self.assertTrue(torch.equal(calls[0][0], q[:2]))
self.assertTrue(torch.equal(calls[1][0], q[2:]))
self.assertTrue(torch.equal(out, q + 100))
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -11,17 +11,17 @@ from sglang.test.test_utils import (
popen_launch_server, popen_launch_server,
) )
register_cuda_ci(est_time=261, stage="extra-b", runner_config="4-gpu-h100") register_cuda_ci(est_time=260, stage="extra-b", runner_config="4-gpu-h100")
QWEN3_30B_MODEL_PATH = "Qwen/Qwen3-30B-A3B-FP8" GQA_MODEL_PATH = "Qwen/Qwen3-30B-A3B-FP8"
GSM8K_BASELINE_ACCURACY = 0.85 GSM8K_BASELINE_ACCURACY = 0.93
class TestQwen330B(CustomTestCase): class TestGQACP2TP2EP2(CustomTestCase):
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
cls.model = QWEN3_30B_MODEL_PATH cls.model = GQA_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server( cls.process = popen_launch_server(
cls.model, cls.model,
@@ -46,10 +46,12 @@ class TestQwen330B(CustomTestCase):
"--model-loader-extra-config", "--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}', '{"enable_multithread_load": true, "num_threads": 64}',
], ],
env={"SGLANG_ENABLE_CP_V2": "0"},
) )
@classmethod @classmethod
def tearDownClass(cls): def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid) kill_process_tree(cls.process.pid)
def test_gsm8k(self): def test_gsm8k(self):
@@ -73,10 +75,10 @@ class TestQwen330B(CustomTestCase):
self.assertGreaterEqual(metrics["score"], GSM8K_BASELINE_ACCURACY) self.assertGreaterEqual(metrics["score"], GSM8K_BASELINE_ACCURACY)
class TestQwen330BCP(CustomTestCase): class TestGQACPTP2CP2EP4(CustomTestCase):
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
cls.model = QWEN3_30B_MODEL_PATH cls.model = GQA_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server( cls.process = popen_launch_server(
cls.model, cls.model,
@@ -101,10 +103,12 @@ class TestQwen330BCP(CustomTestCase):
"--model-loader-extra-config", "--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}', '{"enable_multithread_load": true, "num_threads": 64}',
], ],
env={"SGLANG_ENABLE_CP_V2": "0"},
) )
@classmethod @classmethod
def tearDownClass(cls): def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid) kill_process_tree(cls.process.pid)
def test_gsm8k(self): def test_gsm8k(self):
+200
View File
@@ -0,0 +1,200 @@
import unittest
from types import SimpleNamespace
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.run_eval import run_eval
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
kill_process_tree,
popen_launch_server,
)
register_cuda_ci(est_time=500, stage="extra-b", runner_config="4-gpu-h100")
GQA_MODEL_PATH = "Qwen/Qwen3-30B-A3B-FP8"
GSM8K_BASELINE_ACCURACY = 0.93
class TestGQACP2TP2EP2(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = GQA_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--moe-dp-size",
"2",
"--ep-size",
"2",
"--attn-cp-size",
"2",
"--enable-prefill-cp",
"--cp-strategy",
"zigzag",
"--cuda-graph-max-bs",
"32",
"--max-running-requests",
"32",
"--trust-remote-code",
"--disable-piecewise-cuda-graph",
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
],
env={"SGLANG_ENABLE_CP_V2": "1"},
)
@classmethod
def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
args = SimpleNamespace(
model=self.model,
eval_name="gsm8k",
num_shots=5,
num_examples=200,
max_tokens=16000,
num_threads=128,
repeat=1,
temperature=0.6,
top_p=0.95,
top_k=20,
base_url=self.base_url,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
metrics = run_eval(args)
print(f"{metrics=}")
self.assertGreaterEqual(metrics["score"], GSM8K_BASELINE_ACCURACY)
class TestGQACPTP2CP2EP4(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = GQA_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--moe-dp-size",
"1",
"--ep-size",
"4",
"--attn-cp-size",
"2",
"--enable-prefill-cp",
"--cp-strategy",
"zigzag",
"--cuda-graph-max-bs",
"32",
"--max-running-requests",
"32",
"--trust-remote-code",
"--disable-piecewise-cuda-graph",
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
],
env={"SGLANG_ENABLE_CP_V2": "1"},
)
@classmethod
def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
args = SimpleNamespace(
model=self.model,
eval_name="gsm8k",
num_shots=5,
num_examples=200,
max_tokens=16000,
num_threads=128,
repeat=1,
temperature=0.6,
top_p=0.95,
top_k=20,
base_url=self.base_url,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
metrics = run_eval(args)
print(f"{metrics=}")
self.assertGreaterEqual(metrics["score"], GSM8K_BASELINE_ACCURACY)
class TestGQACPCP4EP4(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = GQA_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--ep",
"4",
"--attn-cp-size",
"4",
"--enable-prefill-cp",
"--cp-strategy",
"zigzag",
"--moe-a2a-backend",
"deepep",
"--attention-backend",
"fa3",
"--cuda-graph-max-bs",
"32",
"--max-running-requests",
"32",
"--trust-remote-code",
"--disable-piecewise-cuda-graph",
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
],
)
@classmethod
def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
args = SimpleNamespace(
model=self.model,
eval_name="gsm8k",
num_shots=5,
num_examples=200,
max_tokens=16000,
num_threads=128,
repeat=1,
temperature=0.6,
top_p=0.95,
top_k=20,
base_url=self.base_url,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
metrics = run_eval(args)
print(f"{metrics=}")
self.assertGreaterEqual(metrics["score"], GSM8K_BASELINE_ACCURACY)
if __name__ == "__main__":
unittest.main()
@@ -159,6 +159,7 @@ class TestContextParallelServerArgs(CustomTestCase):
enable_dsa_prefill_context_parallel=False, enable_dsa_prefill_context_parallel=False,
enable_prefill_cp=False, enable_prefill_cp=False,
cp_strategy=None, cp_strategy=None,
model_path="instance://127.0.0.1:8000/dummy",
dsa_prefill_cp_mode="round-robin-split", dsa_prefill_cp_mode="round-robin-split",
prefill_cp_mode="in-seq-split", prefill_cp_mode="in-seq-split",
attn_cp_size=1, attn_cp_size=1,