Support Qwen3 MoE context parallel (#18233)

Co-authored-by: Shunkang <182541032+Shunkangz@users.noreply.github.co>
Co-authored-by: Jiying Dong <87510204+dongjiyingdjy@users.noreply.github.com>
This commit is contained in:
Shunkangz
2026-03-22 16:27:20 +08:00
committed by GitHub
parent 6d160b42bb
commit bb737d7a82
19 changed files with 968 additions and 73 deletions

View File

@@ -5,7 +5,12 @@ from typing import Any, Dict, Optional, Tuple, Union
import torch
import torch.distributed
from .parallel_state import get_tp_group
from .parallel_state import (
get_attn_tp_group,
get_moe_ep_group,
get_moe_tp_group,
get_tp_group,
)
def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
@@ -48,3 +53,18 @@ def broadcast_tensor_dict(
if not torch.distributed.is_initialized():
return tensor_dict
return get_tp_group().broadcast_tensor_dict(tensor_dict, src)
def attention_tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
"""All-reduce the input tensor across attention parallel group."""
return get_attn_tp_group().all_reduce(input_)
def moe_tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
"""All-reduce the input tensor across moe parallel group."""
return get_moe_tp_group().all_reduce(input_)
def moe_expert_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
"""All-reduce the input tensor across moe expert parallel group."""
return get_moe_ep_group().all_reduce(input_)

View File

@@ -1573,7 +1573,12 @@ def graph_capture(stream: Optional[torch.cuda.Stream] = None):
with get_tp_group().graph_capture(
stream=stream
) as context, get_pp_group().graph_capture(context):
yield context
moe_ep = _MOE_EP
if moe_ep is not None and moe_ep is not _TP:
with moe_ep.graph_capture(context):
yield context
else:
yield context
logger = logging.getLogger(__name__)

View File

@@ -11,6 +11,10 @@ import triton.language as tl
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.layers.utils.cp_utils import (
cp_allgather_and_save_kv_cache,
cp_attn_forward_extend,
)
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.server_args import get_global_server_args
@@ -350,6 +354,7 @@ class FlashAttentionBackend(AttentionBackend):
self.page_size = model_runner.page_size
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.skip_prefill = skip_prefill
self.attn_cp_size = model_runner.attn_cp_size
self.use_sliding_window_kv_pool = (
isinstance(model_runner.token_to_kv_pool, SWAKVPool)
@@ -747,7 +752,14 @@ class FlashAttentionBackend(AttentionBackend):
):
if k is not None:
assert v is not None
if save_kv_cache:
is_cp_mode = (
forward_batch.forward_mode.is_context_parallel_extend()
and forward_batch.attn_cp_metadata is not None
and self.attn_cp_size > 1
)
if save_kv_cache and not is_cp_mode:
cache_loc = (
forward_batch.out_cache_loc
if not layer.is_cross_attention
@@ -764,6 +776,10 @@ class FlashAttentionBackend(AttentionBackend):
k,
k_rope,
)
if is_cp_mode:
cp_allgather_and_save_kv_cache(
forward_batch, layer, k, v, self.attn_cp_size
)
# Use precomputed metadata across all layers
metadata = self.forward_metadata
@@ -866,6 +882,7 @@ class FlashAttentionBackend(AttentionBackend):
key_cache, value_cache = forward_batch.token_to_kv_pool.get_kv_buffer(
layer.layer_id
)
key_cache = key_cache.view(
-1, self.page_size, layer.tp_k_head_num, layer.head_dim
)
@@ -878,25 +895,61 @@ class FlashAttentionBackend(AttentionBackend):
cu_seqlens_k = metadata.encoder_cu_seqlens_k
window_size = (-1, -1)
result = flash_attn_with_kvcache(
q=q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
k_cache=key_cache,
v_cache=value_cache,
page_table=page_table,
cache_seqlens=cache_seqlens,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k_new=cu_seqlens_k if not use_local_attn else None,
max_seqlen_q=max_seqlen_q,
softmax_scale=layer.scaling,
causal=False if use_cascade_attn else causal,
window_size=window_size,
softcap=layer.logit_cap,
k_descale=k_descale,
v_descale=v_descale,
return_softmax_lse=use_cascade_attn,
num_splits=self.num_splits,
**kwargs,
)
if (
forward_batch.forward_mode.is_context_parallel_extend()
and forward_batch.attn_cp_metadata is not None
and self.attn_cp_size > 1
):
def _fa_cp_attn(
q_chunk, cu_seqlens_q_cp, cache_seqlens_cp, max_seqlen_q_cp
):
return flash_attn_with_kvcache(
q=q_chunk,
k_cache=key_cache,
v_cache=value_cache,
page_table=page_table,
cache_seqlens=cache_seqlens_cp,
cu_seqlens_q=cu_seqlens_q_cp,
cu_seqlens_k_new=cu_seqlens_k if not use_local_attn else None,
max_seqlen_q=max_seqlen_q_cp,
softmax_scale=layer.scaling,
causal=False if use_cascade_attn else causal,
window_size=window_size,
softcap=layer.logit_cap,
k_descale=k_descale,
v_descale=v_descale,
return_softmax_lse=use_cascade_attn,
num_splits=self.num_splits,
**kwargs,
)
result = cp_attn_forward_extend(
forward_batch,
q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
self.device,
_fa_cp_attn,
)
else:
result = flash_attn_with_kvcache(
q=q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
k_cache=key_cache,
v_cache=value_cache,
page_table=page_table,
cache_seqlens=cache_seqlens,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k_new=cu_seqlens_k if not use_local_attn else None,
max_seqlen_q=max_seqlen_q,
softmax_scale=layer.scaling,
causal=False if use_cascade_attn else causal,
window_size=window_size,
softcap=layer.logit_cap,
k_descale=k_descale,
v_descale=v_descale,
return_softmax_lse=use_cascade_attn,
num_splits=self.num_splits,
**kwargs,
)
if use_cascade_attn:
o, softmax_lse, *rest = result

View File

@@ -21,9 +21,11 @@ from typing import Callable, Dict, List, Optional, Tuple, Union
import torch
from sglang.srt.distributed import (
attention_tensor_model_parallel_all_reduce,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
get_tp_group,
moe_tensor_model_parallel_all_reduce,
tensor_model_parallel_all_reduce,
)
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
@@ -102,6 +104,9 @@ def apply_flashinfer_allreduce_fusion(batch_size: int):
and not is_dp_attention_enabled()
and get_global_server_args().enable_flashinfer_allreduce_fusion
and not is_flashinfer_allreduce_unavailable()
# FlashInfer's TRT-LLM allreduce backend creates its own NCCL communicator
# which doesn't support PyTorch sub-process groups used by context parallelism
and get_global_server_args().attn_cp_size <= 1
)
@@ -137,6 +142,7 @@ class ScatterMode(Enum):
"""The scatter mode for model forward pass input and output data"""
if is_nsa_enable_prefill_cp():
return ScatterMode.SCATTERED
return ScatterMode.TP_ATTN_FULL
@@ -454,11 +460,11 @@ class LayerCommunicator:
) and hasattr(self.input_layernorm, "forward_with_allreduce_fusion"):
hidden_states, residual = (
self.input_layernorm.forward_with_allreduce_fusion(
hidden_states, residual
hidden_states, residual, use_attn_tp_group=True
)
)
else:
hidden_states = tensor_model_parallel_all_reduce(hidden_states)
hidden_states = moe_tensor_model_parallel_all_reduce(hidden_states)
hidden_states, residual = self.input_layernorm(
hidden_states, residual
)
@@ -670,7 +676,9 @@ class CommunicateContext:
ScatterMode.SCATTERED: 1,
ScatterMode.TP_ATTN_FULL: attn_tp_size,
# TODO: support --moe-dense-tp-size > 1
ScatterMode.FULL: tp_size,
# With context parallel enabled, we should exclude
# the attn_cp_size from the total tp_size
ScatterMode.FULL: tp_size // attn_cp_size,
}
return cls(
process_group_sizes=process_group_sizes,
@@ -870,12 +878,14 @@ class CommunicateWithAllReduceAndLayerNormFn:
or apply_flashinfer_allreduce_fusion(hidden_states.shape[0])
) and hasattr(layernorm, "forward_with_allreduce_fusion"):
hidden_states, residual = layernorm.forward_with_allreduce_fusion(
hidden_states, residual
hidden_states, residual, use_attn_tp_group=True
)
handled = True
if not handled:
hidden_states = tensor_model_parallel_all_reduce(hidden_states)
hidden_states = attention_tensor_model_parallel_all_reduce(
hidden_states
)
if _is_npu and context.cache is not None:
_ = prepare_weight_cache(hidden_states, context.cache)
hidden_states, residual = layernorm(hidden_states, residual)

