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:
co-authored by
Shunkang
Khoa Pham
Baizhou Zhang
parent
ddf0627254
commit
19663aafcd
@@ -1455,10 +1455,14 @@ class Indexer(MultiPlatformOp):
|
||||
forward_batch.attn_cp_metadata is not None
|
||||
and is_dsa_prefill_cp_in_seq_split()
|
||||
):
|
||||
kv_len_prev = forward_batch.attn_cp_metadata.kv_len_prev
|
||||
kv_len_next = forward_batch.attn_cp_metadata.kv_len_next
|
||||
actual_seq_q_prev = forward_batch.attn_cp_metadata.actual_seq_q_prev
|
||||
actual_seq_q_next = forward_batch.attn_cp_metadata.actual_seq_q_next
|
||||
kv_len_prev = forward_batch.attn_cp_metadata.kv_len_prev_list[0]
|
||||
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_list[0]
|
||||
)
|
||||
actual_seq_q_next = (
|
||||
forward_batch.attn_cp_metadata.actual_seq_q_next_list[0]
|
||||
)
|
||||
|
||||
# TODO support mutil-batch
|
||||
# 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)
|
||||
attn_cp_size = get_attention_cp_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(
|
||||
forward_batch.is_extend_in_batch, global_num_tokens
|
||||
)
|
||||
|
||||
@@ -20,24 +20,41 @@ from sglang.srt.server_args import get_global_server_args
|
||||
|
||||
@dataclass
|
||||
class ContextParallelMetadata:
|
||||
# Layout lists have length bs * cp_segment_num (= bs * 2 * cp_size).
|
||||
split_list: List[int] = None
|
||||
max_rank_len: 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
|
||||
reverse_split_len: List[int] = None
|
||||
|
||||
# metadata for attention
|
||||
kv_len_prev: int = -1
|
||||
kv_len_next: int = -1
|
||||
actual_seq_q_prev: int = -1
|
||||
actual_seq_q_next: int = -1
|
||||
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
|
||||
# Per-rank-aggregate lists have length cp_size.
|
||||
# max_rank_len is a list of cp_size copies of max(per_rank_actual_token),
|
||||
# kept as a list for torch.split() bucket sizes.
|
||||
per_rank_actual_token: List[int] = None
|
||||
max_rank_len: List[int] = 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():
|
||||
@@ -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):
|
||||
# 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
|
||||
|
||||
# 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)
|
||||
return (
|
||||
if not (
|
||||
cur_cp_seq_len != 0
|
||||
and cp_size > 1
|
||||
# prepare_context_parallel_metadata hard-codes bs_per_cp_group = 1;
|
||||
# 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()
|
||||
# is_context_parallel_extend() returns True for MIXED (prefill+decode
|
||||
# in one step), but the zigzag split only makes sense on pure extend.
|
||||
and forward_batch.forward_mode != ForwardMode.MIXED
|
||||
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):
|
||||
@@ -134,9 +171,7 @@ def cp_split_and_rebuild_position(forward_batch, positions: torch.Tensor):
|
||||
return positions
|
||||
|
||||
|
||||
def cp_all_gather_reorganized_into_tensor(
|
||||
input_tensor, total_len, cp_size, forward_batch, stream
|
||||
):
|
||||
def cp_all_gather_reorganized_into_tensor(input_tensor, cp_size, forward_batch, stream):
|
||||
"""
|
||||
Allgather communication for context_parallel(kv_cache, index_k, hidden_states).
|
||||
This implementation mainly consists of three parts:
|
||||
@@ -144,10 +179,7 @@ def cp_all_gather_reorganized_into_tensor(
|
||||
Step 2, allgather communication(async).
|
||||
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.
|
||||
# No need to pad again.
|
||||
# step1
|
||||
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]
|
||||
if pad_size > 0:
|
||||
input_tensor = F.pad(
|
||||
@@ -186,13 +218,13 @@ def cp_all_gather_reorganized_into_tensor(
|
||||
|
||||
|
||||
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.
|
||||
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]
|
||||
if pad_size > 0:
|
||||
# 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
|
||||
output_tensor = cp_all_gather_reorganized_into_tensor(
|
||||
input_tensor,
|
||||
forward_batch.attn_cp_metadata.total_seq_lens,
|
||||
cp_size,
|
||||
forward_batch,
|
||||
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(
|
||||
input_tensor,
|
||||
forward_batch.attn_cp_metadata.total_seq_lens,
|
||||
cp_size,
|
||||
forward_batch,
|
||||
stream,
|
||||
@@ -386,6 +416,11 @@ def cp_attn_forward_extend(
|
||||
backend-specific attention function twice with appropriate per-half
|
||||
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(q, cu_seqlens_q, cache_seqlens, max_seqlen_q) -> result
|
||||
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
|
||||
|
||||
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(
|
||||
q_prev, cu_seqlens_q_prev, cp_meta.kv_len_prev_tensor, cp_meta.actual_seq_q_prev
|
||||
)
|
||||
|
||||
cu_seqlens_q_next = torch.tensor(
|
||||
[0, cp_meta.actual_seq_q_next], device=device, dtype=torch.int32
|
||||
q_prev,
|
||||
cp_meta.cu_seqlens_q_prev_tensor,
|
||||
cp_meta.kv_len_prev_tensor,
|
||||
cp_meta.max_seqlen_q_prev,
|
||||
)
|
||||
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)
|
||||
@@ -417,7 +452,8 @@ def prepare_context_parallel_metadata(
|
||||
cp_rank,
|
||||
cp_size,
|
||||
seqs_len,
|
||||
extend_lens,
|
||||
extend_seqs_len=None,
|
||||
device="cuda",
|
||||
):
|
||||
from sglang.srt.layers.attention.dsa.utils import (
|
||||
is_dsa_prefill_cp_round_robin_split,
|
||||
@@ -427,144 +463,191 @@ def prepare_context_parallel_metadata(
|
||||
return ContextParallelMetadata()
|
||||
|
||||
"""prepare_input_dp_with_cp_dsa-zigzag index
|
||||
Example (DP_ATTENT_TP == CP_SIZE == 4):
|
||||
Description:
|
||||
1. Start with a full-length request.
|
||||
2. Split the request into multiple blocks (block0 to block7).
|
||||
3. Rearrange these blocks to balance computational
|
||||
load across different DP ranks.
|
||||
4. Assign the rearranged blocks to different DP attention
|
||||
time points (dp_atten_tp0 to dp_atten_tp3).
|
||||
+---------------------------------+
|
||||
| cp_split_tokens |
|
||||
+---------------------------------+
|
||||
| |
|
||||
| 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.
|
||||
Example (DP_ATTENT_TP == CP_SIZE == 4, single sequence):
|
||||
block0 | block1 | block2 | block3 | block4 | block5 | block6 | block7
|
||||
rank 0: block0, block7
|
||||
rank 1: block1, block6
|
||||
rank 2: block2, block5
|
||||
rank 3: block3, block4
|
||||
For bs > 1, each sequence is split into cp_segment_num = 2 * cp_size
|
||||
blocks independently; per-rank layout becomes:
|
||||
[s0.block_r, s1.block_r, ..., s_{bs-1}.block_r,
|
||||
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.
|
||||
"""
|
||||
# just support batch = 1
|
||||
# kv_len: the number of tokens *computed in this extend pass* (i.e. the
|
||||
# "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
|
||||
assert extend_seqs_len is not None
|
||||
extend_seqs_len = [int(x) for x in extend_seqs_len]
|
||||
|
||||
# Derive prefix offset from unpadded CPU tensors. Both `seqs_len` and `extend_lens` are unpadded by the caller
|
||||
# Using the padded `kv_len` here would undercount `prefix_len` by the padding amount and shift the FA causal horizon.
|
||||
assert (
|
||||
len(seqs_len) == 1 and len(extend_lens) == 1
|
||||
), "Prefill Context Parallel only supports batch_size == 1 for now"
|
||||
prefix_len = max(0, int(seqs_len[0]) - int(extend_lens[0]))
|
||||
# get zigzag index
|
||||
# Update the extend_seqs_len to the padded length.
|
||||
pad_len = int(kv_len) - 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 = 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
|
||||
max_rank_len = seq_max_rank_len.repeat_interleave(cp_size).int().tolist()
|
||||
# Prefix offset (radix cache hit length) per sequence. For non-NSA
|
||||
# (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(
|
||||
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(
|
||||
range(
|
||||
cp_segment_num - cp_rank - 1,
|
||||
bs_per_cp_group * cp_segment_num,
|
||||
bs * cp_segment_num,
|
||||
cp_segment_num,
|
||||
)
|
||||
)
|
||||
|
||||
per_rank_actual_token = list(
|
||||
split_list[i] + split_list[cp_size * 2 - i - 1] for i in range(cp_size)
|
||||
)
|
||||
reverse_split_len = [
|
||||
element
|
||||
for i in range(cp_size)
|
||||
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):
|
||||
# Reverse index: given the post-allgather concatenation
|
||||
# [rank0_prevs_all_seqs, rank0_nexts_all_seqs,
|
||||
# rank1_prevs_all_seqs, rank1_nexts_all_seqs, ...]
|
||||
# produce a permutation that restores [s0_b0..s0_bN, s1_b0..s1_bN, ...].
|
||||
cp_reverse_index: List[int] = []
|
||||
for batch_id in range(bs):
|
||||
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(
|
||||
range(
|
||||
(cp_segment_num - 1) * bs_per_cp_group + batch_id,
|
||||
(cp_segment_num - 1) * bs + batch_id,
|
||||
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
|
||||
# Prefix offset is critical when radix cache hits (prefix_len > 0).
|
||||
# For non-DSA CP (e.g. qwen3-moe), consumers use these values directly as
|
||||
# FlashAttention cache_seqlens, so the prefix must be baked in here.
|
||||
# For DSA CP, `_get_topk_ragged_with_cp` re-adds the cached-prefix offset
|
||||
# from (seq_lens_cpu - extend_seq_lens_cpu); baking prefix_len in here
|
||||
# would silently drop it whenever the scheduler packs multiple requests
|
||||
# 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.
|
||||
# Split sizes matching the post-allgather concatenation order above.
|
||||
reverse_split_len: List[int] = []
|
||||
for r in range(cp_size):
|
||||
for s in range(bs):
|
||||
reverse_split_len.append(per_seq_block_sizes[s][r])
|
||||
for s in range(bs):
|
||||
reverse_split_len.append(per_seq_block_sizes[s][cp_segment_num - 1 - r])
|
||||
|
||||
# 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
|
||||
|
||||
if is_dsa_enable_prefill_cp():
|
||||
kv_len_prev = prefix_sum_list[cp_rank]
|
||||
kv_len_next = prefix_sum_list[cp_size * 2 - cp_rank - 1]
|
||||
else:
|
||||
kv_len_prev = prefix_len + prefix_sum_list[cp_rank]
|
||||
kv_len_next = prefix_len + prefix_sum_list[cp_size * 2 - cp_rank - 1]
|
||||
actual_seq_q_prev = split_list[cp_rank]
|
||||
actual_seq_q_next = split_list[cp_size * 2 - cp_rank - 1]
|
||||
# Flash Attention expects cache_seqlens to have shape (batch_size,), not scalar
|
||||
kv_len_prev_tensor = torch.tensor([kv_len_prev], device="cuda", dtype=torch.int32)
|
||||
kv_len_next_tensor = torch.tensor([kv_len_next], device="cuda", dtype=torch.int32)
|
||||
nsa_mode = is_dsa_enable_prefill_cp()
|
||||
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 s in range(bs):
|
||||
blk = per_seq_block_sizes[s]
|
||||
cum_prev = sum(blk[: cp_rank + 1])
|
||||
cum_next = sum(blk[: cp_segment_num - cp_rank])
|
||||
# NSA indexer re-adds prefix offset itself; leave bare cumulative.
|
||||
# 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], 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], 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,
|
||||
max_rank_len=max_rank_len,
|
||||
zigzag_index=zigzag_index,
|
||||
per_rank_actual_token=per_rank_actual_token,
|
||||
reverse_split_len=reverse_split_len,
|
||||
cp_reverse_index=cp_reverse_index,
|
||||
kv_len_prev=kv_len_prev,
|
||||
kv_len_next=kv_len_next,
|
||||
actual_seq_q_prev=actual_seq_q_prev,
|
||||
actual_seq_q_next=actual_seq_q_next,
|
||||
reverse_split_len=reverse_split_len,
|
||||
per_rank_actual_token=per_rank_actual_token,
|
||||
max_rank_len=max_rank_len,
|
||||
kv_len_prev_tensor=kv_len_prev_tensor,
|
||||
kv_len_next_tensor=kv_len_next_tensor,
|
||||
actual_seq_q_prev_tensor=actual_seq_q_prev_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
|
||||
# Enabling context parallelism currently presents precision issues;
|
||||
# therefore, the prefill-batch setting is temporarily set to 1.
|
||||
if (
|
||||
self.dsa_prefill_cp_in_seq_split or self.prefill_context_parallel_enabled
|
||||
) and len(self.can_run_list) >= 1:
|
||||
if (self.dsa_prefill_cp_in_seq_split) and len(self.can_run_list) >= 1:
|
||||
return AddReqResult.OTHER
|
||||
|
||||
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
|
||||
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()
|
||||
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(
|
||||
self.is_extend_in_batch, global_num_tokens
|
||||
|
||||
@@ -311,7 +311,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
|
||||
self.cp_rank,
|
||||
self.cp_size,
|
||||
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:
|
||||
if can_cp_split(len(input_ids), self.cp_size, forward_batch):
|
||||
@@ -320,7 +320,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
|
||||
self.cp_rank,
|
||||
self.cp_size,
|
||||
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)
|
||||
return self.logits_processor(
|
||||
|
||||
@@ -2609,7 +2609,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
|
||||
self.cp_rank,
|
||||
self.cp_size,
|
||||
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:
|
||||
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_size,
|
||||
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):
|
||||
|
||||
@@ -1577,7 +1577,7 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
self.cp_rank,
|
||||
self.cp_size,
|
||||
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():
|
||||
attn_backend = get_attn_backend()
|
||||
|
||||
@@ -249,7 +249,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM):
|
||||
self.cp_rank,
|
||||
self.cp_size,
|
||||
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():
|
||||
attn_backend = get_attn_backend()
|
||||
|
||||
@@ -1005,7 +1005,7 @@ class Qwen3MoeForCausalLM(nn.Module):
|
||||
self.attn_cp_rank,
|
||||
self.attn_cp_size,
|
||||
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(
|
||||
|
||||
@@ -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):
|
||||
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
|
||||
|
||||
@@ -126,8 +126,15 @@ def _cp_attn_for_rank(
|
||||
kv_len_next_tensor=torch.tensor(
|
||||
[(b_next + 1) * block_size], dtype=torch.int32, device=DEVICE
|
||||
),
|
||||
actual_seq_q_prev=block_size,
|
||||
actual_seq_q_next=block_size,
|
||||
cu_seqlens_q_prev_tensor=torch.tensor(
|
||||
[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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user