diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index d5f724b9b..0a50f0d60 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -13,6 +13,8 @@ from sglang.srt.layers.attention.triton_ops.metadata import ( prepare_swa_spec_page_table_triton, ) 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.utils.cp_utils import ( cp_allgather_and_save_kv_cache, @@ -811,18 +813,26 @@ class FlashAttentionBackend(AttentionBackend): elif is_cp_mode: # Dense-MHA CP: k, v are still rank-local; backend # all-gathers and writes to the per-rank pool. - cp_allgather_and_save_kv_cache( - forward_batch, - layer, - k, - v, - self.attn_cp_size, - swa_loc=( - self.forward_metadata.swa_out_cache_loc - if self.use_sliding_window_kv_pool - else None - ), + 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( + forward_batch, + layer, + k, + v, + self.attn_cp_size, + swa_loc=swa_loc, + ) else: self.token_to_kv_pool.set_kv_buffer( layer, @@ -967,12 +977,24 @@ class FlashAttentionBackend(AttentionBackend): **kwargs, ) - result = cp_attn_forward_extend( - forward_batch, - q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim), - self.device, - _fa_cp_attn, - ) + 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( + forward_batch, + q_cp, + self.device, + _fa_cp_attn, + ) elif self.fa_skip_kv_cache: # Embedding mode: skip KV cache read and use raw K/V tensors # directly via flash_attn_varlen_func. The KV cache write is diff --git a/python/sglang/srt/layers/cp/base.py b/python/sglang/srt/layers/cp/base.py index b63f46933..4be8498a0 100644 --- a/python/sglang/srt/layers/cp/base.py +++ b/python/sglang/srt/layers/cp/base.py @@ -185,6 +185,7 @@ class ContextParallelStrategy(ABC): layer: Any, k: Any, v: Any, + swa_loc: Optional[Any] = None, ) -> None: """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. Subprocesses re-import this module with ``_STRATEGY = None`` and never @@ -257,20 +258,15 @@ def _get_cp_strategy() -> Optional[ContextParallelStrategy]: 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: - strategy = _get_cp_strategy() + strategy = get_cp_strategy() if strategy is None: return ContextParallelStrategyKind.NONE return strategy.kind def is_cp_enabled() -> bool: - return _get_cp_strategy() is not None + return get_cp_strategy() is not None def is_zigzag() -> bool: diff --git a/python/sglang/srt/layers/cp/interleave.py b/python/sglang/srt/layers/cp/interleave.py index 7ce9db896..7db762ab9 100644 --- a/python/sglang/srt/layers/cp/interleave.py +++ b/python/sglang/srt/layers/cp/interleave.py @@ -99,7 +99,9 @@ class InterleaveCPStrategy(ContextParallelStrategy): "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( "Interleave KV materialization will land in a follow-up PR" ) diff --git a/python/sglang/srt/layers/cp/utils.py b/python/sglang/srt/layers/cp/utils.py index a8a8582db..21f980987 100644 --- a/python/sglang/srt/layers/cp/utils.py +++ b/python/sglang/srt/layers/cp/utils.py @@ -12,13 +12,16 @@ # 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 ( BaseContextParallelMetadata, ContextParallelStrategy, ContextParallelStrategyKind, CPAttentionBackendKind, + get_cp_strategy, ) from sglang.srt.layers.cp.interleave import ( InterleaveContextParallelMetadata, @@ -30,6 +33,96 @@ from sglang.srt.layers.cp.zigzag import ( 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__ = [ "BaseContextParallelMetadata", "CPAttentionBackendKind", @@ -40,4 +133,11 @@ __all__ = [ "InterleaveContextParallelMetadata", "ZigzagCPStrategy", "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", ] diff --git a/python/sglang/srt/layers/cp/zigzag.py b/python/sglang/srt/layers/cp/zigzag.py index c25a75898..3e5af00ba 100644 --- a/python/sglang/srt/layers/cp/zigzag.py +++ b/python/sglang/srt/layers/cp/zigzag.py @@ -30,15 +30,29 @@ After all-gather, the blocks are reranged back to their original order: from __future__ import annotations +from contextlib import nullcontext from dataclasses import dataclass +from itertools import accumulate 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 ( BaseContextParallelMetadata, ContextParallelStrategy, ContextParallelStrategyKind, 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 @@ -85,7 +99,13 @@ class ZigzagCPStrategy(ContextParallelStrategy): if self.cp_size <= 1 or num_tokens < self.cp_size * 2: return False 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( self, @@ -93,32 +113,193 @@ class ZigzagCPStrategy(ContextParallelStrategy): seqs_len: Optional[List[int]], extend_seqs_len: Optional[List[int]] = None, ) -> 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( - total_seq_lens=sum(extend_seqs_len or seqs_len or [num_tokens]), - bs=len(extend_seqs_len or seqs_len or [num_tokens]), + split_list=split_list, + 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: - raise NotImplementedError( - "Zigzag hidden-state sharding will land in a follow-up PR" + chunks = torch.split(x, forward_batch.attn_cp_metadata.split_list, dim=0) + 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: - raise NotImplementedError( - "Zigzag position-id sharding will land in a follow-up PR" + chunks = torch.split( + 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( self, x: Any, forward_batch, stream: Optional[Any] = None ) -> Any: - raise NotImplementedError( - "Zigzag hidden-state 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 gather_kv_cache( self, x: Any, forward_batch, stream: Optional[Any] = None ) -> 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( self, @@ -128,11 +309,79 @@ class ZigzagCPStrategy(ContextParallelStrategy): attn_fn, attention_backend: CPAttentionBackendKind = CPAttentionBackendKind.FLASH_ATTENTION, ) -> Any: - raise NotImplementedError( - "Zigzag attention dispatch will land in a follow-up PR" + assert ( + 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: - raise NotImplementedError( - "Zigzag KV materialization will land in a follow-up PR" + def _all_gather_reorganized(self, x: torch.Tensor, forward_batch, stream): + meta = forward_batch.attn_cp_metadata + 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, ) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index fc02f46ed..126ccfae5 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -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.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 ( DpPaddingMode, get_attention_tp_group, @@ -3411,6 +3418,8 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.prefill_cuda_graph_runner is not None 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: # TODO: device_timer.wrap is too broad here — it also includes # 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. self.model.prepare_forward_batch(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 = ( self.device_timer.wrap(metadata={"category": "extend"}) @@ -3441,7 +3465,11 @@ class ModelRunner(ModelRunnerKVCacheMixin): else contextlib.nullcontext() ) 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 # size, run eagerly under enable_tc_piecewise_cuda_graph() and # set_tc_piecewise_forward_context() so that (a) Dynamo guards on @@ -3461,14 +3489,47 @@ class ModelRunner(ModelRunnerKVCacheMixin): ): ret = self.model.forward( forward_batch.input_ids, - forward_batch.positions, + forward_positions, forward_batch, **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: ret = self.model.forward( forward_batch.input_ids, - forward_batch.positions, + forward_positions, forward_batch, **kwargs, ) diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index e82c039c5..66da84645 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -42,6 +42,7 @@ from sglang.srt.layers.communicator import ( LayerScatterModes, ScatterMode, ) +from sglang.srt.layers.cp.utils import is_cp_v2_active from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) @@ -889,6 +890,7 @@ class Qwen2MoeModel(nn.Module): if ( 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.attn_cp_metadata is not None ): @@ -944,6 +946,7 @@ class Qwen2MoeModel(nn.Module): if ( self.pp_group.is_last_rank + and not is_cp_v2_active(forward_batch) and is_prefill_context_parallel_enabled() and forward_batch.forward_mode.is_context_parallel_extend() and forward_batch.attn_cp_metadata is not None diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 90ec529cd..a29df9f46 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -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_dispatch import ExpertLocationDispatchInfo 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.linear import ( QKVParallelLinear, @@ -988,7 +989,7 @@ class Qwen3MoeForCausalLM(nn.Module): input_embeds: torch.Tensor = None, pp_proxy_tensors: Optional[PPProxyTensors] = None, ) -> 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): forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( len(input_ids), diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index a43eed04d..b4098a0dc 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -3436,6 +3436,17 @@ class ServerArgs: self.prefill_cp_mode = mode 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: raise ValueError( "--cp-strategy must be set when --enable-prefill-cp is enabled." diff --git a/test/registered/cp/test_cp_strategy_unit.py b/test/registered/cp/test_cp_strategy_unit.py index 8f306b5ea..270ab9feb 100644 --- a/test/registered/cp/test_cp_strategy_unit.py +++ b/test/registered/cp/test_cp_strategy_unit.py @@ -2,6 +2,8 @@ import unittest from types import SimpleNamespace from unittest.mock import patch +import torch + from sglang.srt.layers.cp.base import ( ContextParallelStrategyKind, get_cp_strategy, @@ -11,10 +13,31 @@ from sglang.srt.layers.cp.base import ( is_interleave, 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.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): @@ -70,5 +93,367 @@ class TestCPStrategyUnit(CustomTestCase): 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__": unittest.main() diff --git a/test/registered/cp/test_qwen3_30b.py b/test/registered/cp/test_gqa_prefill_cp_legacy.py similarity index 85% rename from test/registered/cp/test_qwen3_30b.py rename to test/registered/cp/test_gqa_prefill_cp_legacy.py index e21dafb5c..925e54184 100644 --- a/test/registered/cp/test_qwen3_30b.py +++ b/test/registered/cp/test_gqa_prefill_cp_legacy.py @@ -11,17 +11,17 @@ from sglang.test.test_utils import ( 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 def setUpClass(cls): - cls.model = QWEN3_30B_MODEL_PATH + cls.model = GQA_MODEL_PATH cls.base_url = DEFAULT_URL_FOR_TEST cls.process = popen_launch_server( cls.model, @@ -46,11 +46,13 @@ class TestQwen330B(CustomTestCase): "--model-loader-extra-config", '{"enable_multithread_load": true, "num_threads": 64}', ], + env={"SGLANG_ENABLE_CP_V2": "0"}, ) @classmethod def tearDownClass(cls): - kill_process_tree(cls.process.pid) + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) def test_gsm8k(self): args = SimpleNamespace( @@ -73,10 +75,10 @@ class TestQwen330B(CustomTestCase): self.assertGreaterEqual(metrics["score"], GSM8K_BASELINE_ACCURACY) -class TestQwen330BCP(CustomTestCase): +class TestGQACPTP2CP2EP4(CustomTestCase): @classmethod def setUpClass(cls): - cls.model = QWEN3_30B_MODEL_PATH + cls.model = GQA_MODEL_PATH cls.base_url = DEFAULT_URL_FOR_TEST cls.process = popen_launch_server( cls.model, @@ -101,11 +103,13 @@ class TestQwen330BCP(CustomTestCase): "--model-loader-extra-config", '{"enable_multithread_load": true, "num_threads": 64}', ], + env={"SGLANG_ENABLE_CP_V2": "0"}, ) @classmethod def tearDownClass(cls): - kill_process_tree(cls.process.pid) + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) def test_gsm8k(self): args = SimpleNamespace( diff --git a/test/registered/cp/test_gqa_preill_cp.py b/test/registered/cp/test_gqa_preill_cp.py new file mode 100644 index 000000000..d19847ab5 --- /dev/null +++ b/test/registered/cp/test_gqa_preill_cp.py @@ -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() diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 92a954a7c..53f5b8908 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -159,6 +159,7 @@ class TestContextParallelServerArgs(CustomTestCase): enable_dsa_prefill_context_parallel=False, enable_prefill_cp=False, cp_strategy=None, + model_path="instance://127.0.0.1:8000/dummy", dsa_prefill_cp_mode="round-robin-split", prefill_cp_mode="in-seq-split", attn_cp_size=1,