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.spec_hidden_size = (
|
||||||
self.hidden_size * hc_mult if hc_mult > 1 else self.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_hidden_layers = self.hf_text_config.num_hidden_layers
|
||||||
self.num_attention_layers = self.num_hidden_layers
|
self.num_attention_layers = self.num_hidden_layers
|
||||||
if "LongcatFlashForCausalLM" in self.hf_config.architectures:
|
if "LongcatFlashForCausalLM" in self.hf_config.architectures:
|
||||||
|
|||||||
@@ -47,11 +47,21 @@ class KVArgs:
|
|||||||
kv_head_num: int
|
kv_head_num: int
|
||||||
total_kv_head_num: int
|
total_kv_head_num: int
|
||||||
page_size: int
|
page_size: int
|
||||||
|
# for system dp
|
||||||
|
system_dp_rank: int
|
||||||
# for pp prefill
|
# for pp prefill
|
||||||
pp_rank: int
|
pp_rank: int
|
||||||
prefill_start_layer: int
|
prefill_start_layer: int
|
||||||
# for system dp
|
# Absolute end layer (exclusive) for this prefill PP stage. Needed to
|
||||||
system_dp_rank: int
|
# 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
|
# Only used of npu, for kv buf groups
|
||||||
kv_buf_groups: int
|
kv_buf_groups: int
|
||||||
# Only used of npu, for decode total kv layers
|
# Only used of npu, for decode total kv layers
|
||||||
|
|||||||
@@ -474,15 +474,133 @@ class CommonKVManager(BaseKVManager):
|
|||||||
def get_mla_kv_ptrs_with_pp(
|
def get_mla_kv_ptrs_with_pp(
|
||||||
self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int]
|
self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int]
|
||||||
) -> Tuple[List[int], List[int], 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
|
start_layer = self.kv_args.prefill_start_layer
|
||||||
end_layer = start_layer + len(src_kv_ptrs)
|
end_layer = start_layer + len(src_kv_ptrs)
|
||||||
if len(src_kv_ptrs) == len(dst_kv_ptrs):
|
# Decode pp size should be equal to prefill pp size or 1
|
||||||
sliced_dst_kv_ptrs = dst_kv_ptrs
|
sliced_dst_kv_ptrs = dst_kv_ptrs[start_layer:end_layer]
|
||||||
else:
|
return src_kv_ptrs, sliced_dst_kv_ptrs, len(src_kv_ptrs)
|
||||||
# Decode pp size should be equal to prefill pp size or 1
|
|
||||||
sliced_dst_kv_ptrs = dst_kv_ptrs[start_layer:end_layer]
|
def _mla_slice_ptrs_for_pp(
|
||||||
layers_current_pp_stage = len(src_kv_ptrs)
|
self,
|
||||||
return src_kv_ptrs, sliced_dst_kv_ptrs, layers_current_pp_stage
|
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):
|
class CommonKVSender(BaseKVSender):
|
||||||
|
|||||||
@@ -55,6 +55,7 @@ from sglang.srt.mem_cache.common import (
|
|||||||
maybe_cache_unfinished_req,
|
maybe_cache_unfinished_req,
|
||||||
release_kv_cache,
|
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
|
from sglang.srt.observability.req_time_stats import set_schedule_time_batch
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -146,6 +147,8 @@ class PrefillBootstrapQueue:
|
|||||||
kv_args.pp_rank = self.pp_rank
|
kv_args.pp_rank = self.pp_rank
|
||||||
kv_args.system_dp_rank = self.scheduler.ps.dp_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_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 = (
|
kv_data_ptrs, kv_data_lens, kv_item_lens = (
|
||||||
self.token_to_kv_pool.get_contiguous_buf_infos()
|
self.token_to_kv_pool.get_contiguous_buf_infos()
|
||||||
)
|
)
|
||||||
@@ -185,6 +188,13 @@ class PrefillBootstrapQueue:
|
|||||||
req_to_token_pool=req_to_token_pool,
|
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_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
|
||||||
kv_manager = kv_manager_class(
|
kv_manager = kv_manager_class(
|
||||||
kv_args,
|
kv_args,
|
||||||
|
|||||||
@@ -2322,7 +2322,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Update fields
|
# 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
|
self.output_ids = None
|
||||||
|
|
||||||
if self.model_config.is_encoder_decoder:
|
if self.model_config.is_encoder_decoder:
|
||||||
|
|||||||
@@ -613,9 +613,13 @@ class SchedulerPPMixin:
|
|||||||
batch.global_num_tokens = global_num_tokens
|
batch.global_num_tokens = global_num_tokens
|
||||||
batch.global_num_tokens_for_logprob = 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 = {
|
proxy_tensors = {
|
||||||
"hidden_states": torch.zeros(
|
"hidden_states": torch.zeros(
|
||||||
(current_seq_len, model_config.hidden_size),
|
(current_seq_len, hs),
|
||||||
dtype=model_config.dtype,
|
dtype=model_config.dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -401,6 +401,19 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
self.state_dtype = state_dtype
|
self.state_dtype = state_dtype
|
||||||
self.compression_ratios = compression_ratios
|
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
|
assert page_size % swa_page_size == 0
|
||||||
|
|
||||||
self.swa_size = swa_size
|
self.swa_size = swa_size
|
||||||
@@ -412,8 +425,8 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
self.qk_rope_head_dim = qk_rope_head_dim
|
self.qk_rope_head_dim = qk_rope_head_dim
|
||||||
self.indexer_head_dim = indexer_head_dim
|
self.indexer_head_dim = indexer_head_dim
|
||||||
|
|
||||||
c4_layer_num = sum(1 for r in compression_ratios if r == 4)
|
c4_layer_num = sum(1 for r in stage_ratios if r == 4)
|
||||||
c128_layer_num = sum(1 for r in compression_ratios if r == 128)
|
c128_layer_num = sum(1 for r in stage_ratios if r == 128)
|
||||||
c4_page_size = page_size // 4
|
c4_page_size = page_size // 4
|
||||||
c128_page_size = page_size // 128
|
c128_page_size = page_size // 128
|
||||||
self.swa_kv_pool = DeepSeekV4SingleKVPool(
|
self.swa_kv_pool = DeepSeekV4SingleKVPool(
|
||||||
@@ -467,6 +480,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
self._init_paged_compress_states(enable_memory_saver)
|
self._init_paged_compress_states(enable_memory_saver)
|
||||||
|
|
||||||
self._should_cache_swa = envs.SGLANG_OPT_CACHE_SWA_TRANSLATION.get()
|
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):
|
def register_mapping(self, full_to_swa_index_mapping: torch.Tensor):
|
||||||
self.full_to_swa_index_mapping = full_to_swa_index_mapping
|
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):
|
def _init_paged_compress_states(self, enable_memory_saver: bool):
|
||||||
c4_state_pool_size = self.c4_state_pool_size
|
c4_state_pool_size = self.c4_state_pool_size
|
||||||
c128_state_pool_size = self.c128_state_pool_size
|
c128_state_pool_size = self.c128_state_pool_size
|
||||||
self.compress_state_pools: List[CompressStatePool] = []
|
total_L = len(self.compression_ratios)
|
||||||
self.indexer_compress_state_pools: List[CompressStatePool] = []
|
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
|
overlap = ratio == 4
|
||||||
compress_state_pool = indexer_compress_state_pool = None
|
|
||||||
size = c4_state_pool_size if ratio == 4 else c128_state_pool_size
|
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
|
ring_size = self.get_ring_size(ratio)
|
||||||
if ratio != 0:
|
|
||||||
compress_state_pool = CompressStatePool(
|
self.compress_state_pools[idx] = CompressStatePool(
|
||||||
size=size,
|
size=size,
|
||||||
ring_size=ring_size,
|
ring_size=ring_size,
|
||||||
overlap=overlap,
|
overlap=overlap,
|
||||||
head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim,
|
head_dim=self.qk_nope_head_dim + self.qk_rope_head_dim,
|
||||||
dtype=self.state_dtype,
|
dtype=self.state_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
enable_memory_saver=enable_memory_saver,
|
enable_memory_saver=enable_memory_saver,
|
||||||
ratio=ratio,
|
ratio=ratio,
|
||||||
online=(ratio == 128 and ONLINE_C128),
|
online=(ratio == 128 and ONLINE_C128),
|
||||||
)
|
)
|
||||||
|
|
||||||
if ratio == 4:
|
if ratio == 4:
|
||||||
indexer_compress_state_pool = CompressStatePool(
|
self.indexer_compress_state_pools[idx] = CompressStatePool(
|
||||||
size=size,
|
size=size,
|
||||||
ring_size=ring_size,
|
ring_size=ring_size,
|
||||||
overlap=overlap,
|
overlap=overlap,
|
||||||
@@ -568,38 +587,31 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
ratio=ratio,
|
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):
|
def _init_compressed_layer_mapping(self):
|
||||||
c1_cnt, c4_cnt, c128_cnt = 0, 0, 0
|
c1_cnt = c4_cnt = c128_cnt = 0
|
||||||
self.layer_mapping: List[DeepSeekV4LayerItem] = []
|
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:
|
if ratio == 0:
|
||||||
self.layer_mapping.append(
|
self.layer_mapping[idx] = DeepSeekV4LayerItem(
|
||||||
DeepSeekV4LayerItem(
|
compress_ratio=0,
|
||||||
compress_ratio=0,
|
compress_layer_id=c1_cnt,
|
||||||
compress_layer_id=c1_cnt,
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
c1_cnt += 1
|
c1_cnt += 1
|
||||||
elif ratio == 4:
|
elif ratio == 4:
|
||||||
self.layer_mapping.append(
|
self.layer_mapping[idx] = DeepSeekV4LayerItem(
|
||||||
DeepSeekV4LayerItem(
|
compress_ratio=4,
|
||||||
compress_ratio=4,
|
compress_layer_id=c4_cnt,
|
||||||
compress_layer_id=c4_cnt,
|
compress_kv_pool=self.c4_kv_pool,
|
||||||
compress_kv_pool=self.c4_kv_pool,
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
c4_cnt += 1
|
c4_cnt += 1
|
||||||
elif ratio == 128:
|
elif ratio == 128:
|
||||||
self.layer_mapping.append(
|
self.layer_mapping[idx] = DeepSeekV4LayerItem(
|
||||||
DeepSeekV4LayerItem(
|
compress_ratio=128,
|
||||||
compress_ratio=128,
|
compress_layer_id=c128_cnt,
|
||||||
compress_layer_id=c128_cnt,
|
compress_kv_pool=self.c128_kv_pool,
|
||||||
compress_kv_pool=self.c128_kv_pool,
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
c128_cnt += 1
|
c128_cnt += 1
|
||||||
else:
|
else:
|
||||||
@@ -625,9 +637,13 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
), "Only c4 layers have indexer states."
|
), "Only c4 layers have indexer states."
|
||||||
return indexer_compress_state_pool
|
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:
|
def get_swa_key_buffer(self, layer_id: int) -> torch.Tensor:
|
||||||
self.wait_layer_transfer(layer_id)
|
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(
|
def set_swa_key_buffer(
|
||||||
self,
|
self,
|
||||||
@@ -635,7 +651,9 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
loc: torch.Tensor,
|
loc: torch.Tensor,
|
||||||
cache_nope_fp8_rope_bf16_pack: NopeFp8RopeBf16Pack,
|
cache_nope_fp8_rope_bf16_pack: NopeFp8RopeBf16Pack,
|
||||||
) -> None:
|
) -> 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:
|
def get_extra_key_page_size(self, layer_id: int) -> int:
|
||||||
_, _, compress_kv_pool = self.layer_mapping[layer_id]
|
_, _, compress_kv_pool = self.layer_mapping[layer_id]
|
||||||
@@ -715,12 +733,12 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
) -> None:
|
) -> None:
|
||||||
swa_loc = self.translate_loc_from_full_to_swa(raw_loc)
|
swa_loc = self.translate_loc_from_full_to_swa(raw_loc)
|
||||||
self.swa_kv_pool.set_key_buffer(
|
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:
|
def get_swa_key_buffer_radix(self, layer_id: int) -> torch.Tensor:
|
||||||
self.wait_layer_transfer(layer_id)
|
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(
|
def set_swa_key_buffer_radix_fused(
|
||||||
self,
|
self,
|
||||||
@@ -729,12 +747,14 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
cache_k: torch.Tensor,
|
cache_k: torch.Tensor,
|
||||||
) -> None:
|
) -> None:
|
||||||
if self._should_cache_swa:
|
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)
|
self.cached_loc = self.translate_loc_from_full_to_swa(raw_loc)
|
||||||
swa_loc = self.cached_loc
|
swa_loc = self.cached_loc
|
||||||
else:
|
else:
|
||||||
swa_loc = self.translate_loc_from_full_to_swa(raw_loc)
|
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(
|
def set_swa_key_buffer_radix_fused_norm_rope(
|
||||||
self,
|
self,
|
||||||
@@ -759,7 +779,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
freqs_cis=freqs_cis,
|
freqs_cis=freqs_cis,
|
||||||
positions=positions,
|
positions=positions,
|
||||||
out_loc=swa_loc,
|
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,
|
page_size=self.swa_kv_pool.page_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -170,6 +170,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
enable_mamba_track: bool,
|
enable_mamba_track: bool,
|
||||||
ne_token_table: Optional[torch.Tensor] = None,
|
ne_token_table: Optional[torch.Tensor] = None,
|
||||||
is_hybrid_swa: bool = False,
|
is_hybrid_swa: bool = False,
|
||||||
|
hc_hidden_size: Optional[int] = None,
|
||||||
) -> "DecodeInputBuffers":
|
) -> "DecodeInputBuffers":
|
||||||
with torch.device(device):
|
with torch.device(device):
|
||||||
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
|
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
|
||||||
@@ -203,10 +204,16 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if pp_size > 1:
|
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 = {
|
pp_proxy_tensors = {
|
||||||
"hidden_states": torch.zeros((max_bs, hidden_size), dtype=dtype),
|
"hidden_states": torch.zeros((max_bs, hs), dtype=dtype),
|
||||||
"residual": torch.zeros((max_bs, hidden_size), dtype=dtype),
|
|
||||||
}
|
}
|
||||||
|
if not is_mhc:
|
||||||
|
pp_proxy_tensors["residual"] = torch.zeros(
|
||||||
|
(max_bs, hidden_size), dtype=dtype
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
pp_proxy_tensors = None
|
pp_proxy_tensors = None
|
||||||
|
|
||||||
@@ -686,6 +693,9 @@ class CudaGraphRunner:
|
|||||||
model_runner.token_table if self.use_ngram_embedding else None
|
model_runner.token_table if self.use_ngram_embedding else None
|
||||||
),
|
),
|
||||||
is_hybrid_swa=model_runner.is_hybrid_swa,
|
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()
|
self.buffers.share_buffers()
|
||||||
|
|
||||||
|
|||||||
@@ -2,7 +2,16 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
import logging
|
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
|
||||||
import torch.nn as nn
|
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 import get_moe_a2a_backend
|
||||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
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.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 (
|
from sglang.srt.layers.utils.cp_utils import (
|
||||||
cp_all_gather_rerange_output,
|
cp_all_gather_rerange_output,
|
||||||
cp_split_and_rebuild_data,
|
cp_split_and_rebuild_data,
|
||||||
@@ -62,6 +71,7 @@ from sglang.srt.model_executor.cuda_graph_runner import (
|
|||||||
compile_in_capture_mode,
|
compile_in_capture_mode,
|
||||||
get_is_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.utils import maybe_executor_submit, should_async_load
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.models.dbrx import ReplicatedLinear
|
from sglang.srt.models.dbrx import ReplicatedLinear
|
||||||
@@ -86,10 +96,7 @@ if TYPE_CHECKING:
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization import QuantizationConfig
|
from sglang.srt.layers.quantization import QuantizationConfig
|
||||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import (
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
ForwardBatch,
|
|
||||||
PPProxyTensors,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
@@ -870,11 +877,15 @@ class DeepseekV4Model(nn.Module):
|
|||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.pp_group = get_pp_group()
|
self.pp_group = get_pp_group()
|
||||||
self.embed_tokens = VocabParallelEmbedding(
|
self.hidden_size = config.hidden_size
|
||||||
config.vocab_size,
|
if self.pp_group.is_first_rank:
|
||||||
config.hidden_size,
|
self.embed_tokens = VocabParallelEmbedding(
|
||||||
enable_tp=not is_dp_attention_enabled(),
|
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.rms_norm_eps = config.rms_norm_eps
|
||||||
self.alt_streams = (
|
self.alt_streams = (
|
||||||
[torch.cuda.Stream() for _ in range(5)] if (_is_cuda or _is_hip) else None
|
[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,
|
pp_size=self.pp_group.world_size,
|
||||||
prefix=add_prefix("layers", prefix),
|
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.gemm_output_zero_allocator_size = 0
|
||||||
self.hc_eps = config.hc_eps
|
self.hc_eps = config.hc_eps
|
||||||
self.hc_mult = hc_mult = config.hc_mult
|
self.hc_mult = hc_mult = config.hc_mult
|
||||||
self.norm_eps = config.rms_norm_eps
|
self.norm_eps = config.rms_norm_eps
|
||||||
hc_dim = hc_mult * config.hidden_size
|
if self.pp_group.is_last_rank:
|
||||||
self.hc_head_fn = nn.Parameter(
|
hc_dim = hc_mult * config.hidden_size
|
||||||
torch.empty(hc_mult, hc_dim, dtype=torch.float32)
|
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.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()
|
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
|
||||||
if self.nsa_enable_prefill_cp:
|
if self.nsa_enable_prefill_cp:
|
||||||
@@ -940,9 +955,19 @@ class DeepseekV4Model(nn.Module):
|
|||||||
positions: torch.Tensor,
|
positions: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
input_embeds: Optional[torch.Tensor],
|
input_embeds: Optional[torch.Tensor],
|
||||||
) -> torch.Tensor:
|
pp_proxy_tensors: Optional[PPProxyTensors] = None,
|
||||||
hidden_states = self.embed_tokens(input_ids)
|
) -> Union[torch.Tensor, PPProxyTensors]:
|
||||||
hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1)
|
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():
|
if get_attention_dp_size() > 1 and get_moe_a2a_backend().is_none():
|
||||||
input_ids_global = torch.empty(
|
input_ids_global = torch.empty(
|
||||||
@@ -956,7 +981,8 @@ class DeepseekV4Model(nn.Module):
|
|||||||
input_ids_global = input_ids
|
input_ids_global = input_ids
|
||||||
|
|
||||||
if nsa_use_prefill_cp(forward_batch):
|
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)
|
positions = cp_split_and_rebuild_position(forward_batch, positions)
|
||||||
|
|
||||||
for i in range(self.start_layer, self.end_layer):
|
for i in range(self.start_layer, self.end_layer):
|
||||||
@@ -969,7 +995,8 @@ class DeepseekV4Model(nn.Module):
|
|||||||
input_ids_global=input_ids_global,
|
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 = cp_all_gather_rerange_output(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
self.cp_size,
|
self.cp_size,
|
||||||
@@ -977,6 +1004,10 @@ class DeepseekV4Model(nn.Module):
|
|||||||
torch.cuda.current_stream(),
|
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)
|
pre_hc_head = hidden_states.flatten(1)
|
||||||
|
|
||||||
hidden_states = self.hc_head(
|
hidden_states = self.hc_head(
|
||||||
@@ -1003,28 +1034,37 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
config, quant_config, prefix=add_prefix("model", prefix)
|
config, quant_config, prefix=add_prefix("model", prefix)
|
||||||
)
|
)
|
||||||
self.pp_group = get_pp_group()
|
self.pp_group = get_pp_group()
|
||||||
if config.tie_word_embeddings:
|
if self.pp_group.is_last_rank:
|
||||||
self.lm_head = self.model.embed_tokens
|
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:
|
else:
|
||||||
self.lm_head = ParallelLMHead(
|
self.lm_head = PPMissingLayer()
|
||||||
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.logits_processor = LogitsProcessor(config)
|
self.logits_processor = LogitsProcessor(config)
|
||||||
self.capture_aux_hidden_states = False
|
self.capture_aux_hidden_states = False
|
||||||
get_attn_tp_context().init_context(config.q_lora_rank, is_nsa=True)
|
get_attn_tp_context().init_context(config.q_lora_rank, is_nsa=True)
|
||||||
|
|
||||||
self._routed_experts_weights_of_layer = LazyValue(
|
self._routed_experts_weights_of_layer = LazyValue(
|
||||||
lambda: {
|
lambda: {
|
||||||
layer_id: layer.mlp.get_moe_weights()
|
layer_id: self.model.layers[layer_id].mlp.get_moe_weights()
|
||||||
for layer_id, layer in enumerate(self.model.layers)
|
for layer_id in range(self.model.start_layer, self.model.end_layer)
|
||||||
if isinstance(layer.mlp, deepseek_v2.DeepseekV2MoE)
|
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()
|
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
|
||||||
if self.nsa_enable_prefill_cp:
|
if self.nsa_enable_prefill_cp:
|
||||||
self.cp_rank = get_attention_cp_rank()
|
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):
|
with get_attn_tp_context().maybe_input_scattered(forward_batch):
|
||||||
hidden_states = self.model.forward(
|
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
|
aux_hidden_states = None
|
||||||
if self.capture_aux_hidden_states:
|
if self.capture_aux_hidden_states:
|
||||||
hidden_states, aux_hidden_states = hidden_states
|
hidden_states, aux_hidden_states = hidden_states
|
||||||
@@ -1098,7 +1141,10 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
if is_nextn:
|
if is_nextn:
|
||||||
layers = [self.model.decoder]
|
layers = [self.model.decoder]
|
||||||
else:
|
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:
|
for layer in layers:
|
||||||
attn = layer.self_attn
|
attn = layer.self_attn
|
||||||
G = attn.n_local_groups
|
G = attn.n_local_groups
|
||||||
@@ -1121,7 +1167,8 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
|
|
||||||
if is_nextn:
|
if is_nextn:
|
||||||
return
|
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
|
self_attn = layer.self_attn
|
||||||
if self_attn.compress_ratio != 0 and not self_attn.compressor.ape_converted:
|
if self_attn.compress_ratio != 0 and not self_attn.compressor.ape_converted:
|
||||||
self_attn.compressor.apply_ape_hotfix()
|
self_attn.compressor.apply_ape_hotfix()
|
||||||
@@ -1389,7 +1436,15 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
and not self.pp_group.is_first_rank
|
and not self.pp_group.is_first_rank
|
||||||
):
|
):
|
||||||
continue
|
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
|
continue
|
||||||
elif COMPRESSOR_PART in name:
|
elif COMPRESSOR_PART in name:
|
||||||
is_kv = name.endswith(".wkv.weight")
|
is_kv = name.endswith(".wkv.weight")
|
||||||
@@ -1493,6 +1548,11 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
unloaded_params = params_dict.keys() - loaded_params
|
unloaded_params = params_dict.keys() - loaded_params
|
||||||
|
|
||||||
skipped_checking_patterns = ["attn_mqa.k_scale", "attn_mqa.v_scale"]
|
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:
|
if is_nextn:
|
||||||
skipped_checking_patterns.extend(["lm_head", "embed_tokens"])
|
skipped_checking_patterns.extend(["lm_head", "embed_tokens"])
|
||||||
unloaded_params = {
|
unloaded_params = {
|
||||||
|
|||||||
Reference in New Issue
Block a user