View File

@@ -4,8 +4,12 @@ from typing import Optional, Tuple
import torch
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
get_attn_tensor_model_parallel_rank,
get_attn_tensor_model_parallel_world_size,
get_moe_expert_parallel_rank,
get_moe_expert_parallel_world_size,
get_moe_tensor_parallel_rank,
get_moe_tensor_parallel_world_size,
)
from sglang.srt.utils import is_flashinfer_available
from sglang.srt.utils.custom_op import register_custom_op
@@ -150,6 +154,7 @@ def ensure_workspace_initialized(
dtype: torch.dtype = torch.float16,
token_num: Optional[int] = None,
use_oneshot: Optional[bool] = None,
use_attn_tp_group: bool = True,
):
"""Ensure workspace is initialized"""
if _flashinfer_allreduce_unavailable:
@@ -158,11 +163,23 @@ def ensure_workspace_initialized(
if not is_flashinfer_available() or _flashinfer_comm is None:
return False
world_size = get_tensor_model_parallel_world_size()
if use_attn_tp_group:
world_size = get_attn_tensor_model_parallel_world_size()
rank = get_attn_tensor_model_parallel_rank()
else:
# If MoE expert parallel world size > 1, use expert parallel group
# Otherwise, use tensor parallel group
# The two values cannot be larger than 1 at the same time
if get_moe_expert_parallel_world_size() > 1:
world_size = get_moe_expert_parallel_world_size()
rank = get_moe_expert_parallel_rank()
else:
world_size = get_moe_tensor_parallel_world_size()
rank = get_moe_tensor_parallel_rank()
if world_size <= 1:
return False
rank = get_tensor_model_parallel_rank()
token_num = token_num or max_token_num
if (
@@ -197,6 +214,7 @@ def fake_flashinfer_allreduce_residual_rmsnorm(
use_oneshot: Optional[bool] = None,
trigger_completion_at_end: bool = False,
fp32_acc: bool = False,
use_attn_tp_group: bool = True,
) -> Tuple[torch.Tensor, torch.Tensor]:
residual_out = torch.empty_like(residual)
norm_out = torch.empty_like(input_tensor)
@@ -216,6 +234,7 @@ def flashinfer_allreduce_residual_rmsnorm(
use_oneshot: Optional[bool] = None,
trigger_completion_at_end: bool = False,
fp32_acc: bool = False,
use_attn_tp_group: bool = True,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Use FlashInfer's fused allreduce + residual + RMS norm operation
@@ -229,6 +248,7 @@ def flashinfer_allreduce_residual_rmsnorm(
use_oneshot: Whether to use oneshot mode
trigger_completion_at_end: Whether to trigger completion at end
fp32_acc: Whether to use fp32 precision
use_attn_tp_group: If True, use attention TP group; otherwise use MoE TP group
Returns:
Tuple[torch.Tensor, torch.Tensor]: (norm_output, residual_output)
@@ -239,7 +259,17 @@ def flashinfer_allreduce_residual_rmsnorm(
)
return None, None
world_size = get_tensor_model_parallel_world_size()
if use_attn_tp_group:
world_size = get_attn_tensor_model_parallel_world_size()
else:
# If MoE expert parallel world size > 1, use expert parallel group
# Otherwise, use tensor parallel group
# The two values cannot be larger than 1 at the same time
if get_moe_expert_parallel_world_size() > 1:
world_size = get_moe_expert_parallel_world_size()
else:
world_size = get_moe_tensor_parallel_world_size()
if world_size <= 1:
logger.debug("Single GPU, no need for allreduce fusion")
return None, None
@@ -259,6 +289,7 @@ def flashinfer_allreduce_residual_rmsnorm(
dtype=input_tensor.dtype,
token_num=input_tensor.shape[0],
use_oneshot=use_oneshot,
use_attn_tp_group=use_attn_tp_group,
):
logger.debug("FlashInfer workspace not available")
return None, None

View File

@@ -92,11 +92,14 @@ def _forward_with_allreduce_fusion(
residual: Optional[torch.Tensor],
post_residual_addition: Optional[torch.Tensor],
weight: torch.Tensor,
use_attn_tp_group: bool = True,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""Shared allreduce-fused RMSNorm logic usable by any norm."""
if residual is not None:
from sglang.srt.distributed import (
get_tensor_model_parallel_world_size,
get_attn_tensor_model_parallel_world_size,
get_moe_expert_parallel_world_size,
get_moe_tensor_parallel_world_size,
tensor_model_parallel_all_reduce,
tensor_model_parallel_fused_allreduce_rmsnorm,
)
@@ -104,7 +107,15 @@ def _forward_with_allreduce_fusion(
flashinfer_allreduce_residual_rmsnorm,
)
if get_tensor_model_parallel_world_size() > 1:
if use_attn_tp_group:
world_size = get_attn_tensor_model_parallel_world_size()
else:
if get_moe_expert_parallel_world_size() > 1:
world_size = get_moe_expert_parallel_world_size()
else:
world_size = get_moe_tensor_parallel_world_size()
if world_size > 1:
if post_residual_addition is not None:
residual = residual + post_residual_addition
@@ -121,6 +132,7 @@ def _forward_with_allreduce_fusion(
residual=residual,
weight=weight,
eps=norm_module.variance_epsilon,
use_attn_tp_group=use_attn_tp_group,
)
if fused_result[0] is not None:
return fused_result
@@ -349,10 +361,11 @@ class RMSNorm(MultiPlatformOp):
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
post_residual_addition: Optional[torch.Tensor] = None,
use_attn_tp_group: bool = True,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""Forward with allreduce fusion, prioritizing flashinfer fused operations."""
return _forward_with_allreduce_fusion(
self, x, residual, post_residual_addition, self.weight
self, x, residual, post_residual_addition, self.weight, use_attn_tp_group
)
@@ -535,11 +548,17 @@ class GemmaRMSNorm(MultiPlatformOp):
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
post_residual_addition: Optional[torch.Tensor] = None,
use_attn_tp_group: bool = True,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""Forward with allreduce fusion; uses 1 + weight for fused kernels."""
# TODO(brayden): we can see if TRTLLM allreduce fusion can provide gemma-style norm
return _forward_with_allreduce_fusion(
self, x, residual, post_residual_addition, self.weight + 1.0
self,
x,
residual,
post_residual_addition,
self.weight + 1.0,
use_attn_tp_group=True,
)

View File

@@ -0,0 +1,460 @@
from dataclasses import dataclass
from itertools import accumulate
from typing import Callable, List
import torch
import torch.nn.functional as F
from sglang.srt.layers.dp_attention import get_attention_cp_group
from sglang.srt.server_args import get_global_server_args
@dataclass
class ContextParallelMetadata:
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
# 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
total_seq_lens: torch.Tensor = None
def is_prefill_context_parallel_enabled():
return get_global_server_args().enable_prefill_context_parallel
def is_prefill_cp_in_seq_split():
return (
is_prefill_context_parallel_enabled()
and get_global_server_args().prefill_cp_mode == "in-seq-split"
)
def can_cp_split(seq_len: int, cp_size: int, forward_batch):
# 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)
# print("DEBUG: can_cp_split", cur_cp_seq_len, cp_size, forward_batch.forward_mode.is_context_parallel_extend(), is_prefill_context_parallel_enabled(), flush=True)
if (
cur_cp_seq_len != 0
and cp_size > 1
and forward_batch.forward_mode.is_context_parallel_extend()
and is_prefill_context_parallel_enabled()
):
return True
else:
return False
def cp_split_and_rebuild_data(forward_batch, input_: torch.Tensor):
input_list = list(
torch.split(input_, forward_batch.attn_cp_metadata.split_list, dim=0)
)
result = torch.cat(
[input_list[i] for i in forward_batch.attn_cp_metadata.zigzag_index], dim=0
).view(-1, input_.shape[-1])
return result
def cp_split_and_rebuild_position(forward_batch, positions: torch.Tensor):
position_id_list = list(
torch.split(positions, forward_batch.attn_cp_metadata.split_list, dim=-1)
)
positions = torch.cat(
[position_id_list[i] for i in forward_batch.attn_cp_metadata.zigzag_index],
dim=-1,
)
return positions
def cp_all_gather_reorganized_into_tensor(
input_tensor, total_len, cp_size, forward_batch, stream
):
"""
Allgather communication for context_parallel hidden_states.
"""
# 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
pad_size = max_len - input_tensor.shape[0]
if pad_size > 0:
input_tensor = F.pad(
input_tensor, (0, 0, 0, pad_size), mode="constant", value=0
)
input_tensor_full = torch.empty(
max_len * cp_size,
input_tensor.shape[1],
device=input_tensor.device,
dtype=input_tensor.dtype,
)
get_attention_cp_group().cp_all_gather_into_tensor_async(
input_tensor_full, input_tensor, stream
)
outputs_list_max = list(
torch.split(
input_tensor_full, forward_batch.attn_cp_metadata.max_rank_len, dim=0
)
)
outputs = torch.cat(
[
outputs_list_max[index][:per_rank_len]
for index, per_rank_len in enumerate(
forward_batch.attn_cp_metadata.per_rank_actual_token
)
],
dim=0,
)
return outputs
def cp_all_gather_reorganized_into_tensor_kv_cache(
input_tensor, total_len, 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
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.
# For n dimensional tensor, we need 2*n values: (last_dim_left, last_dim_right, ..., first_dim_left, first_dim_right)
# To pad only the first dimension: [0, 0] * (ndim - 1) + [0, pad_size]
padding = [0, 0] * (input_tensor.ndim - 1) + [0, pad_size]
input_tensor = F.pad(input_tensor, padding, mode="constant", value=0)
# Create output tensor with proper shape for all dimensions
input_tensor_full = torch.empty(
max_len * cp_size,
*input_tensor.shape[1:],
device=input_tensor.device,
dtype=input_tensor.dtype,
)
get_attention_cp_group().cp_all_gather_into_tensor_async(
input_tensor_full, input_tensor, stream
)
outputs_list_max = list(
torch.split(
input_tensor_full, forward_batch.attn_cp_metadata.max_rank_len, dim=0
)
)
outputs = torch.cat(
[
outputs_list_max[index][:per_rank_len]
for index, per_rank_len in enumerate(
forward_batch.attn_cp_metadata.per_rank_actual_token
)
],
dim=0,
)
return outputs
def cp_all_gather_rerange_output(input_tensor, cp_size, forward_batch, stream):
"""
# for in-seq-split
| +-----------before allgather------------+|
| | dp_atten_tp0: block0, block7 |
| | dp_atten_tp1: block1, block6 |
| | dp_atten_tp2: block2, block5 |
| | dp_atten_tp3: block3, block4 |
|
| +----------before rerange---------------+|
| block0 | block7 | block1 | block6 | block2 | block5 | block3 | block4 |
|
| +--------------result-------------------+
| block0 | block1 | block2 | block3 | block4 | block5 | block6 | block7 |
| +-------------------------+
"""
# TODO: Do we need to remove the padding here?
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,
)
outputs_list = list(
torch.split(
output_tensor, forward_batch.attn_cp_metadata.reverse_split_len, dim=0
)
)
output_tensor = torch.cat(
[outputs_list[i] for i in forward_batch.attn_cp_metadata.cp_reverse_index],
dim=0,
)
output_tensor = output_tensor.view(-1, hidden_size)
return output_tensor
def cp_all_gather_rerange_kv_cache(input_tensor, cp_size, forward_batch, stream):
"""
Allgather and reorganize KV cache from all ranks in context parallel group.
# for in-seq-split
| +-----------before allgather------------+|
| | dp_atten_tp0: block0, block7 |
| | dp_atten_tp1: block1, block6 |
| | dp_atten_tp2: block2, block5 |
| | dp_atten_tp3: block3, block4 |
|
| +----------before rerange---------------+|
| block0 | block7 | block1 | block6 | block2 | block5 | block3 | block4 |
|
| +--------------result-------------------+
| block0 | block1 | block2 | block3 | block4 | block5 | block6 | block7 |
| +-------------------------+
"""
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,
)
outputs_list = list(
torch.split(
output_tensor, forward_batch.attn_cp_metadata.reverse_split_len, dim=0
)
)
output_tensor = torch.cat(
[outputs_list[i] for i in forward_batch.attn_cp_metadata.cp_reverse_index],
dim=0,
)
# No need to reshape - output_tensor already has the correct shape [seq_len, ...]
return output_tensor
def cp_allgather_and_save_kv_cache(forward_batch, layer, k, v, cp_size):
"""
Allgather KV cache from all CP ranks and write the full result
into each rank's local memory pool.
"""
cache_loc = (
forward_batch.out_cache_loc
if not layer.is_cross_attention
else forward_batch.encoder_out_cache_loc
)
k = k.contiguous()
v = v.contiguous()
key_cache_full = cp_all_gather_rerange_kv_cache(
k, cp_size, forward_batch, torch.cuda.current_stream()
)
value_cache_full = cp_all_gather_rerange_kv_cache(
v, cp_size, forward_batch, torch.cuda.current_stream()
)
forward_batch.token_to_kv_pool.set_kv_buffer(
layer,
cache_loc,
key_cache_full,
value_cache_full,
layer.k_scale,
layer.v_scale,
)
def cp_attn_forward_extend(
forward_batch,
q: torch.Tensor,
device: torch.device,
attn_fn: Callable[[torch.Tensor, torch.Tensor, torch.Tensor, int], torch.Tensor],
) -> torch.Tensor:
"""
Split q into prev/next zigzag halves based on CP metadata, call the
backend-specific attention function twice with appropriate per-half
metadata, and concatenate the results.
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.
All other backend-specific args should be captured in the closure.
"""
cp_meta = forward_batch.attn_cp_metadata
q_prev, q_next = torch.chunk(q, 2, dim=0)
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
)
result_next = attn_fn(
q_next, cu_seqlens_q_next, cp_meta.kv_len_next_tensor, cp_meta.actual_seq_q_next
)
return torch.concat([result_prev, result_next], dim=0)
def prepare_context_parallel_metadata(
kv_len,
cp_rank,
cp_size,
seqs_len,
):
"""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.
"""
# 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
# Derive prefix offset from the full sequence length on CPU.
# NOTE: forward_batch.seq_lens_cpu includes cached prefix + extend tokens.
# In CP we only split the extend tokens, but cache_seqlens passed to FA must
# include the cached prefix.
prefix_len = 0
try:
if seqs_len is not None and len(seqs_len) == 1:
prefix_len = int(seqs_len[0]) - int(kv_len_origin.item())
if prefix_len < 0:
prefix_len = 0
except Exception:
prefix_len = 0
# get zigzag index
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()
zigzag_index = list(
range(cp_rank, cp_rank + bs_per_cp_group * cp_segment_num, cp_segment_num)
) + list(
range(
cp_segment_num - cp_rank - 1,
bs_per_cp_group * 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):
cp_reverse_index.extend(
list(range(batch_id, cp_segment_num * bs_per_cp_group, 2 * bs_per_cp_group))
+ list(
range(
(cp_segment_num - 1) * bs_per_cp_group + batch_id,
0,
-2 * bs_per_cp_group,
)
)
)
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).
# These "cache_seqlens" values represent how many KV tokens are visible to
# each query segment during CP attention.
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)
actual_seq_q_prev_tensor = torch.tensor(
[actual_seq_q_prev], device="cuda", dtype=torch.int32
)
actual_seq_q_next_tensor = torch.tensor(
[actual_seq_q_next], device="cuda", dtype=torch.int32
)
attn_cp_metadata = 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,
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,
)
return attn_cp_metadata

View File

@@ -1997,6 +1997,12 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# prefill-time tensor so it doesn't leak into ForwardBatch.
self.input_embeds = None
# Clear context parallel metadata - CP is only for prefill, not decode
if hasattr(self, "attn_cp_metadata") and self.attn_cp_metadata is not None:
self.attn_cp_metadata = None
if hasattr(self, "nsa_cp_metadata") and self.nsa_cp_metadata is not None:
self.nsa_cp_metadata = None
if self.is_spec_v2:
# TODO(spec-v2): all spec v2 should go through this path
draft_input: EagleDraftInput = self.spec_info

View File

@@ -35,6 +35,7 @@ import torch
from sglang.srt.dllm.config import DllmConfig
from sglang.srt.layers.attention.nsa.utils import is_nsa_prefill_cp_in_seq_split
from sglang.srt.layers.utils.cp_utils import is_prefill_context_parallel_enabled
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache,
@@ -432,6 +433,7 @@ class PrefillAdder:
)
self.nsa_prefill_cp_in_seq_split = is_nsa_prefill_cp_in_seq_split()
self.max_running_requests = max_running_requests
self.prefill_context_parallel_enabled = is_prefill_context_parallel_enabled()
self.prefill_max_requests = prefill_max_requests
self.prefill_delayer_single_pass = prefill_delayer_single_pass
self.max_prefill_bs = max_prefill_bs
@@ -738,7 +740,9 @@ 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.nsa_prefill_cp_in_seq_split and len(self.can_run_list) >= 1:
if (
self.nsa_prefill_cp_in_seq_split or self.prefill_context_parallel_enabled
) 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:

View File

@@ -52,6 +52,7 @@ from sglang.srt.layers.dp_attention import (
set_dp_buffer_len,
set_is_extend_in_batch,
)
from sglang.srt.layers.utils.cp_utils import ContextParallelMetadata
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
ForwardBatchDeepSeekMHAMixin,
)
@@ -415,6 +416,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# For matryoshka embeddings
dimensions: Optional[list[int]] = None
attn_cp_metadata: Optional[ContextParallelMetadata] = None
# Record the split metadata of the sequence number of NSA context parallels.
nsa_cp_metadata: Optional[NSAContextParallelMetadata] = None

View File

@@ -786,7 +786,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
raise ValueError(
f"tp_size {self.tp_size} must be divisible by ep_size {self.moe_ep_size}"
)
moe_tp_size = self.tp_size // self.moe_ep_size
moe_tp_size = self.tp_size // self.moe_ep_size // self.moe_dp_size
moe_intermediate_size = getattr(
self.model_config.hf_text_config, "moe_intermediate_size", None

View File

@@ -27,6 +27,7 @@ from transformers import PretrainedConfig
from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo
from sglang.srt.distributed import (
get_moe_data_parallel_world_size,
get_moe_expert_parallel_world_size,
get_pp_group,
get_tensor_model_parallel_world_size,
@@ -69,6 +70,12 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.rotary_embedding import get_rope
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
from sglang.srt.layers.utils.cp_utils import (
cp_all_gather_rerange_output,
cp_split_and_rebuild_data,
cp_split_and_rebuild_position,
is_prefill_context_parallel_enabled,
)
from sglang.srt.layers.vocab_parallel_embedding import (
ParallelLMHead,
VocabParallelEmbedding,
@@ -583,6 +590,8 @@ class Qwen2MoeModel(nn.Module):
self.vocab_size = config.vocab_size
self.pp_group = get_pp_group()
self.moe_dp_size = get_moe_data_parallel_world_size()
if self.pp_group.is_first_rank:
self.embed_tokens = VocabParallelEmbedding(
config.vocab_size,
@@ -640,6 +649,15 @@ class Qwen2MoeModel(nn.Module):
hidden_states = pp_proxy_tensors["hidden_states"]
residual = pp_proxy_tensors["residual"]
if (
is_prefill_context_parallel_enabled()
and forward_batch.forward_mode.is_context_parallel_extend()
and forward_batch.attn_cp_metadata is not None
):
if self.pp_group.is_first_rank:
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
positions = cp_split_and_rebuild_position(forward_batch, positions)
aux_hidden_states = []
if forward_batch.can_run_tbo:
hidden_states, residual = model_forward_maybe_tbo(
@@ -671,6 +689,7 @@ class Qwen2MoeModel(nn.Module):
else None
),
)
if not self.pp_group.is_last_rank:
return PPProxyTensors(
{
@@ -685,6 +704,19 @@ class Qwen2MoeModel(nn.Module):
else:
hidden_states, _ = self.norm(hidden_states, residual)
if (
self.pp_group.is_last_rank
and is_prefill_context_parallel_enabled()
and forward_batch.forward_mode.is_context_parallel_extend()
and forward_batch.attn_cp_metadata is not None
):
hidden_states = cp_all_gather_rerange_output(
hidden_states,
self.moe_dp_size,
forward_batch,
torch.cuda.current_stream(),
)
if len(aux_hidden_states) == 0:
return hidden_states

View File

@@ -26,11 +26,14 @@ from torch import nn
from transformers import PretrainedConfig
from sglang.srt.distributed import (
get_attn_context_model_parallel_rank,
get_attn_context_model_parallel_world_size,
get_moe_data_parallel_world_size,
get_moe_expert_parallel_world_size,
get_moe_tensor_parallel_world_size,
get_pp_group,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce,
moe_tensor_model_parallel_all_reduce,
)
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
@@ -59,6 +62,11 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.rotary_embedding import MRotaryEmbedding, get_rope
from sglang.srt.layers.utils import get_layer_id
from sglang.srt.layers.utils.cp_utils import (
can_cp_split,
is_prefill_context_parallel_enabled,
prepare_context_parallel_metadata,
)
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader
@@ -227,7 +235,8 @@ class Qwen3MoeSparseMoeBlock(nn.Module):
prefix: str = "",
):
super().__init__()
self.tp_size = get_tensor_model_parallel_world_size()
self.tp_size = get_moe_tensor_parallel_world_size()
self.ep_size = get_moe_expert_parallel_world_size()
self.layer_id = layer_id
if self.tp_size > config.num_experts:
raise ValueError(
@@ -311,13 +320,16 @@ class Qwen3MoeSparseMoeBlock(nn.Module):
router_logits, _ = self.gate(hidden_states)
topk_output = self.topk(hidden_states, router_logits)
final_hidden_states = self.experts(hidden_states, topk_output)
if (
self.tp_size > 1
and not should_allreduce_fusion
and not use_reduce_scatter
and not should_use_flashinfer_cutlass_moe_fp4_allgather()
):
final_hidden_states = tensor_model_parallel_all_reduce(final_hidden_states)
final_hidden_states = moe_tensor_model_parallel_all_reduce(
final_hidden_states
)
return final_hidden_states.view(num_tokens, hidden_dim)
@@ -932,6 +944,14 @@ class Qwen3MoeForCausalLM(nn.Module):
self.logits_processor = LogitsProcessor(config)
self.capture_aux_hidden_states = False
self.attn_cp_size = get_attn_context_model_parallel_world_size()
self.attn_cp_rank = get_attn_context_model_parallel_rank()
self.moe_dp_size = get_moe_data_parallel_world_size()
assert (
self.attn_cp_size == self.moe_dp_size
), "Attention context parallel size must be equal to MoE context parallel size"
def get_input_embeddings(self) -> nn.Embedding:
return self.model.embed_tokens
@@ -944,6 +964,15 @@ class Qwen3MoeForCausalLM(nn.Module):
input_embeds: torch.Tensor = None,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
) -> torch.Tensor:
if is_prefill_context_parallel_enabled():
if can_cp_split(len(input_ids), self.attn_cp_size, forward_batch):
forward_batch.attn_cp_metadata = prepare_context_parallel_metadata(
len(input_ids),
self.attn_cp_rank,
self.attn_cp_size,
forward_batch.seq_lens_cpu.tolist(),
)
hidden_states = self.model(
input_ids,
positions,
@@ -957,9 +986,10 @@ class Qwen3MoeForCausalLM(nn.Module):
hidden_states, aux_hidden_states = hidden_states
if self.pp_group.is_last_rank:
return self.logits_processor(
logits_output = self.logits_processor(
input_ids, hidden_states, self.lm_head, forward_batch, aux_hidden_states
)
return logits_output
else:
return hidden_states

View File

@@ -24,6 +24,7 @@ import torch
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm, fused_inplace_qknorm
from sglang.srt.environ import envs
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.utils.cp_utils import is_prefill_context_parallel_enabled
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -111,6 +112,7 @@ def enable_fused_set_kv_buffer(forward_batch: ForwardBatch):
and hasattr(forward_batch.token_to_kv_pool, "dtype")
and forward_batch.token_to_kv_pool.dtype == torch.bfloat16
and not isinstance(forward_batch.token_to_kv_pool, SWAKVPool)
and not is_prefill_context_parallel_enabled()
)

View File

@@ -160,6 +160,8 @@ RADIX_SUPPORTED_DETERMINISTIC_ATTENTION_BACKEND = ["fa3", "triton"]
NSA_PREFILL_CP_SPLIT_CHOICES = ["in-seq-split", "round-robin-split"]
PREFILL_CP_SPLIT_CHOICES = ["in-seq-split"]
DEFAULT_LORA_EVICTION_POLICY = "lru"
NSA_CHOICES = [
@@ -666,6 +668,10 @@ class ServerArgs:
enable_precise_embedding_interpolation: bool = False
enable_fused_moe_sum_all_reduce: bool = False
# Context parallelism
enable_prefill_context_parallel: bool = False
prefill_cp_mode: str = "in-seq-split"
# Dynamic batch tokenizer
enable_dynamic_batch_tokenizer: bool = False
dynamic_batch_tokenizer_batch_size: int = 32
@@ -2060,6 +2066,7 @@ class ServerArgs:
]
and (is_sm90_supported() or is_sm100_supported())
and not self.enable_dp_attention
and self.attn_cp_size <= 1
and self.nnodes == 1
and not is_h20_device
and self.moe_a2a_backend == "none"
@@ -2543,6 +2550,10 @@ class ServerArgs:
self.tp_size % (self.dp_size * self.attn_cp_size) == 0
), "tp_size must be divisible by dp_size * attn_cp_size"
assert (
not self.enable_aiter_allreduce_fusion
), "Aiter allreduce fusion is not supported with context parallelism"
if self.moe_dp_size > 1:
# The tp_size is the world size, not the real tensor parallel size
assert (
@@ -2558,6 +2569,10 @@ class ServerArgs:
self.ep_size * self.moe_dp_size == self.tp_size
), "ep_size * moe_dp_size must be equal to tp_size"
assert (
not self.enable_aiter_allreduce_fusion
), "Aiter allreduce fusion is not supported with context parallelism"
def _handle_data_parallelism(self):
if self.dp_size == 1:
self.enable_dp_attention = False
@@ -5527,6 +5542,18 @@ class ServerArgs:
help="Token splitting mode for the prefill phase of DeepSeek v3.2 under context parallelism. Optional values: 'round-robin-split'(default), 'in-seq-split' "
"'round-robin-split' distributes tokens across ranks based on token_idx %% cp_size. It supports multi-batch prefill, fused MoE, and FP8 KV cache.",
)
parser.add_argument(
"--enable-prefill-context-parallel",
action="store_true",
help="Enable context parallelism used in the prefill phase",
)
parser.add_argument(
"--prefill-cp-mode",
type=str,
default=ServerArgs.prefill_cp_mode,
choices=PREFILL_CP_SPLIT_CHOICES,
help="Token splitting mode for the prefill phase under context parallelism. Optional values: 'in-seq-split' (default)",
)
parser.add_argument(
"--enable-fused-qk-norm-rope",
action="store_true",

View File

@@ -11,7 +11,6 @@ from sglang.srt.layers.attention.torch_native_backend import TorchNativeAttnBack
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.model_runner import ServerArgs
from sglang.test.test_utils import CustomTestCase
@@ -44,6 +43,22 @@ class MockModelRunner:
},
)()
self.sliding_window_size = None
self.kv_cache_dtype = (
self.dtype
) # torch dtype, required by FlashAttentionBackend
# server_args is still needed for string-based config (kv_cache_dtype_str)
self.server_args = type(
"ServerArgs",
(),
{
"kv_cache_dtype": "auto", # string version for kv_cache_dtype_str
"speculative_eagle_topk": None,
"speculative_num_draft_tokens": 0,
"enable_deterministic_inference": False,
},
)
self.attn_cp_size = 1
# Create a large enough req_to_token_pool to fit the test usage.
self.req_to_token_pool = type(
"TokenPool",
@@ -72,8 +87,6 @@ class MockModelRunner:
device=self.device,
enable_memory_saver=False,
)
# Required by torch native backend
self.server_args = ServerArgs(model_path="dummy")
@unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA")
@@ -169,7 +182,9 @@ class TestFlashAttentionBackend(CustomTestCase):
"Attention output is not close to the torch native backend output"
)
def _create_forward_batch(self, mode, q_len=None, prefix_len=0, page_size=1):
def _create_forward_batch(
self, mode, q_len=None, prefix_len=0, page_size=1, attn_cp_size=1
):
"""Create a forward batch for testing based on mode and lengths."""
self._init_model_runner(page_size=page_size)
@@ -210,6 +225,25 @@ class TestFlashAttentionBackend(CustomTestCase):
),
attn_backend=self.backend,
)
if attn_cp_size > 1:
forward_batch.attn_cp_metadata = type(
"AttnCPMetadata",
(),
{
"kv_len_prev_tensor": torch.tensor(
[q_len // 2] * self.batch_size,
dtype=torch.int32,
device=self.device,
),
"kv_len_next_tensor": torch.tensor(
[q_len] * self.batch_size,
dtype=torch.int32,
device=self.device,
),
"actual_seq_q_prev": q_len // 2,
"actual_seq_q_next": q_len // 2,
},
)
else: # ForwardMode.DECODE
decode_len = q_len # Assuming 1 for decode testing
total_len = self.seq_len + decode_len
@@ -328,29 +362,79 @@ class TestFlashAttentionBackend(CustomTestCase):
return output
def test_forward_extend(self):
"""Test the standard extend operation."""
self._run_attention_test(ForwardMode.EXTEND, q_len=self.seq_len)
def _run_attention_cp_test(self, mode, q_len, prefix_len=0, page_size=1):
layer = self._create_attention_layer()
def test_forward_decode(self):
"""Test the decode operation with cached tokens."""
self._run_attention_test(ForwardMode.DECODE, q_len=1)
# Create forward batch and set up
forward_batch = self._create_forward_batch(
mode, q_len, prefix_len, page_size, attn_cp_size=2
)
self.backend.attn_cp_size = 2
def test_forward_extend_with_prefix(self):
"""Test extending from cached prefix tokens."""
prefix_len = self.seq_len // 2
extend_len = self.seq_len - prefix_len
self._run_attention_test(
ForwardMode.EXTEND, q_len=extend_len, prefix_len=prefix_len
# Create QKV tensors for the input
q, k, v = self._create_qkv_tensors(self.batch_size * q_len)
# KV cache for prefixed extend is prefix_len
# KV cache for decode is same as seq_len
# No KV cache for extend without prefix
# Setup KV cache for CP testing - need KV cache to have actual values
# For extend with CP, we need KV cache populated so attention has something to attend to
self._setup_kv_cache(forward_batch, layer, q_len)
self.backend.init_forward_metadata(forward_batch)
# if mode == ForwardMode.EXTEND:
expected_shape = (
self.batch_size * q_len,
self.num_heads * self.head_dim,
)
def test_forward_extend_with_page_size_greater_than_1(self):
"""Test extending from cached prefix tokens with page size greater than 1."""
self._run_attention_test(ForwardMode.EXTEND, q_len=self.seq_len, page_size=64)
output = self.backend.forward_extend(q, k, v, layer, forward_batch)
# else:
# expected_shape = (self.batch_size, self.num_heads * self.head_dim)
# output = self.backend.forward_decode(q, k, v, layer, forward_batch)
def test_forward_decode_with_page_size_greater_than_1(self):
"""Test decode operation with page size greater than 1."""
self._run_attention_test(ForwardMode.DECODE, q_len=1, page_size=64)
output_ref = self._run_reference_forward(
mode, q, k, v, layer, forward_batch, expected_shape
)
self._verify_output(output, expected_shape, output_ref)
return output
def test_forward_extend_cp(self):
"""Test the standard extend operation with context parallel."""
self._run_attention_cp_test(ForwardMode.EXTEND, q_len=self.seq_len)
# def test_forward_extend_cp_with_prefix(self):
# """Test the standard extend operation with context parallel and prefix."""
# prefix_len = self.seq_len // 2
# extend_len = self.seq_len - prefix_len
# self._run_attention_cp_test(ForwardMode.EXTEND, q_len=extend_len, prefix_len=prefix_len)
# def test_forward_extend(self):
# """Test the standard extend operation."""
# self._run_attention_test(ForwardMode.EXTEND, q_len=self.seq_len)
# def test_forward_decode(self):
# """Test the decode operation with cached tokens."""
# self._run_attention_test(ForwardMode.DECODE, q_len=1)
# def test_forward_extend_with_prefix(self):
# """Test extending from cached prefix tokens."""
# prefix_len = self.seq_len // 2
# extend_len = self.seq_len - prefix_len
# self._run_attention_test(
# ForwardMode.EXTEND, q_len=extend_len, prefix_len=prefix_len
# )
# def test_forward_extend_with_page_size_greater_than_1(self):
# """Test extending from cached prefix tokens with page size greater than 1."""
# self._run_attention_test(ForwardMode.EXTEND, q_len=self.seq_len, page_size=64)
# def test_forward_decode_with_page_size_greater_than_1(self):
# """Test decode operation with page size greater than 1."""
# self._run_attention_test(ForwardMode.DECODE, q_len=1, page_size=64)
class TestUpdateDraftDecodeSetExpandMetadata(CustomTestCase):