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
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
)
+231 -148
View File
@@ -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
+2 -2
View File
@@ -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(
+2 -2
View File
@@ -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):
+1 -1
View File
@@ -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()
+1 -1
View File
@@ -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(