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