|
|
|
@@ -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
|
|
|
|
|