(1/n)support context parallel with deepseekv3.2-DSA (#12065)

This commit is contained in:
lixiaolx
2025-11-17 12:12:25 +08:00
committed by GitHub
parent 7e626d12b7
commit d368c7451a
17 changed files with 1247 additions and 54 deletions

View File

@@ -209,6 +209,34 @@ class PyNcclCommunicator:
cudaStream_t(stream.cuda_stream),
)
def cp_all_gather_into_tensor(
self,
output_tensor: torch.Tensor,
input_tensor: torch.Tensor,
stream=None,
sizes: Optional[list[int]] = None,
):
"""
Currently, it is mainly used in context parallelism,
primarily leveraging pynccl to implement non-blocking allgather communication.
"""
# nccl communicator created on a specific device
# will only work on tensors on the same device
# otherwise it will cause "illegal memory access"
assert input_tensor.device == self.device, (
f"this nccl communicator is created to work on {self.device}, "
f"but the input tensor is on {input_tensor.device}"
)
stream = self._resolve_stream(stream)
self.nccl.ncclAllGather(
buffer_type(input_tensor.data_ptr()),
buffer_type(output_tensor.data_ptr()),
input_tensor.numel(),
ncclDataTypeEnum.from_torch(input_tensor.dtype),
self.comm,
cudaStream_t(stream.cuda_stream),
)
def reduce_scatter(
self,
output_tensor: torch.Tensor,

View File

@@ -748,6 +748,27 @@ class GroupCoordinator:
output, input, group_name=self.unique_name
)
def cp_all_gather_into_tensor_async(
self, output: torch.Tensor, input: torch.Tensor, stream=None
):
"""
Implement an asynchronous `allgather` operation on a specified stream.
(the default `torch.distributed.all_gather_into_tensor` will trigger event synchronization),
eliminating the CPU-side launch-kernel blocking issue caused by synchronization problems.
The specific implementation uses the interface provided by pynccl to remove the synchronization logic of events.
"""
assert (
stream is not None
), f"Invalid params stream ({stream}, Please specify the stream to use when calling cp_all_gather_into_tensor_async.)"
pynccl_comm = self.pynccl_comm
if pynccl_comm is not None:
pynccl_comm.cp_all_gather_into_tensor(output, input, stream=stream)
else:
logger.warning("not all_gather_into_tensor_async")
torch.ops.sglang.reg_all_gather_into_tensor(
output, input, group_name=self.unique_name
)
def all_gather(
self,
input_: torch.Tensor,

View File

@@ -1,7 +1,7 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Dict, Optional
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
import torch
from einops import rearrange
@@ -16,9 +16,18 @@ if is_cuda():
except ImportError as e:
deep_gemm = e
from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.attention.nsa.utils import NSA_DUAL_STREAM
from sglang.srt.layers.dp_attention import get_attention_tp_group
from sglang.srt.layers.attention.nsa.utils import (
NSA_DUAL_STREAM,
cp_all_gather_rerange_output,
is_nsa_enable_prefill_cp,
)
from sglang.srt.layers.dp_attention import (
get_attention_tp_group,
get_attention_tp_rank,
get_attention_tp_size,
)
from sglang.srt.layers.linear import ReplicatedLinear
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.rotary_embedding import get_rope_wrapper
@@ -112,6 +121,13 @@ class Indexer(CustomOp):
self.layer_id = layer_id
self.alt_stream = alt_stream
self.fuse_wk_and_weights_proj = fuse_wk_and_weights_proj
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
if self.nsa_enable_prefill_cp:
self.cp_size = get_attention_tp_size()
self.cp_rank = get_attention_tp_rank()
else:
self.cp_size = None
self.cp_rank = None
if is_cuda():
self.sm_count = deep_gemm.get_num_sms()
self.half_device_sm_count = ceil_align(self.sm_count // 2, 8)
@@ -171,6 +187,7 @@ class Indexer(CustomOp):
x: torch.Tensor,
positions: torch.Tensor,
enable_dual_stream: bool,
forward_batch: ForwardBatch,
):
weights = None
if enable_dual_stream:
@@ -228,6 +245,15 @@ class Indexer(CustomOp):
query[..., : self.rope_head_dim] = q_rope
key[..., : self.rope_head_dim] = k_rope
# allgather+rerrange
if forward_batch.nsa_cp_metadata is not None and self.nsa_enable_prefill_cp:
key = cp_all_gather_rerange_output(
key.contiguous(),
self.cp_size,
forward_batch,
torch.cuda.current_stream(),
)
if enable_dual_stream:
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
@@ -469,6 +495,153 @@ class Indexer(CustomOp):
)
return metadata.topk_transform(dummy_logits, self.index_topk)
def _get_topk_ragged_with_cp(
self,
forward_batch: ForwardBatch,
layer_id: int,
q_fp8: torch.Tensor,
weights: torch.Tensor,
metadata: BaseIndexerMetadata,
kv_len: int,
actual_seq_q: int,
cp_index: List[Tuple[int, int, int]] = None,
) -> torch.Tensor:
if TYPE_CHECKING:
assert isinstance(forward_batch.token_to_kv_pool, NSATokenToKVPool)
page_size = forward_batch.token_to_kv_pool.page_size
assert page_size == 64, "only support page size 64"
assert len(weights.shape) == 3
weights = weights.squeeze(-1)
k_fp8_list = []
k_scale_list = []
ks_list = []
ke_offset_list = []
offset = 0
actual_seq_q_list = []
batch_idx_list = []
block_tables = metadata.get_page_table_64()
assert (
forward_batch.seq_lens_cpu is not None
and forward_batch.extend_seq_lens_cpu is not None
)
if cp_index is not None:
# TODO Multi-batch support has accuracy issues
for batch_idx, start_seq_position, end_seq_position in cp_index:
pre_chunk_offset = (
forward_batch.seq_lens_cpu[batch_idx].item()
- forward_batch.extend_seq_lens_cpu[batch_idx]
)
start_seq_position += pre_chunk_offset
end_seq_position += pre_chunk_offset
if offset == 0 and batch_idx != 0:
offset += forward_batch.extend_seq_lens_cpu[batch_idx - 1]
k_fp8 = forward_batch.token_to_kv_pool.get_index_k_continuous(
layer_id,
end_seq_position,
block_tables[batch_idx],
)
k_scale = forward_batch.token_to_kv_pool.get_index_k_scale_continuous(
layer_id,
end_seq_position,
block_tables[batch_idx],
)
extend_seq_len = end_seq_position - start_seq_position
ks = torch.full(
(extend_seq_len,), offset, dtype=torch.int32, device="cuda"
)
k_fp8_list.append(k_fp8)
k_scale_list.append(k_scale)
ks_list.append(ks)
ke_offset = torch.arange(
start_seq_position + 1,
end_seq_position + 1,
dtype=torch.int32,
device="cuda",
)
ke_offset_list.append(ke_offset)
actual_seq_q = torch.tensor(
[extend_seq_len], dtype=torch.int32, device="cuda"
)
actual_seq_q_list.append(actual_seq_q)
batch_idx_list.append(batch_idx)
k_fp8 = torch.cat(k_fp8_list, dim=0).view(torch.float8_e4m3fn)
k_scale = torch.cat(k_scale_list, dim=0).view(torch.float32).squeeze(-1)
kv_fp8 = (k_fp8, k_scale)
ks = torch.cat(ks_list, dim=0)
ke_offset = torch.cat(ke_offset_list, dim=0)
ke = ks + ke_offset
actual_seq_q = torch.cat(actual_seq_q_list, dim=0)
logits = deep_gemm.fp8_mqa_logits(
q_fp8,
kv_fp8,
weights,
ks,
ke,
clean_logits=False,
)
topk_result = metadata.topk_transform(
logits,
self.index_topk,
ks=ks,
cu_seqlens_q=actual_seq_q,
ke_offset=ke_offset,
batch_idx_list=batch_idx_list,
)
else:
kv_len = (
forward_batch.seq_lens_cpu[0].item()
- forward_batch.extend_seq_lens_cpu[0]
+ kv_len
)
k_fp8 = forward_batch.token_to_kv_pool.get_index_k_continuous(
layer_id,
kv_len,
block_tables[0],
)
k_scale = forward_batch.token_to_kv_pool.get_index_k_scale_continuous(
layer_id,
kv_len,
block_tables[0],
)
k_fp8 = k_fp8.view(torch.float8_e4m3fn)
k_scale = k_scale.view(torch.float32).squeeze(-1)
kv_fp8 = (k_fp8, k_scale)
ks = torch.full((actual_seq_q,), offset, dtype=torch.int32, device="cuda")
ke_offset = torch.arange(
(kv_len - actual_seq_q) + 1,
kv_len + 1,
dtype=torch.int32,
device="cuda",
)
ke = ks + ke_offset
logits = deep_gemm.fp8_mqa_logits(
q_fp8,
kv_fp8,
weights,
ks,
ke,
clean_logits=False,
)
actual_seq_q = torch.tensor([actual_seq_q], dtype=torch.int32).to(
device="cuda", non_blocking=True
)
topk_result = metadata.topk_transform(
logits,
self.index_topk,
ks=ks,
cu_seqlens_q=actual_seq_q,
ke_offset=ke_offset,
)
return topk_result
def forward_indexer(
self,
q_fp8: torch.Tensor,
@@ -594,7 +767,7 @@ class Indexer(CustomOp):
skip_logits_computation = max_kv_len <= self.index_topk
# Optimization: fast path when skipping topk computation
if skip_logits_computation:
if skip_logits_computation and (not self.nsa_enable_prefill_cp):
return self._forward_cuda_k_only(
x,
positions,
@@ -607,7 +780,7 @@ class Indexer(CustomOp):
)
query, key, weights = self._get_q_k_bf16(
q_lora, x, positions, enable_dual_stream
q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch
)
if enable_dual_stream:
@@ -660,9 +833,49 @@ class Indexer(CustomOp):
forward_batch, layer_id, q_fp8, weights, metadata
)
else:
topk_result = self._get_topk_ragged(
forward_batch, layer_id, q_fp8, weights, metadata
)
if (
forward_batch.nsa_cp_metadata is not None
and self.nsa_enable_prefill_cp
):
kv_len_prev = forward_batch.nsa_cp_metadata.kv_len_prev
kv_len_next = forward_batch.nsa_cp_metadata.kv_len_next
actual_seq_q_prev = forward_batch.nsa_cp_metadata.actual_seq_q_prev
actual_seq_q_next = forward_batch.nsa_cp_metadata.actual_seq_q_next
# TODO support mutil-batch
# cp_batch_seq_index_prev = forward_batch.nsa_cp_metadata["cp_batch_seq_index_prev"]
# cp_batch_seq_index_next = forward_batch.nsa_cp_metadata["cp_batch_seq_index_next"]
# TODO prev, next, combined into a single call
q_fp8_prev, q_fp8_next = torch.split(
q_fp8, (q_fp8.shape[0] + 1) // 2, dim=0
)
weights_prev, weights_next = torch.split(
weights, (weights.shape[0] + 1) // 2, dim=0
)
topk_result_prev = self._get_topk_ragged_with_cp(
forward_batch,
layer_id,
q_fp8_prev,
weights_prev,
metadata,
kv_len_prev,
actual_seq_q_prev,
)
topk_result_next = self._get_topk_ragged_with_cp(
forward_batch,
layer_id,
q_fp8_next,
weights_next,
metadata,
kv_len_next,
actual_seq_q_next,
)
return torch.cat([topk_result_prev, topk_result_next], dim=0)
else:
topk_result = self._get_topk_ragged(
forward_batch, layer_id, q_fp8, weights, metadata
)
else:
topk_result = self.forward_indexer(
q_fp8.contiguous(),

View File

@@ -1,4 +1,13 @@
# temp NSA debugging environ
from dataclasses import dataclass
from itertools import accumulate
from typing import List
import torch
import torch.nn.functional as F
from sglang.srt.layers.dp_attention import get_attention_tp_group
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import get_bool_env_var
NSA_DUAL_STREAM = get_bool_env_var("SGLANG_NSA_DUAL_STREAM", "true")
@@ -21,3 +30,299 @@ def print_nsa_bool_env_vars():
def compute_nsa_seqlens(original_seq_lens, nsa_index_topk: int):
return original_seq_lens.clamp(max=nsa_index_topk)
def is_nsa_enable_prefill_cp():
return get_global_server_args().enable_nsa_prefill_context_parallel
@dataclass
class NSAContextParallelMetadata:
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
kv_len_prev: int = -1
kv_len_next: int = -1
actual_seq_q_prev: int = -1
actual_seq_q_next: int = -1
total_seq_lens: torch.Tensor = None
def can_cp_split(cur_cp_seq_len: int, cp_size: int, use_nsa: bool, forward_batch):
if (
cur_cp_seq_len != 0
and cp_size > 1
and use_nsa
and forward_batch.forward_mode.is_context_parallel_extend()
and is_nsa_enable_prefill_cp()
):
return True
else:
return False
def cp_split_and_rebuild_data(forward_batch, input_: torch.Tensor):
input_list = list(
torch.split(input_, forward_batch.nsa_cp_metadata.split_list, dim=0)
)
result = torch.cat(
[input_list[i] for i in forward_batch.nsa_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.nsa_cp_metadata.split_list, dim=-1)
)
positions = torch.cat(
[position_id_list[i] for i in forward_batch.nsa_cp_metadata.zigzag_index],
dim=-1,
)
return positions
def enable_prefill_cp(forward_batch, nsa_enable_prefill_cp):
if (
forward_batch.nsa_cp_metadata is not None
and nsa_enable_prefill_cp
and forward_batch.forward_mode.is_context_parallel_extend()
):
return True
else:
return False
def cp_attn_tp_all_gather_reorganazied_into_tensor(
input_: torch.Tensor, total_len, attn_tp_size, forward_batch, stream_op
):
"""
Allgather communication for context_parallel(kv_cache, index_k, hidden_states).
This implementation mainly consists of three parts:
Step 1, padding the input shape to unify the shape for allgather communication (the shape must be the same).
Step 2, allgather communication(async).
Step 3, removing the padding and reassembling the data according to the actual tokens.
"""
# step1
max_len = (total_len + attn_tp_size - 1) // attn_tp_size
pad_size = max_len - input_.shape[0]
if pad_size > 0:
input_ = F.pad(input_, (0, 0, 0, pad_size), mode="constant", value=0)
input_tensor_all = torch.empty(
max_len * attn_tp_size,
input_.shape[1],
device=input_.device,
dtype=input_.dtype,
)
# step2
get_attention_tp_group().cp_all_gather_into_tensor_async(
input_tensor_all, input_, stream_op
)
# step3
outputs_list_max = list(
torch.split(input_tensor_all, forward_batch.nsa_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.nsa_cp_metadata.per_rank_actual_token
)
],
dim=0,
)
return outputs
def cp_all_gather_rerange_output(input_tensor, cp_size, forward_batch, stream):
"""
| +-----------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 |
| +-------------------------+
"""
bs_seq_len, hidden_size = input_tensor.shape
output_tensor = cp_attn_tp_all_gather_reorganazied_into_tensor(
input_tensor,
forward_batch.nsa_cp_metadata.total_seq_lens,
cp_size,
forward_batch,
stream,
)
outputs_list = list(
torch.split(
output_tensor, forward_batch.nsa_cp_metadata.reverse_split_len, dim=0
)
)
output_tensor = torch.cat(
[outputs_list[i] for i in forward_batch.nsa_cp_metadata.cp_reverse_index], dim=0
)
output_tensor = output_tensor.view(-1, hidden_size)
return output_tensor
def calculate_cp_seq_idx(cp_chunks_len, seqs_len):
"""Used to obtain the index of the seq corresponding
to each cp block in the forwardbatch, and the starting
and ending positions of the corresponding seq in the cp block"""
j = 0
tuple_len = [] # Only keep this result list
cumulative = {} # Used to track cumulative values for each index
for i in range(len(cp_chunks_len)):
current_dict = {}
current_tuples = []
c_val = cp_chunks_len[i]
while j < len(seqs_len):
s_val = seqs_len[j]
if s_val == c_val:
idx = j
current_dict[idx] = s_val
# Update cumulative value for this index
cumulative[idx] = cumulative.get(idx, 0) + s_val
j += 1
break
elif s_val > c_val:
idx = j
current_dict[idx] = c_val
# Update cumulative value for this index
cumulative[idx] = cumulative.get(idx, 0) + c_val
seqs_len[j] = s_val - c_val
break
else: # s_val < c_val
idx = j
current_dict[idx] = s_val
# Update cumulative value for this index
cumulative[idx] = cumulative.get(idx, 0) + s_val
c_val -= s_val
j += 1
# Build tuple: (index, historical cumulative, historical+current)
for idx, val in current_dict.items():
# Subtract current value to get historical cumulative
prev_cum = cumulative.get(idx, 0) - val
current_cum = prev_cum + val
current_tuples.append((idx, prev_cum, current_cum))
tuple_len.append(current_tuples)
return tuple_len
def prepare_input_dp_with_cp_dsa(
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 hissenstate needs to be sliced by cp_size*2 and rearranged.
"""
# just support batch = 1
bs_per_cp_group = 1
kv_len_origin = kv_len
# 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
# cp_seq_index = calculate_cp_seq_idx(split_list[:], seqs_len[:])
nsa_cp_metadata = NSAContextParallelMetadata(
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=prefix_sum_list[cp_rank],
kv_len_next=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],
total_seq_lens=kv_len_origin,
)
return nsa_cp_metadata

View File

@@ -145,6 +145,9 @@ class NSAIndexerMetadata(BaseIndexerMetadata):
logits: torch.Tensor,
topk: int,
ks: Optional[torch.Tensor] = None,
cu_seqlens_q: torch.Tensor = None,
ke_offset: torch.Tensor = None,
batch_idx_list: List[int] = None,
) -> torch.Tensor:
from sgl_kernel import (
fast_topk_transform_fused,
@@ -152,25 +155,42 @@ class NSAIndexerMetadata(BaseIndexerMetadata):
fast_topk_v2,
)
if not NSA_FUSE_TOPK:
return fast_topk_v2(
logits, self.get_seqlens_expanded(), topk, row_starts=ks
if cu_seqlens_q is not None:
cu_seqlens_q = cu_seqlens_q.to(torch.int32)
cu_seqlens_q_topk = compute_cu_seqlens(cu_seqlens_q)
cu_topk_indices_offset = torch.repeat_interleave(
cu_seqlens_q_topk[:-1],
cu_seqlens_q,
)
else:
cu_seqlens_q_topk = self.attn_metadata.cu_seqlens_q
cu_topk_indices_offset = self.attn_metadata.topk_indices_offset
if ke_offset is not None:
seq_lens_topk = ke_offset
else:
seq_lens_topk = self.get_seqlens_expanded()
if batch_idx_list is not None:
page_table_size_1 = self.attn_metadata.page_table_1[batch_idx_list]
else:
page_table_size_1 = self.attn_metadata.page_table_1
if not NSA_FUSE_TOPK:
return fast_topk_v2(logits, seq_lens_topk, topk, row_starts=ks)
elif self.topk_transform_method == TopkTransformMethod.PAGED:
# NOTE(dark): if fused, we return a transformed page table directly
return fast_topk_transform_fused(
score=logits,
lengths=self.get_seqlens_expanded(),
page_table_size_1=self.attn_metadata.page_table_1,
cu_seqlens_q=self.attn_metadata.cu_seqlens_q,
lengths=seq_lens_topk,
page_table_size_1=page_table_size_1,
cu_seqlens_q=cu_seqlens_q_topk,
topk=topk,
row_starts=ks,
)
elif self.topk_transform_method == TopkTransformMethod.RAGGED:
return fast_topk_transform_ragged_fused(
score=logits,
lengths=self.get_seqlens_expanded(),
topk_indices_offset=self.attn_metadata.topk_indices_offset,
lengths=seq_lens_topk,
topk_indices_offset=cu_topk_indices_offset,
topk=topk,
row_starts=ks,
)

View File

@@ -0,0 +1,284 @@
# Copyright 2023-2024 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
from functools import partial
from typing import Callable, Optional
import torch
from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp
from sglang.srt.layers.communicator import (
CommunicateContext,
CommunicateSimpleFn,
CommunicateSummableTensorPairFn,
CommunicateWithAllReduceAndLayerNormFn,
LayerCommunicator,
LayerScatterModes,
ScatterMode,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
def nsa_enable_prefill_cp():
# After using cp, the communication mode of this part changes.
# The three parts of prepare_attn, prepare_mlp, and postprocess_layer
# no longer require additional communication for reduce, scatter, etc.
return is_nsa_enable_prefill_cp()
class NSACPLayerCommunicator(LayerCommunicator):
def __init__(
self,
layer_scatter_modes: LayerScatterModes,
input_layernorm: torch.nn.Module,
post_attention_layernorm: torch.nn.Module,
# Reduce scatter requires skipping all-reduce in model code after MoE/MLP, so only enable for models which have that implemented. Remove flag once done for all models that use LayerCommunicator.
allow_reduce_scatter: bool = False,
is_last_layer: bool = False,
qkv_latent_func: Optional[Callable] = None,
):
super().__init__(
layer_scatter_modes,
input_layernorm,
post_attention_layernorm,
allow_reduce_scatter,
is_last_layer,
qkv_latent_func,
)
self._communicate_simple_fn = NSACPCommunicateSimpleFn.get_fn(
input_mode=self.layer_scatter_modes.layer_input_mode,
output_mode=self.layer_scatter_modes.attn_mode,
context=self._context,
)
self._communicate_with_all_reduce_and_layer_norm_fn = (
NSACPCommunicateWithAllReduceAndLayerNormFn.get_fn(
hidden_states_input_mode=self.layer_scatter_modes.attn_mode,
residual_input_mode=self.layer_scatter_modes.layer_input_mode,
hidden_states_output_mode=self.layer_scatter_modes.mlp_mode,
residual_output_mode=self.layer_scatter_modes.middle_residual_mode,
context=self._context,
)
)
self._communicate_summable_tensor_pair_fn = (
NSACPCommunicateSummableTensorPairFn.get_fn(
hidden_states_input_mode=self.layer_scatter_modes.mlp_mode,
residual_input_mode=self.layer_scatter_modes.middle_residual_mode,
output_mode=self.layer_scatter_modes.layer_output_mode,
context=self._context,
)
)
class NSACPCommunicateSimpleFn(CommunicateSimpleFn):
@staticmethod
def get_fn(
input_mode: ScatterMode,
output_mode: ScatterMode,
context: CommunicateContext,
):
if context.is_same_group_size(input_mode, output_mode):
return NSACPCommunicateSimpleFn._trivial
if (input_mode == ScatterMode.SCATTERED) and (
output_mode == ScatterMode.TP_ATTN_FULL
):
return NSACPCommunicateSimpleFn._scattered_to_tp_attn_full
raise NotImplementedError(f"{input_mode=} {output_mode=}")
@staticmethod
def _scattered_to_tp_attn_full(
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
context: CommunicateContext,
) -> torch.Tensor:
if nsa_enable_prefill_cp():
return hidden_states
else:
assert False, "Not implemented"
class NSACPCommunicateWithAllReduceAndLayerNormFn(
CommunicateWithAllReduceAndLayerNormFn
):
"""Besides communication, needs to
1. All reduce in tp_attn_group on hidden_states
2. Apply layer norm
"""
@staticmethod
def get_fn(
hidden_states_input_mode: ScatterMode,
residual_input_mode: ScatterMode,
hidden_states_output_mode: ScatterMode,
residual_output_mode: ScatterMode,
context: CommunicateContext,
):
if (
context.is_same_group_size(
hidden_states_input_mode, hidden_states_output_mode
)
and context.is_same_group_size(residual_input_mode, residual_output_mode)
and context.attn_tp_size == 1
):
return NSACPCommunicateWithAllReduceAndLayerNormFn._simple
if (
(hidden_states_input_mode == ScatterMode.TP_ATTN_FULL)
and (
residual_input_mode in [ScatterMode.SCATTERED, ScatterMode.TP_ATTN_FULL]
)
and (hidden_states_output_mode == ScatterMode.FULL)
and (residual_output_mode == ScatterMode.TP_ATTN_FULL)
):
return partial(
NSACPCommunicateWithAllReduceAndLayerNormFn._gather_hidden_states_and_residual,
residual_input_mode=residual_input_mode,
)
if (
(hidden_states_input_mode == ScatterMode.TP_ATTN_FULL)
and (
residual_input_mode in [ScatterMode.SCATTERED, ScatterMode.TP_ATTN_FULL]
)
and (hidden_states_output_mode == ScatterMode.SCATTERED)
and (residual_output_mode == ScatterMode.SCATTERED)
):
return partial(
NSACPCommunicateWithAllReduceAndLayerNormFn._scatter_hidden_states_and_residual,
residual_input_mode=residual_input_mode,
)
raise NotImplementedError(
f"{hidden_states_input_mode=} {residual_input_mode=} {hidden_states_output_mode=} {residual_output_mode=}"
)
@staticmethod
def _gather_hidden_states_and_residual(
hidden_states: torch.Tensor,
residual: torch.Tensor,
forward_batch: ForwardBatch,
layernorm: torch.nn.Module,
context: CommunicateContext,
*,
residual_input_mode,
):
if context.attn_dp_size != 1:
if nsa_enable_prefill_cp():
hidden_states += residual
if hidden_states.shape[0] != 0:
hidden_states = layernorm(hidden_states)
return hidden_states, residual
else:
assert False, "not yet handled"
@staticmethod
def _scatter_hidden_states_and_residual(
hidden_states: torch.Tensor,
residual: torch.Tensor,
forward_batch: ForwardBatch,
layernorm: torch.nn.Module,
context: CommunicateContext,
*,
residual_input_mode,
):
if nsa_enable_prefill_cp():
if hidden_states.shape[0] != 0:
hidden_states, residual = layernorm(hidden_states, residual)
return hidden_states, residual
else:
assert False, "not yet handled"
class NSACPCommunicateSummableTensorPairFn(CommunicateSummableTensorPairFn):
"""It is allowed to make (hidden_states, residual) := (hidden_states + residual, None) if needed."""
@staticmethod
def get_fn(
hidden_states_input_mode: ScatterMode,
residual_input_mode: ScatterMode,
output_mode: ScatterMode,
context: CommunicateContext,
):
if context.is_same_group_size(
hidden_states_input_mode, output_mode
) and context.is_same_group_size(residual_input_mode, output_mode):
return NSACPCommunicateSummableTensorPairFn._trivial
if (
(hidden_states_input_mode == ScatterMode.FULL)
and (residual_input_mode == ScatterMode.TP_ATTN_FULL)
and (output_mode == ScatterMode.TP_ATTN_FULL)
):
return NSACPCommunicateSummableTensorPairFn._scatter_hidden_states
if (
(hidden_states_input_mode == ScatterMode.SCATTERED)
and (residual_input_mode == ScatterMode.SCATTERED)
and (output_mode == ScatterMode.TP_ATTN_FULL)
):
return NSACPCommunicateSummableTensorPairFn._gather
if (
(hidden_states_input_mode == ScatterMode.TP_ATTN_FULL)
and (residual_input_mode == ScatterMode.TP_ATTN_FULL)
and (output_mode == ScatterMode.SCATTERED)
):
return NSACPCommunicateSummableTensorPairFn._scatter
raise NotImplementedError(
f"{hidden_states_input_mode=} {residual_input_mode=} {output_mode=}"
)
@staticmethod
def _scatter_hidden_states(
hidden_states: torch.Tensor,
residual: torch.Tensor,
forward_batch: ForwardBatch,
context: CommunicateContext,
allow_reduce_scatter: bool = False,
):
if nsa_enable_prefill_cp():
return hidden_states, residual
else:
assert False, "not yet handled"
@staticmethod
def _gather(
hidden_states: torch.Tensor,
residual: torch.Tensor,
forward_batch: ForwardBatch,
context: CommunicateContext,
**kwargs,
):
hidden_states += residual
residual = None
if nsa_enable_prefill_cp():
return hidden_states, residual
else:
assert False, "not yet handled"
@staticmethod
def _scatter(
hidden_states: torch.Tensor,
residual: torch.Tensor,
forward_batch: ForwardBatch,
context: CommunicateContext,
):
if nsa_enable_prefill_cp():
return hidden_states, residual
else:
assert False, "not yet handled"

View File

@@ -289,6 +289,10 @@ def initialize_dp_attention(
_LOCAL_ATTN_DP_SIZE = 1
tp_group = get_tp_group()
# Trick to solve circular references
from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp
use_pynccl = True if is_nsa_enable_prefill_cp() else SYNC_TOKEN_IDS_ACROSS_TP
_ATTN_TP_GROUP = GroupCoordinator(
[
list(range(head, head + _ATTN_TP_SIZE))
@@ -296,7 +300,7 @@ def initialize_dp_attention(
],
tp_group.local_rank,
torch.distributed.get_backend(tp_group.device_group),
use_pynccl=SYNC_TOKEN_IDS_ACROSS_TP,
use_pynccl=use_pynccl,
use_pymscclpp=False,
use_custom_allreduce=False,
use_torch_symm_mem_all_reduce=False,

View File

@@ -24,6 +24,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Set, Union
import torch
from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.mem_cache.allocator import SWATokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
@@ -363,6 +364,7 @@ class PrefillAdder:
self.priority_scheduling_preemption_threshold = (
priority_scheduling_preemption_threshold
)
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
def _get_running_request_total_token_offset(self, req: Req) -> int:
return (
@@ -565,6 +567,11 @@ class PrefillAdder:
def add_one_req(
self, req: Req, has_chunked_req: bool, truncation_align_size: Optional[int]
):
# 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_enable_prefill_cp and len(self.can_run_list) >= 1:
return AddReqResult.OTHER
if req.sampling_params.ignore_eos and getattr(self.tree_cache, "disable", True):
return self.add_one_req_ignore_eos(req, has_chunked_req)

View File

@@ -39,6 +39,7 @@ from sglang.srt.distributed.parallel_state import (
graph_capture,
set_pdmux_status,
)
from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp
from sglang.srt.layers.dp_attention import (
DpPaddingMode,
get_attention_tp_rank,
@@ -258,6 +259,7 @@ class CudaGraphRunner:
self.attn_tp_size = get_attention_tp_size()
self.attn_tp_rank = get_attention_tp_rank()
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
self.deepep_adapter = DeepEPCudaGraphRunnerAdapter()
@@ -836,7 +838,7 @@ class CudaGraphRunner:
self.global_num_tokens_for_logprob_gpu.fill_(bs * self.num_tokens_per_bs)
if enable_num_token_non_padded(self.model_runner.server_args):
num_token_non_padded = forward_batch.num_token_non_padded
if self.require_gathered_buffer:
if self.require_gathered_buffer and not self.nsa_enable_prefill_cp:
tokens_per_rank = bs // self.attn_tp_size * self.num_tokens_per_bs
num_local_token_non_padded = torch.clamp(
num_token_non_padded - tokens_per_rank * self.attn_tp_rank,

View File

@@ -42,6 +42,7 @@ from sglang.srt.distributed.parallel_state import (
get_moe_expert_parallel_world_size,
get_tensor_model_parallel_world_size,
)
from sglang.srt.layers.attention.nsa.utils import NSAContextParallelMetadata
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
from sglang.srt.layers.dp_attention import (
DpPaddingMode,
@@ -102,6 +103,17 @@ class ForwardMode(IntEnum):
or self == ForwardMode.SPLIT_PREFILL
)
def is_context_parallel_extend(self, include_draft_extend_v2: bool = False):
return (
self == ForwardMode.EXTEND
or self == ForwardMode.MIXED
or (
self == ForwardMode.DRAFT_EXTEND_V2
if include_draft_extend_v2
else False
)
)
def is_decode(self):
return self == ForwardMode.DECODE
@@ -339,6 +351,9 @@ class ForwardBatch:
# For matryoshka embeddings
dimensions: Optional[list[int]] = None
# Record the split metadata of the sequence number of NSA context parallels.
nsa_cp_metadata: Optional[NSAContextParallelMetadata] = None
@classmethod
def init_new(
cls,

View File

@@ -20,9 +20,22 @@ import torch
from torch import nn
from transformers import PretrainedConfig
from sglang.srt.configs.model_config import is_deepseek_nsa
from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
from sglang.srt.layers.attention.nsa.utils import (
can_cp_split,
cp_all_gather_rerange_output,
cp_split_and_rebuild_data,
enable_prefill_cp,
is_nsa_enable_prefill_cp,
prepare_input_dp_with_cp_dsa,
)
from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
is_dp_attention_enabled,
)
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.quantization import Fp8Config
@@ -55,7 +68,6 @@ class DeepseekModelNextN(nn.Module):
prefix: str = "",
) -> None:
super().__init__()
if enable_nextn_moe_bf16_cast_to_fp8(quant_config):
# refer to real DeepSeek V3 quant config
moe_quant_config = Fp8Config(
@@ -106,6 +118,11 @@ class DeepseekModelNextN(nn.Module):
self.shared_head = nn.Module()
self.shared_head.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
if self.nsa_enable_prefill_cp:
self.cp_size = get_attention_tp_size()
else:
self.cp_size = None
def forward(
self,
@@ -138,10 +155,16 @@ class DeepseekModelNextN(nn.Module):
)
)
if enable_prefill_cp(forward_batch, self.nsa_enable_prefill_cp):
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
residual = None
with get_global_expert_distribution_recorder().disable_this_region():
hidden_states, residual = self.decoder(
positions, hidden_states, forward_batch, residual, zero_allocator
positions,
hidden_states,
forward_batch,
residual,
zero_allocator,
)
if not forward_batch.forward_mode.is_idle():
@@ -150,6 +173,15 @@ class DeepseekModelNextN(nn.Module):
else:
hidden_states = self.shared_head.norm(hidden_states)
if enable_prefill_cp(forward_batch, self.nsa_enable_prefill_cp):
# allgather + rerrange
hidden_states = cp_all_gather_rerange_output(
hidden_states,
self.cp_size,
forward_batch,
torch.cuda.current_stream(),
)
return hidden_states
@@ -168,6 +200,14 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
# if not set, model load will be broken in DeepseekV3ForCausalLM load_weights()
self.pp_group = get_pp_group()
self.determine_num_fused_shared_experts("DeepseekV3ForCausalLMNextN")
self.use_nsa = is_deepseek_nsa(config)
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
if self.nsa_enable_prefill_cp:
self.cp_rank = get_attention_tp_rank()
self.cp_size = get_attention_tp_size()
else:
self.cp_rank = None
self.cp_size = None
self.model = DeepseekModelNextN(
config, quant_config, prefix=add_prefix("model", prefix)
@@ -188,6 +228,16 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
positions: torch.Tensor,
forward_batch: ForwardBatch,
) -> torch.Tensor:
# TODO current just support prefill batch=1 and len(input_ids) > self.cp_size * 2
if self.nsa_enable_prefill_cp:
cur_cp_seq_len = len(input_ids) // (self.cp_size * 2)
if can_cp_split(cur_cp_seq_len, self.cp_size, self.use_nsa, forward_batch):
forward_batch.nsa_cp_metadata = prepare_input_dp_with_cp_dsa(
torch.tensor(len(input_ids)),
self.cp_rank,
self.cp_size,
forward_batch.seq_lens_cpu.tolist(),
)
hidden_states = self.model(input_ids, positions, forward_batch)
return self.logits_processor(
input_ids, hidden_states, self.lm_head, forward_batch

View File

@@ -54,6 +54,15 @@ from sglang.srt.layers.attention.npu_ops.mla_preprocess import (
is_mla_preprocess_enabled,
)
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
from sglang.srt.layers.attention.nsa.utils import (
can_cp_split,
cp_all_gather_rerange_output,
cp_split_and_rebuild_data,
cp_split_and_rebuild_position,
enable_prefill_cp,
is_nsa_enable_prefill_cp,
prepare_input_dp_with_cp_dsa,
)
from sglang.srt.layers.attention.utils import concat_and_cast_mha_k_triton
from sglang.srt.layers.communicator import (
LayerCommunicator,
@@ -61,6 +70,7 @@ from sglang.srt.layers.communicator import (
enable_moe_dense_fully_dp,
get_attn_tp_context,
)
from sglang.srt.layers.communicator_nsa_cp import NSACPLayerCommunicator
from sglang.srt.layers.dp_attention import (
get_attention_tp_rank,
get_attention_tp_size,
@@ -412,7 +422,9 @@ def handle_attention_nsa(attn, forward_batch):
if forward_batch.forward_mode.is_decode_or_idle():
return AttnForwardMethod.MLA
if forward_batch.forward_mode.is_extend_without_speculative():
if forward_batch.forward_mode.is_extend_without_speculative() and (
not is_nsa_enable_prefill_cp()
):
assert forward_batch.seq_lens_cpu is not None
max_kv_len = forward_batch.seq_lens_cpu.max().item()
@@ -550,8 +562,14 @@ class MoEGate(nn.Module):
self.e_score_correction_bias = None
if _is_cpu and _is_cpu_amx_available:
self.quant_method = PackWeightMethod(weight_names=["weight"])
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
def forward(self, hidden_states, gemm_output_zero_allocator: BumpAllocator = None):
def forward(
self,
hidden_states,
gemm_output_zero_allocator: BumpAllocator = None,
forward_batch: ForwardBatch = None,
):
if use_intel_amx_backend(self):
return torch.ops.sgl_kernel.weight_packed_linear(
hidden_states,
@@ -563,24 +581,30 @@ class MoEGate(nn.Module):
if get_global_server_args().enable_deterministic_inference:
return F.linear(hidden_states, self.weight, None)
# NOTE: For some unknown reason, router_gemm seems degrade accept length.
if (
_is_cuda
and hidden_states.shape[0] <= 16
and hidden_states.shape[1] == 7168
and (self.weight.shape[0] == 256 or self.weight.shape[0] == 384)
and _device_sm >= 90
if forward_batch is not None and enable_prefill_cp(
forward_batch, self.nsa_enable_prefill_cp
):
# router gemm output float32
logits = dsv3_router_gemm(
hidden_states, self.weight, out_dtype=torch.float32
)
elif _use_aiter_gfx95 and hidden_states.shape[0] <= 256:
logits = aiter_dsv3_router_gemm(
hidden_states, self.weight, gemm_output_zero_allocator
)
else:
logits = F.linear(hidden_states, self.weight, None)
else:
# NOTE: For some unknown reason, router_gemm seems degrade accept length.
if (
_is_cuda
and hidden_states.shape[0] <= 16
and hidden_states.shape[1] == 7168
and (self.weight.shape[0] == 256 or self.weight.shape[0] == 384)
and _device_sm >= 90
):
# router gemm output float32
logits = dsv3_router_gemm(
hidden_states, self.weight, out_dtype=torch.float32
)
elif _use_aiter_gfx95 and hidden_states.shape[0] <= 256:
logits = aiter_dsv3_router_gemm(
hidden_states, self.weight, gemm_output_zero_allocator
)
else:
logits = F.linear(hidden_states, self.weight, None)
return logits
@@ -938,7 +962,7 @@ class DeepseekV2MoE(nn.Module):
shared_output = None
if hidden_states.shape[0] > 0:
# router_logits: (num_tokens, n_experts)
router_logits = self.gate(hidden_states)
router_logits = self.gate(hidden_states, forward_batch=forward_batch)
if not self._fuse_shared_experts_inside_sbo:
shared_output = self._forward_shared_experts(hidden_states)
topk_output = self.topk(
@@ -1124,7 +1148,15 @@ class DeepseekV2AttentionMLA(nn.Module):
self.kv_lora_rank = kv_lora_rank
attn_tp_rank = get_attention_tp_rank()
attn_tp_size = get_attention_tp_size()
self.use_nsa = is_deepseek_nsa(config)
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
if self.nsa_enable_prefill_cp:
assert self.use_nsa, "CP currently only supports deepseek v3.2 model"
# cp reuse the attn_tp comm group but need to duplicate the weights
if self.nsa_enable_prefill_cp and self.use_nsa:
attn_tp_rank = 0
attn_tp_size = 1
self.cp_size = get_attention_tp_size()
self.num_heads = num_heads
assert num_heads % attn_tp_size == 0
self.num_local_heads = num_heads // attn_tp_size
@@ -1174,7 +1206,6 @@ class DeepseekV2AttentionMLA(nn.Module):
prefix=add_prefix("kv_a_proj_with_mqa", prefix),
)
self.use_nsa = is_deepseek_nsa(config)
if self.use_nsa:
self.indexer = Indexer(
hidden_size=hidden_size,
@@ -1652,6 +1683,20 @@ class DeepseekV2AttentionMLA(nn.Module):
and forward_batch.attn_backend.data_type == torch.float8_e4m3fn
)
def rebuild_cp_kv_cache(self, latent_cache, forward_batch, k_nope, k_pe):
# support allgather+rerrange
latent_cache[..., : self.kv_lora_rank] = k_nope.squeeze(1)
latent_cache[..., self.kv_lora_rank :] = k_pe.squeeze(1)
latent_cache_output = cp_all_gather_rerange_output(
latent_cache.contiguous(),
self.cp_size,
forward_batch,
torch.cuda.current_stream(),
)
k_nope = latent_cache_output[..., : self.kv_lora_rank].unsqueeze(1)
k_pe = latent_cache_output[..., self.kv_lora_rank :].unsqueeze(1)
return k_nope, k_pe
def forward_absorb_prepare(
self,
positions: torch.Tensor,
@@ -1787,6 +1832,8 @@ class DeepseekV2AttentionMLA(nn.Module):
q_nope_out = q_nope_out.transpose(0, 1)
if enable_prefill_cp(forward_batch, self.nsa_enable_prefill_cp):
positions = cp_split_and_rebuild_position(forward_batch, positions)
if (
self.rotary_emb is not None
and (not self._fuse_rope_for_trtllm_mla(forward_batch))
@@ -1794,6 +1841,11 @@ class DeepseekV2AttentionMLA(nn.Module):
):
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
if enable_prefill_cp(forward_batch, self.nsa_enable_prefill_cp):
# support allgather+rerrange
k_nope, k_pe = self.rebuild_cp_kv_cache(
latent_cache, forward_batch, k_nope, k_pe
)
topk_indices = None
if q_lora is not None:
topk_indices = self.indexer(
@@ -2696,6 +2748,7 @@ class DeepseekV2DecoderLayer(nn.Module):
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
get_global_server_args().speculative_algorithm
)
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
self.layer_id = layer_id
self.is_nextn = is_nextn
self.self_attn = DeepseekV2AttentionMLA(
@@ -2758,16 +2811,28 @@ class DeepseekV2DecoderLayer(nn.Module):
config.hidden_size, eps=config.rms_norm_eps
)
self.layer_communicator = LayerCommunicator(
layer_scatter_modes=self.layer_scatter_modes,
input_layernorm=self.input_layernorm,
post_attention_layernorm=self.post_attention_layernorm,
allow_reduce_scatter=True,
is_last_layer=(
is_nextn or (self.layer_id == self.config.num_hidden_layers - 1)
),
qkv_latent_func=self.self_attn.prepare_qkv_latent,
)
if self.nsa_enable_prefill_cp:
self.layer_communicator = NSACPLayerCommunicator(
layer_scatter_modes=self.layer_scatter_modes,
input_layernorm=self.input_layernorm,
post_attention_layernorm=self.post_attention_layernorm,
allow_reduce_scatter=True,
is_last_layer=(
is_nextn or (self.layer_id == self.config.num_hidden_layers - 1)
),
qkv_latent_func=self.self_attn.prepare_qkv_latent,
)
else:
self.layer_communicator = LayerCommunicator(
layer_scatter_modes=self.layer_scatter_modes,
input_layernorm=self.input_layernorm,
post_attention_layernorm=self.post_attention_layernorm,
allow_reduce_scatter=True,
is_last_layer=(
is_nextn or (self.layer_id == self.config.num_hidden_layers - 1)
),
qkv_latent_func=self.self_attn.prepare_qkv_latent,
)
def _is_layer_sparse(self, layer_id: int, is_nextn: bool) -> bool:
return is_nextn or (
@@ -2852,7 +2917,7 @@ class DeepseekV2DecoderLayer(nn.Module):
gemm_output_zero_allocator,
)
if should_allreduce_fusion:
if not self.nsa_enable_prefill_cp and should_allreduce_fusion:
hidden_states._sglang_needs_allreduce_fusion = True
if not should_allreduce_fusion:
@@ -2947,6 +3012,11 @@ class DeepseekV2Model(nn.Module):
self.vocab_size = config.vocab_size
self.first_k_dense_replace = config.first_k_dense_replace
self.pp_group = get_pp_group()
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
if self.nsa_enable_prefill_cp:
self.cp_size = get_attention_tp_size()
else:
self.cp_size = None
if self.pp_group.is_first_rank:
self.embed_tokens = VocabParallelEmbedding(
@@ -3082,6 +3152,9 @@ class DeepseekV2Model(nn.Module):
hidden_states = pp_proxy_tensors["hidden_states"]
residual = pp_proxy_tensors["residual"]
if enable_prefill_cp(forward_batch, self.nsa_enable_prefill_cp):
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
normal_start_layer = self.start_layer
normal_end_layer = self.end_layer
if forward_batch.can_run_tbo:
@@ -3134,6 +3207,15 @@ class DeepseekV2Model(nn.Module):
hidden_states = self.norm(hidden_states)
else:
hidden_states, _ = self.norm(hidden_states, residual)
if enable_prefill_cp(forward_batch, self.nsa_enable_prefill_cp):
# allgather + rerrange
hidden_states = cp_all_gather_rerange_output(
hidden_states,
self.cp_size,
forward_batch,
torch.cuda.current_stream(),
)
if len(aux_hidden_states) == 0:
return hidden_states
return hidden_states, aux_hidden_states
@@ -3167,6 +3249,7 @@ class DeepseekV2ForCausalLM(nn.Module):
self.tp_size = get_tensor_model_parallel_world_size()
self.quant_config = quant_config
self.determine_num_fused_shared_experts()
self.use_nsa = is_deepseek_nsa(config)
self.model = DeepseekV2Model(
config, quant_config, prefix=add_prefix("model", prefix)
)
@@ -3187,6 +3270,12 @@ class DeepseekV2ForCausalLM(nn.Module):
}
)
self.capture_aux_hidden_states = False
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
if self.nsa_enable_prefill_cp:
self.cp_rank = get_attention_tp_rank()
self.cp_size = get_attention_tp_size()
else:
self.cp_rank = self.cp_size = None
q_lora_rank = config.q_lora_rank if hasattr(config, "q_lora_rank") else None
get_attn_tp_context().init_context(q_lora_rank, is_deepseek_nsa(config))
@@ -3249,6 +3338,19 @@ class DeepseekV2ForCausalLM(nn.Module):
input_embeds: torch.Tensor = None,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
) -> torch.Tensor:
if self.nsa_enable_prefill_cp:
# 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 = len(input_ids) // (self.cp_size * 2)
if can_cp_split(cur_cp_seq_len, self.cp_size, self.use_nsa, forward_batch):
forward_batch.nsa_cp_metadata = prepare_input_dp_with_cp_dsa(
torch.tensor(len(input_ids)),
self.cp_rank,
self.cp_size,
forward_batch.seq_lens_cpu.tolist(),
)
with get_attn_tp_context().maybe_input_scattered(forward_batch):
hidden_states = self.model(
input_ids, positions, forward_batch, input_embeds, pp_proxy_tensors

View File

@@ -526,6 +526,8 @@ class ServerArgs:
enable_deterministic_inference: bool = False
rl_on_policy_target: Optional[str] = None
enable_attn_tp_input_scattered: bool = False
# Context parallelism used in the long sequence prefill phase of DeepSeek v3.2
enable_nsa_prefill_context_parallel: bool = False
# Dynamic batch tokenizer
enable_dynamic_batch_tokenizer: bool = False
@@ -1105,8 +1107,22 @@ class ServerArgs:
if not is_npu():
self.enable_dp_attention = True
self.dp_size = self.tp_size
logger.warning("DP attention is enabled for DeepSeek NSA.")
if self.enable_nsa_prefill_context_parallel:
# TODO Supports moe_dense_tp_size != 1, kv cache dtype = "fp8",moe_a2a_backend non-deepep and cross-machine operation .
self.moe_dense_tp_size = 1
self.moe_a2a_backend = "deepep"
self.ep_size = self.tp_size
self.kv_cache_dtype = "bf16"
assert (
self.tp_size == 8
), "Current multi-machine CP support suffers from precision issues. So context parallel only support Single machine(tp_size == 8)"
logger.warning(
f"Enable Context Parallel opt for deeeseekv3.2-DSA, Setting dp_size == {self.dp_size} and moe_dense_tp_size == {self.moe_dense_tp_size}, ep_size == {self.ep_size}, tp_size == {self.tp_size}, kv_cache_dtype == {self.kv_cache_dtype}, moe_a2a_backend {self.moe_a2a_backend} "
)
else:
self.dp_size = self.tp_size
self.page_size = 64
logger.warning("Setting page size to 64 for DeepSeek NSA.")
@@ -3493,6 +3509,11 @@ class ServerArgs:
action="store_true",
help="Allow input of attention to be scattered when only using tensor parallelism, to reduce the computational load of operations such as qkv latent.",
)
parser.add_argument(
"--enable-nsa-prefill-context-parallel",
action="store_true",
help="Enable context parallelism used in the long sequence prefill phase of DeepSeek v3.2.",
)
# Dynamic batch tokenizer
parser.add_argument(