[DeepSeek 3.2] Support and optimize pipeline parallelis when context pipeline enabled (#16380)

Co-authored-by: ybyang <ybyang7@iflytek.com>
This commit is contained in:
Yongfei Xu
2026-01-09 11:01:49 +08:00
committed by GitHub
co-authored by ybyang
parent 7460240737
commit 05dfef92a1
2 changed files with 72 additions and 36 deletions
@@ -1,5 +1,6 @@
from __future__ import annotations from __future__ import annotations
import contextlib
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
@@ -23,6 +24,7 @@ if is_npu():
import torch_npu import torch_npu
from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream
from sglang.srt.distributed.parallel_state import get_pp_group
from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.attention.nsa.utils import ( from sglang.srt.layers.attention.nsa.utils import (
cp_all_gather_rerange_output, cp_all_gather_rerange_output,
@@ -146,6 +148,10 @@ class Indexer(MultiPlatformOp):
if is_cuda(): if is_cuda():
self.sm_count = deep_gemm.get_num_sms() self.sm_count = deep_gemm.get_num_sms()
self.half_device_sm_count = ceil_align(self.sm_count // 2, 8) self.half_device_sm_count = ceil_align(self.sm_count // 2, 8)
pp_size = get_global_server_args().pp_size
self.logits_with_pp_recv = pp_size > 1 and not get_pp_group().is_last_rank
else:
self.logits_with_pp_recv = False
self.wq_b = ReplicatedLinear( self.wq_b = ReplicatedLinear(
self.q_lora_rank, self.q_lora_rank,
@@ -184,6 +190,21 @@ class Indexer(MultiPlatformOp):
self.scale_fmt = scale_fmt self.scale_fmt = scale_fmt
self.softmax_scale = self.head_dim**-0.5 self.softmax_scale = self.head_dim**-0.5
@contextlib.contextmanager
def _with_real_sm_count(self):
# When pipeline parallelism is enabled, each PP rank initiates a recv operation after the _pp_launch_batch
# request to receive the PP proxy tensor or output from the previous stage, occupying one SM resource.
# Model execution runs in parallel with the recv operation, so the SMs available to the indexer must be reduced
# by 1. Currently, the last rank starts the send result + recv request only after waiting for execution results.
if self.logits_with_pp_recv:
pp_recv_sm_count = 1
with deep_gemm_wrapper.configure_deep_gemm_num_sms(
self.sm_count - pp_recv_sm_count
):
yield
else:
yield
@torch.compile(dynamic=True) @torch.compile(dynamic=True)
def _get_logits_head_gate(self, x: torch.Tensor, q_scale: torch.Tensor): def _get_logits_head_gate(self, x: torch.Tensor, q_scale: torch.Tensor):
weights, _ = self.weights_proj(x.float()) weights, _ = self.weights_proj(x.float())
@@ -333,7 +354,6 @@ class Indexer(MultiPlatformOp):
) )
assert len(weights.shape) == 3 assert len(weights.shape) == 3
weights = weights.squeeze(2) weights = weights.squeeze(2)
logits = deep_gemm.fp8_paged_mqa_logits( logits = deep_gemm.fp8_paged_mqa_logits(
q_fp8, q_fp8,
kv_cache_fp8, kv_cache_fp8,
@@ -432,14 +452,15 @@ class Indexer(MultiPlatformOp):
if not need_chunk: if not need_chunk:
assert q_fp8[:q_offset].shape[0] != 0 assert q_fp8[:q_offset].shape[0] != 0
logits = deep_gemm.fp8_mqa_logits( with self._with_real_sm_count():
q_fp8[:q_offset], logits = deep_gemm.fp8_mqa_logits(
kv_fp8, q_fp8[:q_offset],
weights[:q_offset], kv_fp8,
ks, weights[:q_offset],
ke, ks,
clean_logits=False, ke,
) clean_logits=False,
)
assert logits.shape[0] == len(seq_lens_expanded) assert logits.shape[0] == len(seq_lens_expanded)
assert logits.shape[1] == k_offset assert logits.shape[1] == k_offset
@@ -468,14 +489,15 @@ class Indexer(MultiPlatformOp):
while start < q_offset: while start < q_offset:
end = min(start + max_rows, q_offset) end = min(start + max_rows, q_offset)
logits_chunk = deep_gemm.fp8_mqa_logits( with self._with_real_sm_count():
q_fp8[start:end], logits_chunk = deep_gemm.fp8_mqa_logits(
kv_fp8, q_fp8[start:end],
weights[start:end], kv_fp8,
ks[start:end], weights[start:end],
ke[start:end], ks[start:end],
clean_logits=False, ke[start:end],
) clean_logits=False,
)
lengths_chunk = seq_lens_expanded[start:end] lengths_chunk = seq_lens_expanded[start:end]
@@ -630,14 +652,15 @@ class Indexer(MultiPlatformOp):
ke_offset = torch.cat(ke_offset_list, dim=0) ke_offset = torch.cat(ke_offset_list, dim=0)
ke = ks + ke_offset ke = ks + ke_offset
actual_seq_q = torch.cat(actual_seq_q_list, dim=0) actual_seq_q = torch.cat(actual_seq_q_list, dim=0)
logits = deep_gemm.fp8_mqa_logits( with self._with_real_sm_count():
q_fp8, logits = deep_gemm.fp8_mqa_logits(
kv_fp8, q_fp8,
weights, kv_fp8,
ks, weights,
ke, ks,
clean_logits=False, ke,
) clean_logits=False,
)
topk_result = metadata.topk_transform( topk_result = metadata.topk_transform(
logits, logits,
self.index_topk, self.index_topk,
@@ -675,14 +698,15 @@ class Indexer(MultiPlatformOp):
) )
ke = ks + ke_offset ke = ks + ke_offset
logits = deep_gemm.fp8_mqa_logits( with self._with_real_sm_count():
q_fp8, logits = deep_gemm.fp8_mqa_logits(
kv_fp8, q_fp8,
weights, kv_fp8,
ks, weights,
ke, ks,
clean_logits=False, ke,
) clean_logits=False,
)
actual_seq_q = torch.tensor([actual_seq_q], dtype=torch.int32).to( actual_seq_q = torch.tensor([actual_seq_q], dtype=torch.int32).to(
device="cuda", non_blocking=True device="cuda", non_blocking=True
) )
@@ -504,6 +504,10 @@ class SchedulerPPMixin:
def init_pp_loop_state(self: Scheduler): def init_pp_loop_state(self: Scheduler):
self.pp_loop_size: int = self.pp_size + self.server_args.pp_async_batch_depth self.pp_loop_size: int = self.pp_size + self.server_args.pp_async_batch_depth
# In CP mode, attention weights are duplicated, eliminating the need for the attention TP all-gather operation.
self.require_attn_tp_allgather = (
not self.server_args.enable_nsa_prefill_context_parallel
)
self.mbs = [None] * self.pp_loop_size self.mbs = [None] * self.pp_loop_size
self.last_mbs = [None] * self.pp_loop_size self.last_mbs = [None] * self.pp_loop_size
self.running_mbs = [ self.running_mbs = [
@@ -906,7 +910,9 @@ class SchedulerPPMixin:
p2p_work.extend( p2p_work.extend(
self.pp_group.send_tensor_dict( self.pp_group.send_tensor_dict(
tensor_dict=tensor_dict, tensor_dict=tensor_dict,
all_gather_group=self.attn_tp_group, all_gather_group=(
self.attn_tp_group if self.require_attn_tp_allgather else None
),
async_send=async_send, async_send=async_send,
) )
) )
@@ -916,7 +922,11 @@ class SchedulerPPMixin:
pp_proxy_tensors = None pp_proxy_tensors = None
if not self.pp_group.is_first_rank: if not self.pp_group.is_first_rank:
pp_proxy_tensors = PPProxyTensors( pp_proxy_tensors = PPProxyTensors(
self.pp_group.recv_tensor_dict(all_gather_group=self.attn_tp_group) self.pp_group.recv_tensor_dict(
all_gather_group=(
self.attn_tp_group if self.require_attn_tp_allgather else None
)
)
) )
return pp_proxy_tensors return pp_proxy_tensors
@@ -924,7 +934,9 @@ class SchedulerPPMixin:
self: Scheduler, self: Scheduler,
) -> Dict[str, torch.Tensor]: ) -> Dict[str, torch.Tensor]:
res = self.pp_group.recv_tensor_dict( res = self.pp_group.recv_tensor_dict(
all_gather_group=self.attn_tp_group, all_gather_group=(
self.attn_tp_group if self.require_attn_tp_allgather else None
),
) )
return res return res