diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 5a9a93101..111145ef6 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -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: diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 264f3ca0d..1e1b4b4f5 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -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 diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 596b58303..555ef5215 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -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): diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index e2cc107de..85273ead2 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -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, diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 3e0327e32..e215da305 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -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: diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index ab1cc6d2a..154ad94be 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -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, ), diff --git a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py index 94b930b89..a185b02c5 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -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, ) diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index 569f294e8..1afcfb286 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -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() diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index f4128884a..b8df0c313 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -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 = {