diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index d18e4dd69..72fec28aa 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -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"] diff --git a/python/sglang/srt/layers/attention/dsa/utils.py b/python/sglang/srt/layers/attention/dsa/utils.py index b49631c97..9daccd870 100644 --- a/python/sglang/srt/layers/attention/dsa/utils.py +++ b/python/sglang/srt/layers/attention/dsa/utils.py @@ -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 ) diff --git a/python/sglang/srt/layers/utils/cp_utils.py b/python/sglang/srt/layers/utils/cp_utils.py index b87225dc7..d599c5491 100644 --- a/python/sglang/srt/layers/utils/cp_utils.py +++ b/python/sglang/srt/layers/utils/cp_utils.py @@ -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 diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index 1ed7bd9ff..5b0c41b3a 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -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: diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 6698429aa..f59b90ffe 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -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 diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index 1e1697252..56e661e80 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -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( diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 8b254348a..564f38d6f 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -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): diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 2a89d22f4..e041e3655 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -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() diff --git a/python/sglang/srt/models/deepseek_v4_nextn.py b/python/sglang/srt/models/deepseek_v4_nextn.py index 6b5c89e50..069a4ec8c 100644 --- a/python/sglang/srt/models/deepseek_v4_nextn.py +++ b/python/sglang/srt/models/deepseek_v4_nextn.py @@ -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() diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 1e31930fb..6179a30ee 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -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( diff --git a/test/registered/cp/test_qwen3_30b.py b/test/registered/cp/test_qwen3_30b.py deleted file mode 100644 index e21dafb5c..000000000 --- a/test/registered/cp/test_qwen3_30b.py +++ /dev/null @@ -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() diff --git a/test/registered/kernels/test_cp_prefix_len_fa3_parity.py b/test/registered/kernels/test_cp_prefix_len_fa3_parity.py index 11d7acbe5..d44f99552 100644 --- a/test/registered/kernels/test_cp_prefix_len_fa3_parity.py +++ b/test/registered/kernels/test_cp_prefix_len_fa3_parity.py @@ -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 diff --git a/test/registered/kernels/test_mla_cp_fa3_parity.py b/test/registered/kernels/test_mla_cp_fa3_parity.py index 9e2a4564a..e4d841bb2 100644 --- a/test/registered/kernels/test_mla_cp_fa3_parity.py +++ b/test/registered/kernels/test_mla_cp_fa3_parity.py @@ -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)