Fix IndexCache PP topk handoff (#28532)

This commit is contained in:
Mohammad Miadh Angkad
2026-06-19 23:39:09 +08:00
committed by GitHub
parent 3af991fb3e
commit 88c261c3f3
7 changed files with 99 additions and 43 deletions
+23
View File
@@ -132,6 +132,29 @@ def get_dsa_index_topk(config: PretrainedConfig) -> int:
return config.index_topk
def dsa_layer_skips_topk(config: PretrainedConfig, layer_id: int) -> bool:
"""Return whether a DSA layer reuses the previous layer's top-k indices."""
assert is_deepseek_dsa(config)
pattern = getattr(config, "index_topk_pattern", None)
if pattern is not None:
return layer_id < len(pattern) and pattern[layer_id] == "S"
freq = getattr(config, "index_topk_freq", 1)
if freq is None:
freq = 1
assert freq > 0, f"index_topk_freq must be positive, got {freq}"
offset = getattr(config, "index_skip_topk_offset", None)
if offset is not None:
assert offset > 0, (
"index_skip_topk_offset must be positive; offset <= 0 "
"marks layer 0 as skip_topk with no prior topk to reuse"
)
return max(layer_id - offset + 1, 0) % freq != 0
return max(layer_id - 1, 0) % freq != 0
def get_dsa_index_n_heads(config: PretrainedConfig) -> int:
assert is_deepseek_dsa(config)
return config.index_n_heads
@@ -650,6 +650,13 @@ class SchedulerPPMixin:
device=self.device,
),
}
pp_proxy_topk_size = model_runner.get_pp_proxy_topk_size()
if pp_proxy_topk_size is not None:
proxy_tensors["topk_indices"] = torch.zeros(
(current_seq_len, pp_proxy_topk_size),
dtype=torch.int32,
device=self.device,
)
pp_proxy = PPProxyTensors(proxy_tensors)
@@ -61,7 +61,9 @@ from sglang.srt.configs.model_config import (
AttentionArch,
ModelConfig,
ModelImpl,
dsa_layer_skips_topk,
get_num_indexer_layers,
is_deepseek_dsa,
)
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS
@@ -839,6 +841,17 @@ class ModelRunner(ModelRunnerKVCacheMixin):
cpu_group=get_world_group().cpu_group,
)
def get_pp_proxy_topk_size(self) -> Optional[int]:
hf_config = self.model_config.hf_text_config
if (
self.pp_size <= 1
or self.pp_rank == 0
or not is_deepseek_dsa(hf_config)
or not dsa_layer_skips_topk(hf_config, self.start_layer)
):
return None
return getattr(hf_config, "index_topk", None)
def alloc_memory_pool(self, memory_pool_config: Optional[MemoryPoolConfig] = None):
"""Allocate KV cache memory pools only (no backends or cuda graphs)."""
if memory_pool_config is not None:
@@ -2758,6 +2771,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
cache_loc_dtype=torch.int64,
enable_mamba_track=False,
hc_hidden_size=getattr(self.model_config, "hc_hidden_size", None),
pp_proxy_topk_size=self.get_pp_proxy_topk_size(),
)
buffers.num_token_non_padded[...] = num_tokens
@@ -182,6 +182,7 @@ def _allocate_decode_buffers(
enable_mamba_track: bool,
ne_token_table: Optional[torch.Tensor] = None,
hc_hidden_size: Optional[int] = None,
pp_proxy_topk_size: Optional[int] = None,
) -> SimpleNamespace:
"""Allocate the FB-shared decode buffers as a namespace adopted by
``build_decode_registry(source=...)``."""
@@ -220,6 +221,10 @@ def _allocate_decode_buffers(
pp_proxy_tensors["residual"] = torch.zeros(
(max_bs, hidden_size), dtype=dtype
)
if pp_proxy_topk_size is not None:
pp_proxy_tensors["topk_indices"] = torch.zeros(
(max_num_token, pp_proxy_topk_size), dtype=torch.int32
)
else:
pp_proxy_tensors = None
@@ -450,6 +455,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
hc_hidden_size=getattr(
self.model_runner.model_config, "hc_hidden_size", None
),
pp_proxy_topk_size=self.model_runner.get_pp_proxy_topk_size(),
)
self.buffers.share_buffers()
# FB-shared slot registry adopting DecodeInputBuffers storage (same
@@ -107,6 +107,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
ne_token_table: Optional[torch.Tensor] = None,
is_hybrid_swa: bool = False,
hc_hidden_size: Optional[int] = None,
pp_proxy_topk_size: Optional[int] = None,
) -> DecodeInputBuffers:
with torch.device(device):
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
@@ -149,6 +150,10 @@ class DecodeInputBuffers(ForwardInputBuffers):
pp_proxy_tensors["residual"] = torch.zeros(
(max_bs, hidden_size), dtype=dtype
)
if pp_proxy_topk_size is not None:
pp_proxy_tensors["topk_indices"] = torch.zeros(
(max_num_token, pp_proxy_topk_size), dtype=torch.int32
)
else:
pp_proxy_tensors = None
+43 -42
View File
@@ -40,6 +40,7 @@ from sglang.srt.batch_overlap.two_batch_overlap import (
)
from sglang.srt.configs.model_config import (
compute_mla_mscale_scaling,
dsa_layer_skips_topk,
get_dsa_index_head_dim,
get_dsa_index_n_heads,
get_dsa_index_topk,
@@ -1604,40 +1605,8 @@ class DeepseekV2AttentionMLA(
self.skip_topk = True
self.next_skip_topk = True
else:
self.index_topk_freq = getattr(config, "index_topk_freq", 1)
self.index_topk_pattern = getattr(config, "index_topk_pattern", None)
self.index_skip_topk_offset = getattr(
config, "index_skip_topk_offset", None
)
if (
self.index_topk_pattern is None
and self.index_skip_topk_offset is not None
):
assert self.index_skip_topk_offset > 0, (
"index_skip_topk_offset must be positive; offset <= 0 "
"marks layer 0 as skip_topk with no prior topk to reuse"
)
self.skip_topk = (
max(layer_id - self.index_skip_topk_offset + 1, 0)
% self.index_topk_freq
!= 0
)
self.next_skip_topk = (
max(layer_id - self.index_skip_topk_offset + 2, 0)
% self.index_topk_freq
!= 0
)
elif self.index_topk_pattern is None:
self.skip_topk = max(layer_id - 1, 0) % self.index_topk_freq != 0
self.next_skip_topk = layer_id % self.index_topk_freq != 0
else:
self.skip_topk = self.index_topk_pattern[layer_id] == "S"
if layer_id < len(self.index_topk_pattern) - 1:
self.next_skip_topk = (
self.index_topk_pattern[layer_id + 1] == "S"
)
else:
self.next_skip_topk = False
self.skip_topk = dsa_layer_skips_topk(config, layer_id)
self.next_skip_topk = dsa_layer_skips_topk(config, layer_id + 1)
self.kv_b_proj = ColumnParallelLinear(
self.kv_lora_rank,
@@ -2290,13 +2259,15 @@ class DeepseekV2Model(nn.Module):
prefix: str = "",
) -> None:
super().__init__()
self.config = config
self.use_dsa = is_deepseek_dsa(config)
self.padding_id = config.pad_token_id
self.vocab_size = config.vocab_size
self.first_k_dense_replace = config.first_k_dense_replace
self.pp_group = get_pp_group()
self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp()
self.mla_enable_prefill_cp = (
is_prefill_context_parallel_enabled() and not is_deepseek_dsa(config)
is_prefill_context_parallel_enabled() and not self.use_dsa
)
if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp:
self.cp_size = get_parallel().attn_cp_size
@@ -2435,6 +2406,17 @@ class DeepseekV2Model(nn.Module):
assert pp_proxy_tensors is not None
hidden_states = pp_proxy_tensors["hidden_states"]
residual = pp_proxy_tensors["residual"]
topk_indices = pp_proxy_tensors.tensors.get("topk_indices")
assert not (
not forward_batch.forward_mode.is_idle()
and hidden_states.shape[0] != 0
and self.use_dsa
and dsa_layer_skips_topk(self.config, self.start_layer)
and topk_indices is None
), (
f"PP stage starting at layer {self.start_layer} requires DSA "
"topk_indices from the previous stage."
)
device = hidden_states.device
zero_allocator = BumpAllocator(
buffer_size=total_num_layers * 2 * (2 if forward_batch.can_run_tbo else 1),
@@ -2487,7 +2469,8 @@ class DeepseekV2Model(nn.Module):
elif self.first_k_dense_replace < normal_start_layer:
normal_end_layer = normal_start_layer = 0
aux_hidden_states = []
topk_indices = None
if self.pp_group.is_first_rank:
topk_indices = None
for i in range(normal_start_layer, normal_end_layer):
# NOTE: torch dynamo does not support graph break in context manager
ctx = (
@@ -2526,12 +2509,30 @@ class DeepseekV2Model(nn.Module):
)
if not self.pp_group.is_last_rank:
return PPProxyTensors(
{
"hidden_states": hidden_states,
"residual": residual,
}
)
proxy_tensors = {
"hidden_states": hidden_states,
"residual": residual,
}
if (
self.use_dsa
and self.end_layer < self.config.num_hidden_layers
and dsa_layer_skips_topk(self.config, self.end_layer)
):
if (
not forward_batch.forward_mode.is_idle()
and hidden_states.shape[0] != 0
):
assert topk_indices is not None, (
f"PP stage ending at layer {self.end_layer} must forward "
"DSA topk_indices because the next stage starts on a "
"skip-topk layer."
)
if topk_indices is None:
topk_indices = hidden_states.new_empty(
(0, get_dsa_index_topk(self.config)), dtype=torch.int32
)
proxy_tensors["topk_indices"] = topk_indices
return PPProxyTensors(proxy_tensors)
else:
if not forward_batch.forward_mode.is_idle():
if residual is None:
+1 -1
View File
@@ -2007,7 +2007,7 @@ class ServerArgs:
self.attention_backend = "dsa"
logger.info("Use dsa attention backend for DeepSeek with DSA.")
index_topk_freq = getattr(hf_config, "index_topk_freq", 1)
index_topk_freq = getattr(hf_config, "index_topk_freq", 1) or 1
index_topk_pattern = getattr(hf_config, "index_topk_pattern", None)
if self.enable_two_batch_overlap and (
index_topk_freq > 1