feat: add Pipeline Parallelism (PP) and PD support for DeepSeek-V4 (#24704)

Co-authored-by: Shangming Cai <csmthu@gmail.com>
Co-authored-by: xuyongfei <xuyongfei.xyf@antgroup.com>
This commit is contained in:
ybyang
2026-05-15 22:54:32 -07:00
committed by GitHub
co-authored by Shangming Cai xuyongfei
parent bda01d2435
commit 162540e0a8
9 changed files with 340 additions and 102 deletions
@@ -766,6 +766,9 @@ class ModelConfig:
self.spec_hidden_size = (
self.hidden_size * hc_mult if hc_mult > 1 else self.hidden_size
)
# mHC-flattened hidden size; None when not running an mHC model
# (e.g. non-DeepSeek-V4 configs without ``hc_mult``).
self.hc_hidden_size = self.spec_hidden_size if hc_mult > 1 else None
self.num_hidden_layers = self.hf_text_config.num_hidden_layers
self.num_attention_layers = self.num_hidden_layers
if "LongcatFlashForCausalLM" in self.hf_config.architectures:
+12 -2
View File
@@ -47,11 +47,21 @@ class KVArgs:
kv_head_num: int
total_kv_head_num: int
page_size: int
# for system dp
system_dp_rank: int
# for pp prefill
pp_rank: int
prefill_start_layer: int
# for system dp
system_dp_rank: int
# Absolute end layer (exclusive) for this prefill PP stage. Needed to
# reconstruct PP sub-ranges when kv_data_ptrs does not use a flat
# layer-indexed layout (e.g. DeepSeek V4's buffer-type-organized flat
# list).
prefill_end_layer: int
# For DeepSeek V4 (and other compressed-MLA) memory pools only.
# Full-model compression ratio per layer (entries are 0/4/128). Used by
# the connection layer to slice the buffer-type-organized flat list in a
# PP-aware manner.
mla_compression_ratios: Optional[List[int]]
# Only used of npu, for kv buf groups
kv_buf_groups: int
# Only used of npu, for decode total kv layers
+125 -7
View File
@@ -474,15 +474,133 @@ class CommonKVManager(BaseKVManager):
def get_mla_kv_ptrs_with_pp(
self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int]
) -> Tuple[List[int], List[int], int]:
# Fast path: both sides use exactly the same PP layout
if len(src_kv_ptrs) == len(dst_kv_ptrs):
return src_kv_ptrs, dst_kv_ptrs, len(src_kv_ptrs)
mla_ratios = getattr(self.kv_args, "mla_compression_ratios", None)
if mla_ratios:
# Compressed-MLA (e.g. DeepSeek V4): the flat list is organized
# by buffer type (compression-ratio bucket) rather than by
# layer, so we locate the sub-range for this PP stage inside each
# section of the dst flat list.
sliced_src_kv_ptrs, sliced_dst_kv_ptrs = self._mla_slice_ptrs_for_pp(
src_kv_ptrs, dst_kv_ptrs, mla_ratios
)
return (
sliced_src_kv_ptrs,
sliced_dst_kv_ptrs,
len(sliced_src_kv_ptrs),
)
# Regular MLA PP slicing
start_layer = self.kv_args.prefill_start_layer
end_layer = start_layer + len(src_kv_ptrs)
if len(src_kv_ptrs) == len(dst_kv_ptrs):
sliced_dst_kv_ptrs = dst_kv_ptrs
else:
# Decode pp size should be equal to prefill pp size or 1
sliced_dst_kv_ptrs = dst_kv_ptrs[start_layer:end_layer]
layers_current_pp_stage = len(src_kv_ptrs)
return src_kv_ptrs, sliced_dst_kv_ptrs, layers_current_pp_stage
# Decode pp size should be equal to prefill pp size or 1
sliced_dst_kv_ptrs = dst_kv_ptrs[start_layer:end_layer]
return src_kv_ptrs, sliced_dst_kv_ptrs, len(src_kv_ptrs)
def _mla_slice_ptrs_for_pp(
self,
src_kv_ptrs: List[int],
dst_kv_ptrs: List[int],
mla_ratios: List[int],
) -> Tuple[List[int], List[int]]:
"""Produce aligned (src, dst) pointer lists for compressed-MLA
pools (e.g. DeepSeek V4) under PP.
The pool produces two possible flat-list layouts (selected via dst
length):
- kv_data layout, length = 2 * c4_L + c128_L:
[c4_layer_{0..c4_L-1},
c4_indexer_layer_{0..c4_L-1},
c128_layer_{0..c128_L-1}]
Each section is indexed by compressed-layer id within that
compression bucket.
- state_data layout, length = swa_L + 2 * c4_L + c128_L:
[swa_layer_{0..swa_L-1},
compress_state_{non-None, c4_L + c128_L},
indexer_compress_state_{non-None, c4_L}]
``swa_L`` is the SWA pool's actual buffer count
(``num_effective_layers``), which can be smaller than
``len(mla_ratios)`` when the HF config's ``compress_ratios``
list contains entries for layers not materialized into the SWA
pool (e.g. an MTP/nextn slot at the tail).
src is already PP-filtered on the prefill side. dst is the
decode-side full-model list (when decode is PP=1). We slice dst to
match src's PP stage. If src itself is also full-model, it is
returned unchanged.
"""
start_layer = self.kv_args.prefill_start_layer
end_layer = getattr(self.kv_args, "prefill_end_layer", None)
assert end_layer is not None, (
"KVArgs.prefill_end_layer must be set when using "
"compressed-MLA PD with PP"
)
c4_full = sum(1 for r in mla_ratios if r == 4)
c128_full = sum(1 for r in mla_ratios if r == 128)
kv_layout_len = 2 * c4_full + c128_full
c4_off_s = sum(1 for r in mla_ratios[:start_layer] if r == 4)
c4_off_e = sum(1 for r in mla_ratios[:end_layer] if r == 4)
c128_off_s = sum(1 for r in mla_ratios[:start_layer] if r == 128)
c128_off_e = sum(1 for r in mla_ratios[:end_layer] if r == 128)
if len(dst_kv_ptrs) == kv_layout_len:
sliced_dst = (
list(dst_kv_ptrs[c4_off_s:c4_off_e])
+ list(dst_kv_ptrs[c4_full + c4_off_s : c4_full + c4_off_e])
+ list(dst_kv_ptrs[2 * c4_full + c128_off_s : 2 * c4_full + c128_off_e])
)
return src_kv_ptrs, sliced_dst
# State-data layout. ``swa_L`` is derived from the actual dst
# length so we tolerate cases where the SWA pool has fewer
# buffers than ``len(mla_ratios)`` (e.g. nextn padding).
swa_L = len(dst_kv_ptrs) - 2 * c4_full - c128_full
if swa_L < 0 or swa_L > len(mla_ratios):
raise ValueError(
f"Unexpected compressed-MLA dst_kv_ptrs length "
f"{len(dst_kv_ptrs)}; expected either {kv_layout_len} "
f"(kv_data) or swa_L + {2 * c4_full + c128_full} "
f"(state_data) given compression_ratios "
f"(c4={c4_full}, c128={c128_full}, "
f"total={len(mla_ratios)})."
)
# Guard against asking the prefill side to read past the SWA
# pool boundary.
assert end_layer <= swa_L, (
f"prefill_end_layer ({end_layer}) exceeds dst SWA pool "
f"buffer count ({swa_L}); compression_ratios may include "
f"layers (e.g. nextn) that the SWA pool does not cover."
)
# compress_state non-None count up to L = count(r != 0).
c_non_zero_s = sum(1 for r in mla_ratios[:start_layer] if r != 0)
c_non_zero_e = sum(1 for r in mla_ratios[:end_layer] if r != 0)
compress_section_start = swa_L
indexer_section_start = swa_L + (c4_full + c128_full)
sliced_dst = (
list(dst_kv_ptrs[start_layer:end_layer])
+ list(
dst_kv_ptrs[
compress_section_start
+ c_non_zero_s : compress_section_start
+ c_non_zero_e
]
)
+ list(
dst_kv_ptrs[
indexer_section_start + c4_off_s : indexer_section_start + c4_off_e
]
)
)
return src_kv_ptrs, sliced_dst
class CommonKVSender(BaseKVSender):
@@ -55,6 +55,7 @@ from sglang.srt.mem_cache.common import (
maybe_cache_unfinished_req,
release_kv_cache,
)
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.observability.req_time_stats import set_schedule_time_batch
if TYPE_CHECKING:
@@ -146,6 +147,8 @@ class PrefillBootstrapQueue:
kv_args.pp_rank = self.pp_rank
kv_args.system_dp_rank = self.scheduler.ps.dp_rank
kv_args.prefill_start_layer = self.token_to_kv_pool.start_layer
kv_args.prefill_end_layer = self.token_to_kv_pool.end_layer
kv_args.mla_compression_ratios = None
kv_data_ptrs, kv_data_lens, kv_item_lens = (
self.token_to_kv_pool.get_contiguous_buf_infos()
)
@@ -185,6 +188,13 @@ class PrefillBootstrapQueue:
req_to_token_pool=req_to_token_pool,
)
if isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool):
# V4's KVCache is organized by compression-ratio
# buckets rather than by layer.
kv_args.mla_compression_ratios = list(
self.token_to_kv_pool.compression_ratios
)
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
kv_manager = kv_manager_class(
kv_args,
+4 -1
View File
@@ -2322,7 +2322,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
)
# Update fields
self.input_ids = self.output_ids
# Coerce to int64: torch sampling helpers (sampling_from_probs_torch /
# top_k_top_p_min_p_sampling_from_probs_torch) return int32 token ids,
# but downstream kernels enforce int64 (e.g. DeepSeek-V4 hash_topk).
self.input_ids = self.output_ids.to(torch.int64)
self.output_ids = None
if self.model_config.is_encoder_decoder:
@@ -613,9 +613,13 @@ class SchedulerPPMixin:
batch.global_num_tokens = global_num_tokens
batch.global_num_tokens_for_logprob = global_num_tokens
hs = (
getattr(model_config, "hc_hidden_size", None)
or model_config.hidden_size
)
proxy_tensors = {
"hidden_states": torch.zeros(
(current_seq_len, model_config.hidden_size),
(current_seq_len, hs),
dtype=model_config.dtype,
device=self.device,
),
@@ -401,6 +401,19 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
self.state_dtype = state_dtype
self.compression_ratios = compression_ratios
# Determine this PP stage's absolute layer range
if (
start_layer is not None
and end_layer is not None
and len(compression_ratios) >= end_layer
):
self._stage_start = start_layer
self._stage_end = end_layer
else:
self._stage_start = 0
self._stage_end = len(compression_ratios)
stage_ratios = compression_ratios[self._stage_start : self._stage_end]
assert page_size % swa_page_size == 0
self.swa_size = swa_size
@@ -412,8 +425,8 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
self.qk_rope_head_dim = qk_rope_head_dim
self.indexer_head_dim = indexer_head_dim
c4_layer_num = sum(1 for r in compression_ratios if r == 4)
c128_layer_num = sum(1 for r in compression_ratios if r == 128)
c4_layer_num = sum(1 for r in stage_ratios if r == 4)
c128_layer_num = sum(1 for r in stage_ratios if r == 128)
c4_page_size = page_size // 4
c128_page_size = page_size // 128
self.swa_kv_pool = DeepSeekV4SingleKVPool(
@@ -467,6 +480,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
self._init_paged_compress_states(enable_memory_saver)
self._should_cache_swa = envs.SGLANG_OPT_CACHE_SWA_TRANSLATION.get()
self.cached_loc = None
def register_mapping(self, full_to_swa_index_mapping: torch.Tensor):
self.full_to_swa_index_mapping = full_to_swa_index_mapping
@@ -535,29 +549,34 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
def _init_paged_compress_states(self, enable_memory_saver: bool):
c4_state_pool_size = self.c4_state_pool_size
c128_state_pool_size = self.c128_state_pool_size
self.compress_state_pools: List[CompressStatePool] = []
self.indexer_compress_state_pools: List[CompressStatePool] = []
total_L = len(self.compression_ratios)
self.compress_state_pools: List[Optional[CompressStatePool]] = [None] * total_L
self.indexer_compress_state_pools: List[Optional[CompressStatePool]] = [
None
] * total_L
for ratio in self.compression_ratios:
for idx in range(self._stage_start, self._stage_end):
ratio = self.compression_ratios[idx]
if ratio == 0:
continue
overlap = ratio == 4
compress_state_pool = indexer_compress_state_pool = None
size = c4_state_pool_size if ratio == 4 else c128_state_pool_size
ring_size = self.get_ring_size(ratio) if ratio != 0 else 0
if ratio != 0:
compress_state_pool = CompressStatePool(
size=size,
ring_size=ring_size,
overlap=overlap,
head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim,
dtype=self.state_dtype,
device=self.device,
enable_memory_saver=enable_memory_saver,
ratio=ratio,
online=(ratio == 128 and ONLINE_C128),
)
ring_size = self.get_ring_size(ratio)
self.compress_state_pools[idx] = CompressStatePool(
size=size,
ring_size=ring_size,
overlap=overlap,
head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim,
dtype=self.state_dtype,
device=self.device,
enable_memory_saver=enable_memory_saver,
ratio=ratio,
online=(ratio == 128 and ONLINE_C128),
)
if ratio == 4:
indexer_compress_state_pool = CompressStatePool(
self.indexer_compress_state_pools[idx] = CompressStatePool(
size=size,
ring_size=ring_size,
overlap=overlap,
@@ -568,38 +587,31 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
ratio=ratio,
)
self.compress_state_pools.append(compress_state_pool)
self.indexer_compress_state_pools.append(indexer_compress_state_pool)
def _init_compressed_layer_mapping(self):
c1_cnt, c4_cnt, c128_cnt = 0, 0, 0
self.layer_mapping: List[DeepSeekV4LayerItem] = []
c1_cnt = c4_cnt = c128_cnt = 0
total_L = len(self.compression_ratios)
self.layer_mapping: List[Optional[DeepSeekV4LayerItem]] = [None] * total_L
for ratio in self.compression_ratios:
for idx in range(self._stage_start, self._stage_end):
ratio = self.compression_ratios[idx]
if ratio == 0:
self.layer_mapping.append(
DeepSeekV4LayerItem(
compress_ratio=0,
compress_layer_id=c1_cnt,
)
self.layer_mapping[idx] = DeepSeekV4LayerItem(
compress_ratio=0,
compress_layer_id=c1_cnt,
)
c1_cnt += 1
elif ratio == 4:
self.layer_mapping.append(
DeepSeekV4LayerItem(
compress_ratio=4,
compress_layer_id=c4_cnt,
compress_kv_pool=self.c4_kv_pool,
)
self.layer_mapping[idx] = DeepSeekV4LayerItem(
compress_ratio=4,
compress_layer_id=c4_cnt,
compress_kv_pool=self.c4_kv_pool,
)
c4_cnt += 1
elif ratio == 128:
self.layer_mapping.append(
DeepSeekV4LayerItem(
compress_ratio=128,
compress_layer_id=c128_cnt,
compress_kv_pool=self.c128_kv_pool,
)
self.layer_mapping[idx] = DeepSeekV4LayerItem(
compress_ratio=128,
compress_layer_id=c128_cnt,
compress_kv_pool=self.c128_kv_pool,
)
c128_cnt += 1
else:
@@ -625,9 +637,13 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
), "Only c4 layers have indexer states."
return indexer_compress_state_pool
def _swa_local_layer_id(self, layer_id: int) -> int:
"""Convert absolute model layer_id to SWA-pool-local (PP-stage-local) index."""
return layer_id - self._stage_start
def get_swa_key_buffer(self, layer_id: int) -> torch.Tensor:
self.wait_layer_transfer(layer_id)
return self.swa_kv_pool.get_key_buffer(layer_id)
return self.swa_kv_pool.get_key_buffer(self._swa_local_layer_id(layer_id))
def set_swa_key_buffer(
self,
@@ -635,7 +651,9 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
loc: torch.Tensor,
cache_nope_fp8_rope_bf16_pack: NopeFp8RopeBf16Pack,
) -> None:
self.swa_kv_pool.set_key_buffer(layer_id, loc, cache_nope_fp8_rope_bf16_pack)
self.swa_kv_pool.set_key_buffer(
self._swa_local_layer_id(layer_id), loc, cache_nope_fp8_rope_bf16_pack
)
def get_extra_key_page_size(self, layer_id: int) -> int:
_, _, compress_kv_pool = self.layer_mapping[layer_id]
@@ -715,12 +733,12 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
) -> None:
swa_loc = self.translate_loc_from_full_to_swa(raw_loc)
self.swa_kv_pool.set_key_buffer(
layer_id, swa_loc, cache_nope_fp8_rope_bf16_pack
self._swa_local_layer_id(layer_id), swa_loc, cache_nope_fp8_rope_bf16_pack
)
def get_swa_key_buffer_radix(self, layer_id: int) -> torch.Tensor:
self.wait_layer_transfer(layer_id)
return self.swa_kv_pool.get_key_buffer(layer_id)
return self.swa_kv_pool.get_key_buffer(self._swa_local_layer_id(layer_id))
def set_swa_key_buffer_radix_fused(
self,
@@ -729,12 +747,14 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
cache_k: torch.Tensor,
) -> None:
if self._should_cache_swa:
if layer_id == 0:
if layer_id == self.start_layer or self.cached_loc is None:
self.cached_loc = self.translate_loc_from_full_to_swa(raw_loc)
swa_loc = self.cached_loc
else:
swa_loc = self.translate_loc_from_full_to_swa(raw_loc)
return self.swa_kv_pool.set_key_buffer_fused(layer_id, swa_loc, cache_k)
return self.swa_kv_pool.set_key_buffer_fused(
self._swa_local_layer_id(layer_id), swa_loc, cache_k
)
def set_swa_key_buffer_radix_fused_norm_rope(
self,
@@ -759,7 +779,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
freqs_cis=freqs_cis,
positions=positions,
out_loc=swa_loc,
kvcache=self.swa_kv_pool.kv_buffer[layer_id],
kvcache=self.swa_kv_pool.kv_buffer[self._swa_local_layer_id(layer_id)],
page_size=self.swa_kv_pool.page_size,
)
@@ -170,6 +170,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
enable_mamba_track: bool,
ne_token_table: Optional[torch.Tensor] = None,
is_hybrid_swa: bool = False,
hc_hidden_size: Optional[int] = None,
) -> "DecodeInputBuffers":
with torch.device(device):
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
@@ -203,10 +204,16 @@ class DecodeInputBuffers(ForwardInputBuffers):
)
if pp_size > 1:
# mHC (e.g. DSV4) flattens residual into hidden_states (size = hc_hidden_size).
is_mhc = hc_hidden_size is not None
hs = hc_hidden_size if is_mhc else hidden_size
pp_proxy_tensors = {
"hidden_states": torch.zeros((max_bs, hidden_size), dtype=dtype),
"residual": torch.zeros((max_bs, hidden_size), dtype=dtype),
"hidden_states": torch.zeros((max_bs, hs), dtype=dtype),
}
if not is_mhc:
pp_proxy_tensors["residual"] = torch.zeros(
(max_bs, hidden_size), dtype=dtype
)
else:
pp_proxy_tensors = None
@@ -686,6 +693,9 @@ class CudaGraphRunner:
model_runner.token_table if self.use_ngram_embedding else None
),
is_hybrid_swa=model_runner.is_hybrid_swa,
hc_hidden_size=getattr(
self.model_runner.model_config, "hc_hidden_size", None
),
)
self.buffers.share_buffers()
+99 -39
View File
@@ -2,7 +2,16 @@ from __future__ import annotations
import concurrent.futures
import logging
from typing import TYPE_CHECKING, Iterable, List, Literal, Optional, Set, Tuple
from typing import (
TYPE_CHECKING,
Iterable,
List,
Literal,
Optional,
Set,
Tuple,
Union,
)
import torch
import torch.nn as nn
@@ -49,7 +58,7 @@ from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.moe import get_moe_a2a_backend
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8
from sglang.srt.layers.utils import get_layer_id
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,
@@ -62,6 +71,7 @@ from sglang.srt.model_executor.cuda_graph_runner import (
compile_in_capture_mode,
get_is_capture_mode,
)
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
from sglang.srt.model_loader.utils import maybe_executor_submit, should_async_load
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.dbrx import ReplicatedLinear
@@ -86,10 +96,7 @@ if TYPE_CHECKING:
)
from sglang.srt.layers.quantization import QuantizationConfig
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch,
PPProxyTensors,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@triton.jit
@@ -870,11 +877,15 @@ class DeepseekV4Model(nn.Module):
) -> None:
super().__init__()
self.pp_group = get_pp_group()
self.embed_tokens = VocabParallelEmbedding(
config.vocab_size,
config.hidden_size,
enable_tp=not is_dp_attention_enabled(),
)
self.hidden_size = config.hidden_size
if self.pp_group.is_first_rank:
self.embed_tokens = VocabParallelEmbedding(
config.vocab_size,
config.hidden_size,
enable_tp=not is_dp_attention_enabled(),
)
else:
self.embed_tokens = PPMissingLayer()
self.rms_norm_eps = config.rms_norm_eps
self.alt_streams = (
[torch.cuda.Stream() for _ in range(5)] if (_is_cuda or _is_hip) else None
@@ -892,17 +903,21 @@ class DeepseekV4Model(nn.Module):
pp_size=self.pp_group.world_size,
prefix=add_prefix("layers", prefix),
)
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
if self.pp_group.is_last_rank:
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
else:
self.norm = PPMissingLayer()
self.gemm_output_zero_allocator_size = 0
self.hc_eps = config.hc_eps
self.hc_mult = hc_mult = config.hc_mult
self.norm_eps = config.rms_norm_eps
hc_dim = hc_mult * config.hidden_size
self.hc_head_fn = nn.Parameter(
torch.empty(hc_mult, hc_dim, dtype=torch.float32)
)
self.hc_head_base = nn.Parameter(torch.empty(hc_mult, dtype=torch.float32))
self.hc_head_scale = nn.Parameter(torch.empty(1, dtype=torch.float32))
if self.pp_group.is_last_rank:
hc_dim = hc_mult * config.hidden_size
self.hc_head_fn = nn.Parameter(
torch.empty(hc_mult, hc_dim, dtype=torch.float32)
)
self.hc_head_base = nn.Parameter(torch.empty(hc_mult, dtype=torch.float32))
self.hc_head_scale = nn.Parameter(torch.empty(1, dtype=torch.float32))
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
if self.nsa_enable_prefill_cp:
@@ -940,9 +955,19 @@ class DeepseekV4Model(nn.Module):
positions: torch.Tensor,
forward_batch: ForwardBatch,
input_embeds: Optional[torch.Tensor],
) -> torch.Tensor:
hidden_states = self.embed_tokens(input_ids)
hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1)
pp_proxy_tensors: Optional[PPProxyTensors] = None,
) -> Union[torch.Tensor, PPProxyTensors]:
if self.pp_group.is_first_rank:
hidden_states = self.embed_tokens(input_ids)
hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1)
else:
assert pp_proxy_tensors is not None
hidden_states = pp_proxy_tensors["hidden_states"]
# Unflatten 2D PP IPC tensor back to 3D mHC shape.
if hidden_states.ndim == 2:
hidden_states = hidden_states.view(
hidden_states.shape[0], self.hc_mult, self.hidden_size
)
if get_attention_dp_size() > 1 and get_moe_a2a_backend().is_none():
input_ids_global = torch.empty(
@@ -956,7 +981,8 @@ class DeepseekV4Model(nn.Module):
input_ids_global = input_ids
if nsa_use_prefill_cp(forward_batch):
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
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)
for i in range(self.start_layer, self.end_layer):
@@ -969,7 +995,8 @@ class DeepseekV4Model(nn.Module):
input_ids_global=input_ids_global,
)
if nsa_use_prefill_cp(forward_batch):
# CP all-gather only on the last PP rank; PP IPC carries CP-split tensors.
if self.pp_group.is_last_rank and nsa_use_prefill_cp(forward_batch):
hidden_states = cp_all_gather_rerange_output(
hidden_states,
self.cp_size,
@@ -977,6 +1004,10 @@ class DeepseekV4Model(nn.Module):
torch.cuda.current_stream(),
)
if not self.pp_group.is_last_rank:
# Flatten 3D mHC tensor for PP IPC.
return PPProxyTensors({"hidden_states": hidden_states.flatten(1)})
pre_hc_head = hidden_states.flatten(1)
hidden_states = self.hc_head(
@@ -1003,28 +1034,37 @@ class DeepseekV4ForCausalLM(nn.Module):
config, quant_config, prefix=add_prefix("model", prefix)
)
self.pp_group = get_pp_group()
if config.tie_word_embeddings:
self.lm_head = self.model.embed_tokens
if self.pp_group.is_last_rank:
if self.pp_group.world_size == 1 and config.tie_word_embeddings:
self.lm_head = self.model.embed_tokens
else:
self.lm_head = ParallelLMHead(
config.vocab_size,
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_global_server_args().enable_dp_lm_head,
)
else:
self.lm_head = ParallelLMHead(
config.vocab_size,
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_global_server_args().enable_dp_lm_head,
)
self.lm_head = PPMissingLayer()
self.logits_processor = LogitsProcessor(config)
self.capture_aux_hidden_states = False
get_attn_tp_context().init_context(config.q_lora_rank, is_nsa=True)
self._routed_experts_weights_of_layer = LazyValue(
lambda: {
layer_id: layer.mlp.get_moe_weights()
for layer_id, layer in enumerate(self.model.layers)
if isinstance(layer.mlp, deepseek_v2.DeepseekV2MoE)
layer_id: self.model.layers[layer_id].mlp.get_moe_weights()
for layer_id in range(self.model.start_layer, self.model.end_layer)
if isinstance(
self.model.layers[layer_id].mlp, deepseek_v2.DeepseekV2MoE
)
}
)
# Expose start_layer/end_layer for model_runner PP support
self.start_layer = self.model.start_layer
self.end_layer = self.model.end_layer
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
if self.nsa_enable_prefill_cp:
self.cp_rank = get_attention_cp_rank()
@@ -1077,8 +1117,11 @@ class DeepseekV4ForCausalLM(nn.Module):
with get_attn_tp_context().maybe_input_scattered(forward_batch):
hidden_states = self.model.forward(
input_ids, positions, forward_batch, input_embeds
input_ids, positions, forward_batch, input_embeds, pp_proxy_tensors
)
if not self.pp_group.is_last_rank:
return hidden_states
aux_hidden_states = None
if self.capture_aux_hidden_states:
hidden_states, aux_hidden_states = hidden_states
@@ -1098,7 +1141,10 @@ class DeepseekV4ForCausalLM(nn.Module):
if is_nextn:
layers = [self.model.decoder]
else:
layers = self.model.layers
layers = [
self.model.layers[layer_id]
for layer_id in range(self.model.start_layer, self.model.end_layer)
]
for layer in layers:
attn = layer.self_attn
G = attn.n_local_groups
@@ -1121,7 +1167,8 @@ class DeepseekV4ForCausalLM(nn.Module):
if is_nextn:
return
for layer in self.model.layers:
for layer_id in range(self.model.start_layer, self.model.end_layer):
layer = self.model.layers[layer_id]
self_attn = layer.self_attn
if self_attn.compress_ratio != 0 and not self_attn.compressor.ape_converted:
self_attn.compressor.apply_ape_hotfix()
@@ -1389,7 +1436,15 @@ class DeepseekV4ForCausalLM(nn.Module):
and not self.pp_group.is_first_rank
):
continue
if ".norm." in name and not self.pp_group.is_last_rank:
if (
name == "model.norm.weight"
and not self.pp_group.is_last_rank
):
continue
if (
name.startswith("model.hc_head_")
or name == "lm_head.weight"
) and not self.pp_group.is_last_rank:
continue
elif COMPRESSOR_PART in name:
is_kv = name.endswith(".wkv.weight")
@@ -1493,6 +1548,11 @@ class DeepseekV4ForCausalLM(nn.Module):
unloaded_params = params_dict.keys() - loaded_params
skipped_checking_patterns = ["attn_mqa.k_scale", "attn_mqa.v_scale"]
if not self.pp_group.is_first_rank:
skipped_checking_patterns.append("embed_tokens")
if not self.pp_group.is_last_rank:
skipped_checking_patterns.append("model.norm.")
skipped_checking_patterns.extend(["lm_head", "hc_head_"])
if is_nextn:
skipped_checking_patterns.extend(["lm_head", "embed_tokens"])
unloaded_params = {