Support batch size > 1 when enable CP (#23269)

Co-authored-by: Shunkang <182541032+Shunkangz@users.noreply.github.co>
Co-authored-by: Khoa Pham <khoa.pham@radixark.ai>
Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com>
This commit is contained in:
Shunkangz
2026-05-27 14:11:17 -07:00
committed by GitHub
co-authored by Shunkang Khoa Pham Baizhou Zhang
parent ddf0627254
commit 19663aafcd
13 changed files with 263 additions and 300 deletions
@@ -1455,10 +1455,14 @@ class Indexer(MultiPlatformOp):
forward_batch.attn_cp_metadata is not None forward_batch.attn_cp_metadata is not None
and is_dsa_prefill_cp_in_seq_split() and is_dsa_prefill_cp_in_seq_split()
): ):
kv_len_prev = forward_batch.attn_cp_metadata.kv_len_prev kv_len_prev = forward_batch.attn_cp_metadata.kv_len_prev_list[0]
kv_len_next = forward_batch.attn_cp_metadata.kv_len_next kv_len_next = forward_batch.attn_cp_metadata.kv_len_next_list[0]
actual_seq_q_prev = forward_batch.attn_cp_metadata.actual_seq_q_prev actual_seq_q_prev = (
actual_seq_q_next = forward_batch.attn_cp_metadata.actual_seq_q_next forward_batch.attn_cp_metadata.actual_seq_q_prev_list[0]
)
actual_seq_q_next = (
forward_batch.attn_cp_metadata.actual_seq_q_next_list[0]
)
# TODO support mutil-batch # TODO support mutil-batch
# cp_batch_seq_index_prev = forward_batch.attn_cp_metadata["cp_batch_seq_index_prev"] # cp_batch_seq_index_prev = forward_batch.attn_cp_metadata["cp_batch_seq_index_prev"]
@@ -133,7 +133,9 @@ def cal_padded_tokens(forward_batch: "ForwardBatch"):
sync_group_size = len(global_num_tokens) sync_group_size = len(global_num_tokens)
attn_cp_size = get_attention_cp_size() attn_cp_size = get_attention_cp_size()
for i in range(sync_group_size): for i in range(sync_group_size):
global_num_tokens[i] = ceil_align(global_num_tokens[i], attn_cp_size) # Must match ForwardBatch.prepare_mlp_sync_batch, which pads to
# attn_cp_size * 2 (tokens are split into 2 * CP chunks for load balance).
global_num_tokens[i] = ceil_align(global_num_tokens[i], attn_cp_size * 2)
dp_padding_mode = DpPaddingMode.get_dp_padding_mode( dp_padding_mode = DpPaddingMode.get_dp_padding_mode(
forward_batch.is_extend_in_batch, global_num_tokens forward_batch.is_extend_in_batch, global_num_tokens
) )
+231 -148
View File
@@ -20,24 +20,41 @@ from sglang.srt.server_args import get_global_server_args
@dataclass @dataclass
class ContextParallelMetadata: class ContextParallelMetadata:
# Layout lists have length bs * cp_segment_num (= bs * 2 * cp_size).
split_list: List[int] = None split_list: List[int] = None
max_rank_len: List[int] = None
zigzag_index: List[int] = None zigzag_index: List[int] = None
per_rank_actual_token: List[int] = None
reverse_split_len: List[int] = None
cp_reverse_index: List[int] = None cp_reverse_index: List[int] = None
reverse_split_len: List[int] = None
# metadata for attention # Per-rank-aggregate lists have length cp_size.
kv_len_prev: int = -1 # max_rank_len is a list of cp_size copies of max(per_rank_actual_token),
kv_len_next: int = -1 # kept as a list for torch.split() bucket sizes.
actual_seq_q_prev: int = -1 per_rank_actual_token: List[int] = None
actual_seq_q_next: int = -1 max_rank_len: List[int] = None
kv_len_prev_tensor: torch.Tensor = None
kv_len_next_tensor: torch.Tensor = None
actual_seq_q_prev_tensor: torch.Tensor = None
actual_seq_q_next_tensor: torch.Tensor = None
total_seq_lens: torch.Tensor = None # Per-sequence FlashAttention tensors (shape [bs] or [bs+1]).
kv_len_prev_tensor: torch.Tensor = None # [bs] int32 CUDA
kv_len_next_tensor: torch.Tensor = None # [bs] int32 CUDA
actual_seq_q_prev_tensor: torch.Tensor = None # [bs] int32 CUDA
actual_seq_q_next_tensor: torch.Tensor = None # [bs] int32 CUDA
cu_seqlens_q_prev_tensor: torch.Tensor = None # [bs+1] int32 CUDA
cu_seqlens_q_next_tensor: torch.Tensor = None # [bs+1] int32 CUDA
# Scalars derived from the per-sequence lists above.
total_q_prev_tokens: int = 0
total_q_next_tokens: int = 0
max_seqlen_q_prev: int = 0
max_seqlen_q_next: int = 0
# Per-seq CPU lists (useful for NSA indexer and diagnostics).
kv_len_prev_list: List[int] = None
kv_len_next_list: List[int] = None
actual_seq_q_prev_list: List[int] = None
actual_seq_q_next_list: List[int] = None
# Aggregate sum of extend_seq_lens across the batch.
total_seq_lens: int = 0
bs: int = 1
def is_prefill_context_parallel_enabled(): def is_prefill_context_parallel_enabled():
@@ -67,25 +84,45 @@ def mla_use_prefill_cp(forward_batch, mla_enable_prefill_cp=None):
def can_cp_split(seq_len: int, cp_size: int, forward_batch): def can_cp_split(seq_len: int, cp_size: int, forward_batch):
# Base conditions: CP must be enabled, size > 1, and this must be a
# CP-extend (prefill) step. The seq_len // (cp_size * 2) check ensures
# the load-balancing split into 2 * cp_size blocks is non-degenerate.
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
# TODO current just support prefill batch=1 and len(input_ids) > self.cp_size * 2
# Note: (self.cp_size * 2) To achieve load balancing for seq computation,
# the seq data needs to be divided and recombined at twice the size of cp_size.
cur_cp_seq_len = seq_len // (cp_size * 2) cur_cp_seq_len = seq_len // (cp_size * 2)
return ( if not (
cur_cp_seq_len != 0 cur_cp_seq_len != 0
and cp_size > 1 and cp_size > 1
# prepare_context_parallel_metadata hard-codes bs_per_cp_group = 1; # prepare_context_parallel_metadata hard-codes bs_per_cp_group = 1;
# guard explicitly to avoid silent mis-partitioning under continuous batching. # guard explicitly to avoid silent mis-partitioning under continuous batching.
# TODO: remove this guard once we support multi-batch-cp-split
and forward_batch.batch_size == 1
and forward_batch.forward_mode.is_context_parallel_extend() and forward_batch.forward_mode.is_context_parallel_extend()
# is_context_parallel_extend() returns True for MIXED (prefill+decode # is_context_parallel_extend() returns True for MIXED (prefill+decode
# in one step), but the zigzag split only makes sense on pure extend. # in one step), but the zigzag split only makes sense on pure extend.
and forward_batch.forward_mode != ForwardMode.MIXED and forward_batch.forward_mode != ForwardMode.MIXED
and is_prefill_context_parallel_enabled() and is_prefill_context_parallel_enabled()
) ):
return False
# Per-sequence guards for bs > 1. Every sequence must be long enough for
# the 2*cp_size-way split. A sub-threshold request reaching this point
# means the scheduler failed to filter it out and a silent non-CP
# fallback would have masked the bug -- raise instead. Per-sequence
# radix-cache prefix is supported: prefix is baked into kv_len_prev/next
# via prefix_offsets[s] inside prepare_context_parallel_metadata.
extend_lens = getattr(forward_batch, "extend_seq_lens_cpu", None)
if extend_lens is None:
return True
cp_min = cp_size * 2
for L in extend_lens:
if L < cp_min:
# A sub-threshold request cannot be zigzag-split into 2*cp_size
# blocks; fall back to a normal (non-CP) prefill for this batch
# instead of failing. Happens e.g. when a radix-cache prefix hit
# leaves only a few unique extend tokens.
return False
return True
def cp_split_and_rebuild_data(forward_batch, input_: torch.Tensor): def cp_split_and_rebuild_data(forward_batch, input_: torch.Tensor):
@@ -134,9 +171,7 @@ def cp_split_and_rebuild_position(forward_batch, positions: torch.Tensor):
return positions return positions
def cp_all_gather_reorganized_into_tensor( def cp_all_gather_reorganized_into_tensor(input_tensor, cp_size, forward_batch, stream):
input_tensor, total_len, cp_size, forward_batch, stream
):
""" """
Allgather communication for context_parallel(kv_cache, index_k, hidden_states). Allgather communication for context_parallel(kv_cache, index_k, hidden_states).
This implementation mainly consists of three parts: This implementation mainly consists of three parts:
@@ -144,10 +179,7 @@ def cp_all_gather_reorganized_into_tensor(
Step 2, allgather communication(async). Step 2, allgather communication(async).
Step 3, removing the padding and reassembling the data according to the actual tokens. Step 3, removing the padding and reassembling the data according to the actual tokens.
""" """
# The input tensor should already be padded to the same length for allgather communication. max_len = forward_batch.attn_cp_metadata.max_rank_len[0]
# No need to pad again.
# step1
max_len = (total_len + cp_size - 1) // cp_size
pad_size = max_len - input_tensor.shape[0] pad_size = max_len - input_tensor.shape[0]
if pad_size > 0: if pad_size > 0:
input_tensor = F.pad( input_tensor = F.pad(
@@ -186,13 +218,13 @@ def cp_all_gather_reorganized_into_tensor(
def cp_all_gather_reorganized_into_tensor_kv_cache( def cp_all_gather_reorganized_into_tensor_kv_cache(
input_tensor, total_len, cp_size, forward_batch, stream input_tensor, cp_size, forward_batch, stream
): ):
""" """
Allgather communication for context_parallel KV cache. Allgather communication for context_parallel KV cache.
Handles multi-dimensional tensors (e.g., [seq_len, num_heads, head_dim]). Handles multi-dimensional tensors (e.g., [seq_len, num_heads, head_dim]).
""" """
max_len = (total_len + cp_size - 1) // cp_size max_len = forward_batch.attn_cp_metadata.max_rank_len[0]
pad_size = max_len - input_tensor.shape[0] pad_size = max_len - input_tensor.shape[0]
if pad_size > 0: if pad_size > 0:
# Pad the first dimension (seq_len). F.pad expects padding in reverse dimension order. # Pad the first dimension (seq_len). F.pad expects padding in reverse dimension order.
@@ -288,7 +320,6 @@ def cp_all_gather_rerange_output(input_tensor, cp_size, forward_batch, stream):
bs_seq_len, hidden_size = input_tensor.shape bs_seq_len, hidden_size = input_tensor.shape
output_tensor = cp_all_gather_reorganized_into_tensor( output_tensor = cp_all_gather_reorganized_into_tensor(
input_tensor, input_tensor,
forward_batch.attn_cp_metadata.total_seq_lens,
cp_size, cp_size,
forward_batch, forward_batch,
stream, stream,
@@ -326,7 +357,6 @@ def cp_all_gather_rerange_kv_cache(input_tensor, cp_size, forward_batch, stream)
""" """
output_tensor = cp_all_gather_reorganized_into_tensor_kv_cache( output_tensor = cp_all_gather_reorganized_into_tensor_kv_cache(
input_tensor, input_tensor,
forward_batch.attn_cp_metadata.total_seq_lens,
cp_size, cp_size,
forward_batch, forward_batch,
stream, stream,
@@ -386,6 +416,11 @@ def cp_attn_forward_extend(
backend-specific attention function twice with appropriate per-half backend-specific attention function twice with appropriate per-half
metadata, and concatenate the results. metadata, and concatenate the results.
For bs > 1, q is laid out as [all_prev_tokens_across_seqs,
all_next_tokens_across_seqs]; the split point is total_q_prev_tokens.
cu_seqlens_q_prev/next tensors have shape [bs+1] and carry the
per-sequence boundaries through FlashAttention's variable-length API.
attn_fn signature: attn_fn signature:
attn_fn(q, cu_seqlens_q, cache_seqlens, max_seqlen_q) -> result attn_fn(q, cu_seqlens_q, cache_seqlens, max_seqlen_q) -> result
where only these four CP-varying parameters differ between halves. where only these four CP-varying parameters differ between halves.
@@ -393,20 +428,20 @@ def cp_attn_forward_extend(
""" """
cp_meta = forward_batch.attn_cp_metadata cp_meta = forward_batch.attn_cp_metadata
q_prev, q_next = torch.chunk(q, 2, dim=0) q_prev = q[: cp_meta.total_q_prev_tokens]
q_next = q[cp_meta.total_q_prev_tokens :]
cu_seqlens_q_prev = torch.tensor(
[0, cp_meta.actual_seq_q_prev], device=device, dtype=torch.int32
)
result_prev = attn_fn( result_prev = attn_fn(
q_prev, cu_seqlens_q_prev, cp_meta.kv_len_prev_tensor, cp_meta.actual_seq_q_prev q_prev,
) cp_meta.cu_seqlens_q_prev_tensor,
cp_meta.kv_len_prev_tensor,
cu_seqlens_q_next = torch.tensor( cp_meta.max_seqlen_q_prev,
[0, cp_meta.actual_seq_q_next], device=device, dtype=torch.int32
) )
result_next = attn_fn( result_next = attn_fn(
q_next, cu_seqlens_q_next, cp_meta.kv_len_next_tensor, cp_meta.actual_seq_q_next q_next,
cp_meta.cu_seqlens_q_next_tensor,
cp_meta.kv_len_next_tensor,
cp_meta.max_seqlen_q_next,
) )
return torch.concat([result_prev, result_next], dim=0) return torch.concat([result_prev, result_next], dim=0)
@@ -417,7 +452,8 @@ def prepare_context_parallel_metadata(
cp_rank, cp_rank,
cp_size, cp_size,
seqs_len, seqs_len,
extend_lens, extend_seqs_len=None,
device="cuda",
): ):
from sglang.srt.layers.attention.dsa.utils import ( from sglang.srt.layers.attention.dsa.utils import (
is_dsa_prefill_cp_round_robin_split, is_dsa_prefill_cp_round_robin_split,
@@ -427,144 +463,191 @@ def prepare_context_parallel_metadata(
return ContextParallelMetadata() return ContextParallelMetadata()
"""prepare_input_dp_with_cp_dsa-zigzag index """prepare_input_dp_with_cp_dsa-zigzag index
Example (DP_ATTENT_TP == CP_SIZE == 4): Example (DP_ATTENT_TP == CP_SIZE == 4, single sequence):
Description: block0 | block1 | block2 | block3 | block4 | block5 | block6 | block7
1. Start with a full-length request. rank 0: block0, block7
2. Split the request into multiple blocks (block0 to block7). rank 1: block1, block6
3. Rearrange these blocks to balance computational rank 2: block2, block5
load across different DP ranks. rank 3: block3, block4
4. Assign the rearranged blocks to different DP attention For bs > 1, each sequence is split into cp_segment_num = 2 * cp_size
time points (dp_atten_tp0 to dp_atten_tp3). blocks independently; per-rank layout becomes:
+---------------------------------+ [s0.block_r, s1.block_r, ..., s_{bs-1}.block_r,
| cp_split_tokens | s0.block_{2*cp_size-1-r}, ..., s_{bs-1}.block_{2*cp_size-1-r}]
+---------------------------------+ i.e. all prev blocks first, then all next blocks -- so torch.split at
| | total_q_prev_tokens cleanly separates them.
| request_with_full_length |
| | split (cp_size * 2) |
| +-------------------------+ |
| | block0 | block1 | block2 | block3 | block4 | block5 | block6 | block7 |
| +-------------------------+ |
| | rerange |
| +---------------------------------+
| | block0 | block7 | block1 | block6 | block2 | block5 | block3 | block4 |
| +---------------------------------+
| |
| +-------------------------+
| | dp_atten_tp0: block0, block7 |
| | dp_atten_tp1: block1, block6 |
| | dp_atten_tp2: block2, block5 |
| | dp_atten_tp3: block3, block4 |
| +-------------------------+
Why zigzag rearrange?
- Attention calculations must follow causal attention principles.
- Simply slicing by rank order can lead to computational load imbalance:
* First rank may focus on fewer historical key-value tokens (less computation)
* Last rank may focus on more tokens (more computation)
- To mitigate uneven load, the input hidden states needs to be sliced by cp_size*2 and rearranged.
""" """
# just support batch = 1 assert extend_seqs_len is not None
# kv_len: the number of tokens *computed in this extend pass* (i.e. the extend_seqs_len = [int(x) for x in extend_seqs_len]
# "new" tokens). When radix/prefix cache hits, the effective KV length
# visible to attention is: prefix_len + kv_len. CP attention must use the
# full visible KV length, otherwise queries won't attend to cached prefix.
kv_len = torch.tensor(kv_len)
bs_per_cp_group = 1
kv_len_origin = kv_len
# Derive prefix offset from unpadded CPU tensors. Both `seqs_len` and `extend_lens` are unpadded by the caller # Update the extend_seqs_len to the padded length.
# Using the padded `kv_len` here would undercount `prefix_len` by the padding amount and shift the FA causal horizon. pad_len = int(kv_len) - sum(extend_seqs_len)
assert ( if pad_len > 0:
len(seqs_len) == 1 and len(extend_lens) == 1 extend_seqs_len[-1] += pad_len
), "Prefill Context Parallel only supports batch_size == 1 for now" if seqs_len is not None and len(seqs_len) == len(extend_seqs_len):
prefix_len = max(0, int(seqs_len[0]) - int(extend_lens[0])) seqs_len = list(seqs_len)
# get zigzag index seqs_len[-1] += pad_len
bs = len(extend_seqs_len)
cp_segment_num = cp_size * 2 cp_segment_num = cp_size * 2
seq_per_batch = kv_len // cp_segment_num # seq_len for each batch and segment
split_list = seq_per_batch.repeat_interleave(cp_segment_num).int().tolist()
remainder = kv_len % (cp_segment_num)
if remainder > 0:
split_list[:remainder] = [x + 1 for x in split_list[:remainder]]
seq_max_rank_len = (kv_len + cp_size - 1) // cp_size # Prefix offset (radix cache hit length) per sequence. For non-NSA
max_rank_len = seq_max_rank_len.repeat_interleave(cp_size).int().tolist() # (FlashAttention) the prefix is baked into kv_len_prev/next via
# prefix_offsets[s] below, so cache_seqlens correctly covers the cached
# prefix. NSA leaves bare cumulatives so its indexer can re-add the
# offset itself.
if seqs_len is not None and len(seqs_len) == bs:
prefix_offsets = [
max(int(seqs_len[s]) - extend_seqs_len[s], 0) for s in range(bs)
]
else:
prefix_offsets = [0] * bs
# Per-sequence block sizes: first (L % cp_segment_num) blocks get +1.
per_seq_block_sizes: List[List[int]] = []
split_list: List[int] = []
for s in range(bs):
L = extend_seqs_len[s]
base = L // cp_segment_num
rem = L % cp_segment_num
blk = [base + 1 if i < rem else base for i in range(cp_segment_num)]
per_seq_block_sizes.append(blk)
split_list.extend(blk)
# Per-rank aggregate: this rank owns block r and block (2*cp_size-1-r)
# of every sequence.
per_rank_actual_token = [0] * cp_size
for r in range(cp_size):
total = 0
for s in range(bs):
total += (
per_seq_block_sizes[s][r]
+ per_seq_block_sizes[s][cp_segment_num - 1 - r]
)
per_rank_actual_token[r] = total
max_single_rank = max(per_rank_actual_token) if per_rank_actual_token else 0
# Kept as cp_size copies so downstream torch.split(x, max_rank_len) still
# works directly. All entries intentionally identical.
max_rank_len = [max_single_rank] * cp_size
# Zigzag index selecting which of split_list's bs * cp_segment_num pieces
# this rank owns, in the order [all_prevs, all_nexts].
zigzag_index = list( zigzag_index = list(
range(cp_rank, cp_rank + bs_per_cp_group * cp_segment_num, cp_segment_num) range(cp_rank, cp_rank + bs * cp_segment_num, cp_segment_num)
) + list( ) + list(
range( range(
cp_segment_num - cp_rank - 1, cp_segment_num - cp_rank - 1,
bs_per_cp_group * cp_segment_num, bs * cp_segment_num,
cp_segment_num, cp_segment_num,
) )
) )
per_rank_actual_token = list( # Reverse index: given the post-allgather concatenation
split_list[i] + split_list[cp_size * 2 - i - 1] for i in range(cp_size) # [rank0_prevs_all_seqs, rank0_nexts_all_seqs,
) # rank1_prevs_all_seqs, rank1_nexts_all_seqs, ...]
reverse_split_len = [ # produce a permutation that restores [s0_b0..s0_bN, s1_b0..s1_bN, ...].
element cp_reverse_index: List[int] = []
for i in range(cp_size) for batch_id in range(bs):
for element in (split_list[i], split_list[cp_size * 2 - i - 1])
]
# get zigzag reverse index
cp_reverse_index = []
for batch_id in range(bs_per_cp_group):
cp_reverse_index.extend( cp_reverse_index.extend(
list(range(batch_id, cp_segment_num * bs_per_cp_group, 2 * bs_per_cp_group)) list(range(batch_id, cp_segment_num * bs, 2 * bs))
+ list( + list(
range( range(
(cp_segment_num - 1) * bs_per_cp_group + batch_id, (cp_segment_num - 1) * bs + batch_id,
0, 0,
-2 * bs_per_cp_group, -2 * bs,
) )
) )
) )
prefix_sum_list = list(accumulate(split_list))
# TODO Support multi-batch-cp-split, multi-batch-cp support has accuracy issues # Split sizes matching the post-allgather concatenation order above.
# Prefix offset is critical when radix cache hits (prefix_len > 0). reverse_split_len: List[int] = []
# For non-DSA CP (e.g. qwen3-moe), consumers use these values directly as for r in range(cp_size):
# FlashAttention cache_seqlens, so the prefix must be baked in here. for s in range(bs):
# For DSA CP, `_get_topk_ragged_with_cp` re-adds the cached-prefix offset reverse_split_len.append(per_seq_block_sizes[s][r])
# from (seq_lens_cpu - extend_seq_lens_cpu); baking prefix_len in here for s in range(bs):
# would silently drop it whenever the scheduler packs multiple requests reverse_split_len.append(per_seq_block_sizes[s][cp_segment_num - 1 - r])
# into a single CP extend (len(seqs_len) != 1 -> prefix_len falls back
# to 0), corrupting the indexer's ke_offset on prefix-cache hits. # Per-sequence cumulatives used for FA cache_seqlens.
# kv_len_prev[s] = sum of seq s's blocks [0..cp_rank] (inclusive).
# kv_len_next[s] = sum of seq s's blocks [0..cp_segment_num-cp_rank-1] (inclusive).
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
if is_dsa_enable_prefill_cp(): nsa_mode = is_dsa_enable_prefill_cp()
kv_len_prev = prefix_sum_list[cp_rank] kv_len_prev_list: List[int] = []
kv_len_next = prefix_sum_list[cp_size * 2 - cp_rank - 1] kv_len_next_list: List[int] = []
else: actual_seq_q_prev_list: List[int] = []
kv_len_prev = prefix_len + prefix_sum_list[cp_rank] actual_seq_q_next_list: List[int] = []
kv_len_next = prefix_len + prefix_sum_list[cp_size * 2 - cp_rank - 1] for s in range(bs):
actual_seq_q_prev = split_list[cp_rank] blk = per_seq_block_sizes[s]
actual_seq_q_next = split_list[cp_size * 2 - cp_rank - 1] cum_prev = sum(blk[: cp_rank + 1])
# Flash Attention expects cache_seqlens to have shape (batch_size,), not scalar cum_next = sum(blk[: cp_segment_num - cp_rank])
kv_len_prev_tensor = torch.tensor([kv_len_prev], device="cuda", dtype=torch.int32) # NSA indexer re-adds prefix offset itself; leave bare cumulative.
kv_len_next_tensor = torch.tensor([kv_len_next], device="cuda", dtype=torch.int32) # For non-NSA (FlashAttention), bake prefix into cache_seqlens.
if nsa_mode:
kv_len_prev_list.append(cum_prev)
kv_len_next_list.append(cum_next)
else:
kv_len_prev_list.append(prefix_offsets[s] + cum_prev)
kv_len_next_list.append(prefix_offsets[s] + cum_next)
actual_seq_q_prev_list.append(blk[cp_rank])
actual_seq_q_next_list.append(blk[cp_segment_num - cp_rank - 1])
# FlashAttention CUDA tensors (device parameterized for unit tests).
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_tensor = torch.tensor(
[actual_seq_q_prev], device="cuda", dtype=torch.int32 actual_seq_q_prev_list, device=device, dtype=torch.int32
) )
actual_seq_q_next_tensor = torch.tensor( actual_seq_q_next_tensor = torch.tensor(
[actual_seq_q_next], device="cuda", dtype=torch.int32 actual_seq_q_next_list, device=device, dtype=torch.int32
) )
cu_prev = [0] + list(accumulate(actual_seq_q_prev_list))
cu_next = [0] + list(accumulate(actual_seq_q_next_list))
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)
attn_cp_metadata = ContextParallelMetadata( 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
total_seq_lens = sum(extend_seqs_len)
# Cheap invariants: metadata must be a valid permutation spec.
# - split_list has bs * cp_segment_num pieces (all blocks, all seqs).
# - zigzag_index has 2 * bs entries (this rank's prev + next per seq).
# - cp_reverse_index has bs * cp_segment_num entries (reorders the
# full allgathered stream back to per-seq-original order).
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 ContextParallelMetadata(
split_list=split_list, split_list=split_list,
max_rank_len=max_rank_len,
zigzag_index=zigzag_index, zigzag_index=zigzag_index,
per_rank_actual_token=per_rank_actual_token,
reverse_split_len=reverse_split_len,
cp_reverse_index=cp_reverse_index, cp_reverse_index=cp_reverse_index,
kv_len_prev=kv_len_prev, reverse_split_len=reverse_split_len,
kv_len_next=kv_len_next, per_rank_actual_token=per_rank_actual_token,
actual_seq_q_prev=actual_seq_q_prev, max_rank_len=max_rank_len,
actual_seq_q_next=actual_seq_q_next,
kv_len_prev_tensor=kv_len_prev_tensor, kv_len_prev_tensor=kv_len_prev_tensor,
kv_len_next_tensor=kv_len_next_tensor, kv_len_next_tensor=kv_len_next_tensor,
actual_seq_q_prev_tensor=actual_seq_q_prev_tensor, actual_seq_q_prev_tensor=actual_seq_q_prev_tensor,
actual_seq_q_next_tensor=actual_seq_q_next_tensor, actual_seq_q_next_tensor=actual_seq_q_next_tensor,
total_seq_lens=kv_len_origin, cu_seqlens_q_prev_tensor=cu_seqlens_q_prev_tensor,
cu_seqlens_q_next_tensor=cu_seqlens_q_next_tensor,
total_q_prev_tokens=total_q_prev_tokens,
total_q_next_tokens=total_q_next_tokens,
max_seqlen_q_prev=max_seqlen_q_prev,
max_seqlen_q_next=max_seqlen_q_next,
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,
) )
return attn_cp_metadata
@@ -826,9 +826,7 @@ class PrefillAdder:
# TODO support cp with multiple requests # TODO support cp with multiple requests
# Enabling context parallelism currently presents precision issues; # Enabling context parallelism currently presents precision issues;
# therefore, the prefill-batch setting is temporarily set to 1. # therefore, the prefill-batch setting is temporarily set to 1.
if ( if (self.dsa_prefill_cp_in_seq_split) and len(self.can_run_list) >= 1:
self.dsa_prefill_cp_in_seq_split or self.prefill_context_parallel_enabled
) and len(self.can_run_list) >= 1:
return AddReqResult.OTHER return AddReqResult.OTHER
if (x := self.prefill_max_requests) is not None and len(self.can_run_list) >= x: if (x := self.prefill_max_requests) is not None and len(self.can_run_list) >= x:
@@ -889,10 +889,11 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# there is no reduce-scatter in LM logprob, so we do not need to adjust the padded length for logprob # there is no reduce-scatter in LM logprob, so we do not need to adjust the padded length for logprob
global_num_tokens[i] = ceil_align(global_num_tokens[i], attn_tp_size) global_num_tokens[i] = ceil_align(global_num_tokens[i], attn_tp_size)
# make sure that each rank has the same number of tokens to do collective communication. # make sure that each rank has the same number of tokens to do collective communication and
# we can divide the tokens into 2 * CP chunks for load balance.
attn_cp_size = get_attention_cp_size() attn_cp_size = get_attention_cp_size()
for i in range(sync_group_size): for i in range(sync_group_size):
global_num_tokens[i] = ceil_align(global_num_tokens[i], attn_cp_size) global_num_tokens[i] = ceil_align(global_num_tokens[i], attn_cp_size * 2)
dp_padding_mode = DpPaddingMode.get_dp_padding_mode( dp_padding_mode = DpPaddingMode.get_dp_padding_mode(
self.is_extend_in_batch, global_num_tokens self.is_extend_in_batch, global_num_tokens
+2 -2
View File
@@ -311,7 +311,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
self.cp_rank, self.cp_rank,
self.cp_size, self.cp_size,
forward_batch.seq_lens_cpu.tolist(), forward_batch.seq_lens_cpu.tolist(),
extend_lens=forward_batch.extend_seq_lens_cpu, extend_seqs_len=forward_batch.extend_seq_lens_cpu,
) )
elif self.mla_enable_prefill_cp: elif self.mla_enable_prefill_cp:
if can_cp_split(len(input_ids), self.cp_size, forward_batch): if can_cp_split(len(input_ids), self.cp_size, forward_batch):
@@ -320,7 +320,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
self.cp_rank, self.cp_rank,
self.cp_size, self.cp_size,
forward_batch.seq_lens_cpu.tolist(), forward_batch.seq_lens_cpu.tolist(),
extend_lens=forward_batch.extend_seq_lens_cpu, extend_seqs_len=forward_batch.extend_seq_lens_cpu,
) )
hidden_states = self.model(input_ids, positions, forward_batch) hidden_states = self.model(input_ids, positions, forward_batch)
return self.logits_processor( return self.logits_processor(
+2 -2
View File
@@ -2609,7 +2609,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
self.cp_rank, self.cp_rank,
self.cp_size, self.cp_size,
forward_batch.seq_lens_cpu.tolist(), forward_batch.seq_lens_cpu.tolist(),
extend_lens=forward_batch.extend_seq_lens_cpu, extend_seqs_len=forward_batch.extend_seq_lens_cpu,
) )
elif self.mla_enable_prefill_cp: elif self.mla_enable_prefill_cp:
if can_cp_split(len_input_ids, self.cp_size, forward_batch): if can_cp_split(len_input_ids, self.cp_size, forward_batch):
@@ -2618,7 +2618,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
self.cp_rank, self.cp_rank,
self.cp_size, self.cp_size,
forward_batch.seq_lens_cpu.tolist(), forward_batch.seq_lens_cpu.tolist(),
extend_lens=forward_batch.extend_seq_lens_cpu, extend_seqs_len=forward_batch.extend_seq_lens_cpu,
) )
with get_attn_tp_context().maybe_input_scattered(forward_batch): with get_attn_tp_context().maybe_input_scattered(forward_batch):
+1 -1
View File
@@ -1577,7 +1577,7 @@ class DeepseekV4ForCausalLM(nn.Module):
self.cp_rank, self.cp_rank,
self.cp_size, self.cp_size,
forward_batch.seq_lens_cpu.tolist(), forward_batch.seq_lens_cpu.tolist(),
extend_lens=forward_batch.extend_seq_lens_cpu, extend_seqs_len=forward_batch.extend_seq_lens_cpu,
) )
if is_dsa_prefill_cp_round_robin_split(): if is_dsa_prefill_cp_round_robin_split():
attn_backend = get_attn_backend() attn_backend = get_attn_backend()
@@ -249,7 +249,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM):
self.cp_rank, self.cp_rank,
self.cp_size, self.cp_size,
forward_batch.seq_lens_cpu.tolist(), forward_batch.seq_lens_cpu.tolist(),
extend_lens=forward_batch.extend_seq_lens_cpu, extend_seqs_len=forward_batch.extend_seq_lens_cpu,
) )
if is_dsa_prefill_cp_round_robin_split(): if is_dsa_prefill_cp_round_robin_split():
attn_backend = get_attn_backend() attn_backend = get_attn_backend()
+1 -1
View File
@@ -1005,7 +1005,7 @@ class Qwen3MoeForCausalLM(nn.Module):
self.attn_cp_rank, self.attn_cp_rank,
self.attn_cp_size, self.attn_cp_size,
forward_batch.seq_lens_cpu.tolist(), forward_batch.seq_lens_cpu.tolist(),
extend_lens=forward_batch.extend_seq_lens_cpu, extend_seqs_len=forward_batch.extend_seq_lens_cpu,
) )
hidden_states = self.model( hidden_states = self.model(
-132
View File
@@ -1,132 +0,0 @@
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=261, stage="extra-b", runner_config="4-gpu-h100")
QWEN3_30B_MODEL_PATH = "Qwen/Qwen3-30B-A3B-FP8"
GSM8K_BASELINE_ACCURACY = 0.85
class TestQwen330B(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = QWEN3_30B_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-context-parallel",
"--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):
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 TestQwen330BCP(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = QWEN3_30B_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-context-parallel",
"--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):
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()
@@ -77,7 +77,7 @@ class TestCPPrefixLenFA3Parity(CustomTestCase):
def _call_meta(rank: int): def _call_meta(rank: int):
return prepare_context_parallel_metadata( return prepare_context_parallel_metadata(
padded_extend, rank, cp_size, seqs_len, extend_lens=extend_lens padded_extend, rank, cp_size, seqs_len, extend_seqs_len=extend_lens
) )
# Exercise the non-DSA branch; the DSA branch uses a separate # Exercise the non-DSA branch; the DSA branch uses a separate
@@ -126,8 +126,15 @@ def _cp_attn_for_rank(
kv_len_next_tensor=torch.tensor( kv_len_next_tensor=torch.tensor(
[(b_next + 1) * block_size], dtype=torch.int32, device=DEVICE [(b_next + 1) * block_size], dtype=torch.int32, device=DEVICE
), ),
actual_seq_q_prev=block_size, cu_seqlens_q_prev_tensor=torch.tensor(
actual_seq_q_next=block_size, [0, block_size], dtype=torch.int32, device=DEVICE
),
cu_seqlens_q_next_tensor=torch.tensor(
[0, block_size], dtype=torch.int32, device=DEVICE
),
max_seqlen_q_prev=block_size,
max_seqlen_q_next=block_size,
total_q_prev_tokens=block_size,
) )
fb = SimpleNamespace(attn_cp_metadata=cp_meta) fb = SimpleNamespace(attn_cp_metadata=cp_meta)