[3/N][CP] Implement zigzag CP strategy (#28421)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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."
|
||||
|
||||
@@ -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()
|
||||
|
||||
+13
-9
@@ -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(
|
||||
@@ -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_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,
|
||||
|
||||
Reference in New Issue
Block a user