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:
co-authored by
Shangming Cai
xuyongfei
parent
bda01d2435
commit
162540e0a8
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user