Extract layer-index setup into a module (#31156)

This commit is contained in:
fzyzcjy
2026-07-14 15:59:07 +08:00
committed by GitHub
parent 5b540b16de
commit 39e508b7fc
4 changed files with 257 additions and 140 deletions
@@ -16,6 +16,9 @@ from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import (
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.model_executor.model_runner_components.layer_setup import (
ModelLayerInfo,
)
logger = logging.getLogger(__name__)
@@ -141,9 +144,11 @@ class MlxModelRunnerStub(ModelRunner):
self.model_config.num_hidden_layers,
self.model_config.num_attention_layers,
)
self.start_layer = 0
self.end_layer = model_num_layers
self.num_effective_layers = model_num_layers
self.layer_info = ModelLayerInfo(
start_layer=0,
end_layer=model_num_layers,
num_effective_layers=model_num_layers,
)
# KV cache dtype
self.kv_cache_dtype = self.dtype
+22 -129
View File
@@ -117,6 +117,12 @@ from sglang.srt.model_executor.forward_context import (
)
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
from sglang.srt.model_executor.hook_manager import register_forward_hooks
from sglang.srt.model_executor.model_runner_components.layer_setup import (
ModelLayerInfo,
adjust_hybrid_swa_layer_ids,
compute_attention_and_moe_layers,
resolve_layer_indices,
)
from sglang.srt.model_executor.model_runner_components.load_model_utils import (
build_load_config,
dist_barrier_after_load,
@@ -180,7 +186,6 @@ from sglang.srt.utils import (
get_available_gpu_memory,
get_bool_env_var,
init_cublas,
is_hip,
is_host_cpu_arm64,
is_npu,
log_info_on_rank0,
@@ -202,7 +207,6 @@ from sglang.srt.utils.profile_utils import build_step_span_name
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
from sglang.srt.utils.weight_checker import WeightChecker
_is_hip = is_hip()
_is_npu = is_npu()
_is_cpu_amx_available = cpu_has_amx_support()
_is_cpu_arm64 = is_host_cpu_arm64()
@@ -663,48 +667,19 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.remote_instance_weight_transporter.maybe_register_and_publish_weight_info()
# For MTP models like DeepSeek-V3 or GLM-4.5, the MTP layer(s) are used separately as draft
# models for speculative decoding. In those cases, `num_nextn_predict_layers` is used to
# determine the number of layers.
# Some EAGLE3 drafts (e.g. nvidia/Kimi-K2.5-Thinking-Eagle3) carry the full DeepSeek-V3
# config schema and explicitly set `num_nextn_predict_layers: 0`. Treat that the same as
# the field being absent — otherwise the draft worker takes the MTP branch below with
# model_num_layers=0, sizing the draft KV pool to zero and producing an IndexError on
# the first forward (`set_mla_kv_buffer` -> `self.kv_buffer[layer_id - self.start_layer]`).
_nnpl = self.model_config.num_nextn_predict_layers
model_has_mtp_layers = _nnpl is not None and _nnpl > 0
if self.is_draft_worker and model_has_mtp_layers:
model_num_layers = getattr(
self.model, "num_stages", self.model_config.num_nextn_predict_layers
)
else:
model_num_layers = max(
self.model_config.num_hidden_layers,
self.model_config.num_attention_layers,
)
if self.model_config.hf_config.architectures[0] == "MiMoV2MTP":
model_num_layers = 1
elif self.model_config.hf_config.architectures[0] == "Step3p5MTP":
model_num_layers = 1
self.start_layer = getattr(self.model, "start_layer", 0)
self.end_layer = getattr(self.model, "end_layer", model_num_layers)
self.num_effective_layers = self.end_layer - self.start_layer
self.layer_info: ModelLayerInfo = resolve_layer_indices(
model=self.model,
model_config=self.model_config,
is_draft_worker=self.is_draft_worker,
spec_algorithm=self.spec_algorithm,
)
self.adjust_hybrid_swa_layers_for_pp()
# For LoopCoder models, each loop has its own layer_id, so we need to multiply by loop_num
loop_num = getattr(self.model_config.hf_config, "loop_num", 1)
if loop_num > 1:
self.num_effective_layers = self.num_effective_layers * loop_num
assert (
(not model_has_mtp_layers)
or (self.spec_algorithm.is_none())
or (
(not self.spec_algorithm.is_none())
and (self.num_effective_layers == model_num_layers)
)
), "PP is not compatible with MTP models."
adjust_hybrid_swa_layer_ids(
model_config=self.model_config,
start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer,
is_hybrid_swa=self.is_hybrid_swa,
)
# Apply torchao quantization
torchao_applied = getattr(self.model, "torchao_applied", False)
@@ -735,7 +710,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
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)
or not dsa_layer_skips_topk(hf_config, self.layer_info.start_layer)
):
return None
return getattr(hf_config, "index_topk", None)
@@ -888,28 +863,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
if self.canary_manager is not None and not self.is_draft_worker:
self.canary_manager.mark_init_finished()
def adjust_hybrid_swa_layers_for_pp(self):
if not self.is_hybrid_swa:
return
if self.model_config.is_deepseek_v4_arch:
return
full_attention_layer_ids = [
layer_idx
for layer_idx in range(self.start_layer, self.end_layer + 1)
if hasattr(self.model_config, "full_attention_layer_ids")
and layer_idx in self.model_config.full_attention_layer_ids
]
swa_attention_layer_ids = [
layer_idx
for layer_idx in range(self.start_layer, self.end_layer + 1)
if hasattr(self.model_config, "swa_attention_layer_ids")
and layer_idx in self.model_config.swa_attention_layer_ids
]
self.model_config.swa_attention_layer_ids = swa_attention_layer_ids
self.model_config.full_attention_layer_ids = full_attention_layer_ids
def init_routed_experts_capturer(self):
if self.is_draft_worker:
# Capture is target-only. The draft worker runs in the same process
@@ -1667,69 +1620,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
)
return
self.attention_layers = []
self.moe_layers = []
self.moe_fusions = []
self.dsa_indexers = []
for layer in layer_model.layers:
attn_layer = None
if hasattr(layer, "self_attn"):
if hasattr(layer.self_attn, "attn"):
attn_layer = layer.self_attn.attn
elif hasattr(layer.self_attn, "attn_mqa"):
# For DeepSeek model
attn_layer = layer.self_attn.attn_mqa
if _is_hip and hasattr(layer.self_attn, "attn_mha"):
attn_layer._pcg_mha_companion = layer.self_attn.attn_mha
# For hybrid model
elif hasattr(layer, "attn"):
attn_layer = layer.attn
elif hasattr(layer, "linear_attn"):
if hasattr(layer.linear_attn, "attn"):
attn_layer = layer.linear_attn.attn
else:
attn_layer = layer.linear_attn
# For InternVL model
elif hasattr(layer, "attention"):
if hasattr(layer.attention, "attn"):
attn_layer = layer.attention.attn
# For NemotronH and similar hybrid models using 'mixer' attribute
elif hasattr(layer, "mixer"):
if hasattr(layer.mixer, "attn"):
attn_layer = layer.mixer.attn
elif hasattr(layer, "_forward_mamba"):
# Mamba layer with split op support - store the layer itself
attn_layer = layer
if attn_layer is not None:
self.attention_layers.append(attn_layer)
elif hasattr(layer, "mixer"):
self.attention_layers.append(None)
moe_block = None
moe_fusion = None
if hasattr(layer, "mlp") and hasattr(layer.mlp, "experts"):
moe_block = layer.mlp.experts
moe_fusion = layer.mlp
if hasattr(layer, "block_sparse_moe") and hasattr(
layer.block_sparse_moe, "experts"
):
moe_block = layer.block_sparse_moe.experts
moe_fusion = layer.block_sparse_moe
if hasattr(layer, "moe") and hasattr(layer.moe, "experts"):
moe_block = layer.moe.experts
moe_fusion = layer.moe
# For NemotronH MoE layers using 'mixer' attribute
if hasattr(layer, "mixer") and hasattr(layer.mixer, "experts"):
moe_block = layer.mixer.experts
moe_fusion = layer.mixer
self.moe_layers.append(moe_block)
self.moe_fusions.append(moe_fusion)
# NSA indexers (None for layers without NSA)
dsa_indexer = None
if hasattr(layer, "self_attn") and hasattr(layer.self_attn, "indexer"):
dsa_indexer = layer.self_attn.indexer
self.dsa_indexers.append(dsa_indexer)
self.attention_layers, self.moe_layers, self.moe_fusions, self.dsa_indexers = (
compute_attention_and_moe_layers(layer_model)
)
if len(self.attention_layers) < self.model_config.num_hidden_layers:
# TODO(yuwei): support Non-Standard GQA
@@ -0,0 +1,219 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, NamedTuple
import msgspec
from sglang.srt.utils import is_hip
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
_is_hip = is_hip()
class AttentionAndMoeLayers(NamedTuple):
attention_layers: list[Any]
moe_layers: list[Any]
moe_fusions: list[Any]
dsa_indexers: list[Any]
def compute_attention_and_moe_layers(layer_model: Any) -> AttentionAndMoeLayers:
attention_layers: list[Any] = []
moe_layers: list[Any] = []
moe_fusions: list[Any] = []
dsa_indexers: list[Any] = []
for layer in layer_model.layers:
attn_layer = None
if hasattr(layer, "self_attn"):
if hasattr(layer.self_attn, "attn"):
attn_layer = layer.self_attn.attn
elif hasattr(layer.self_attn, "attn_mqa"):
# For DeepSeek model
attn_layer = layer.self_attn.attn_mqa
if _is_hip and hasattr(layer.self_attn, "attn_mha"):
attn_layer._pcg_mha_companion = layer.self_attn.attn_mha
# For hybrid model
elif hasattr(layer, "attn"):
attn_layer = layer.attn
elif hasattr(layer, "linear_attn"):
if hasattr(layer.linear_attn, "attn"):
attn_layer = layer.linear_attn.attn
else:
attn_layer = layer.linear_attn
# For InternVL model
elif hasattr(layer, "attention"):
if hasattr(layer.attention, "attn"):
attn_layer = layer.attention.attn
# For NemotronH and similar hybrid models using 'mixer' attribute
elif hasattr(layer, "mixer"):
if hasattr(layer.mixer, "attn"):
attn_layer = layer.mixer.attn
elif hasattr(layer, "_forward_mamba"):
# Mamba layer with split op support - store the layer itself
attn_layer = layer
if attn_layer is not None:
attention_layers.append(attn_layer)
elif hasattr(layer, "mixer"):
attention_layers.append(None)
moe_block = None
moe_fusion = None
if hasattr(layer, "mlp") and hasattr(layer.mlp, "experts"):
moe_block = layer.mlp.experts
moe_fusion = layer.mlp
if hasattr(layer, "block_sparse_moe") and hasattr(
layer.block_sparse_moe, "experts"
):
moe_block = layer.block_sparse_moe.experts
moe_fusion = layer.block_sparse_moe
if hasattr(layer, "moe") and hasattr(layer.moe, "experts"):
moe_block = layer.moe.experts
moe_fusion = layer.moe
# For NemotronH MoE layers using 'mixer' attribute
if hasattr(layer, "mixer") and hasattr(layer.mixer, "experts"):
moe_block = layer.mixer.experts
moe_fusion = layer.mixer
moe_layers.append(moe_block)
moe_fusions.append(moe_fusion)
# NSA indexers (None for layers without NSA)
dsa_indexer = None
if hasattr(layer, "self_attn") and hasattr(layer.self_attn, "indexer"):
dsa_indexer = layer.self_attn.indexer
dsa_indexers.append(dsa_indexer)
return AttentionAndMoeLayers(
attention_layers, moe_layers, moe_fusions, dsa_indexers
)
class _PPLayerRange(msgspec.Struct, frozen=True, kw_only=True):
start_layer: int
end_layer: int
class ModelLayerInfo(msgspec.Struct, frozen=True, kw_only=True):
start_layer: int
end_layer: int
num_effective_layers: int
def resolve_layer_indices(
*,
model: Any,
model_config: ModelConfig,
is_draft_worker: bool,
spec_algorithm: SpeculativeAlgorithm,
) -> ModelLayerInfo:
# For MTP models like DeepSeek-V3 or GLM-4.5, the MTP layer(s) are used separately as draft
# models for speculative decoding. In those cases, `num_nextn_predict_layers` is used to
# determine the number of layers.
model_num_layers = _compute_model_num_layers(
model=model, model_config=model_config, is_draft_worker=is_draft_worker
)
_nnpl = model_config.num_nextn_predict_layers
model_has_mtp_layers = _nnpl is not None and _nnpl > 0
pp_range = _resolve_pp_layer_range(model=model, model_num_layers=model_num_layers)
num_effective_layers = pp_range.end_layer - pp_range.start_layer
# For LoopCoder models, each loop has its own layer_id, so we need to multiply by loop_num
loop_num = getattr(model_config.hf_config, "loop_num", 1)
if loop_num > 1:
num_effective_layers = num_effective_layers * loop_num
_assert_pp_mtp_compat(
model_has_mtp_layers=model_has_mtp_layers,
spec_algorithm=spec_algorithm,
num_effective_layers=num_effective_layers,
model_num_layers=model_num_layers,
)
return ModelLayerInfo(
start_layer=pp_range.start_layer,
end_layer=pp_range.end_layer,
num_effective_layers=num_effective_layers,
)
def _compute_model_num_layers(
*,
model: Any,
model_config: ModelConfig,
is_draft_worker: bool,
) -> int:
# Some EAGLE3 drafts (e.g. nvidia/Kimi-K2.5-Thinking-Eagle3) carry the full DeepSeek-V3
# config schema and explicitly set `num_nextn_predict_layers: 0`. Treat that the same as
# the field being absent — otherwise the draft worker takes the MTP branch below with
# model_num_layers=0, sizing the draft KV pool to zero and producing an IndexError on
# the first forward (`set_mla_kv_buffer` -> `self.kv_buffer[layer_id - self.start_layer]`).
_nnpl = model_config.num_nextn_predict_layers
model_has_mtp_layers = _nnpl is not None and _nnpl > 0
model_num_layers = (
getattr(model, "num_stages", model_config.num_nextn_predict_layers)
if is_draft_worker and model_has_mtp_layers
else max(
model_config.num_hidden_layers,
model_config.num_attention_layers,
)
)
if model_config.hf_config.architectures[0] == "MiMoV2MTP":
model_num_layers = 1
elif model_config.hf_config.architectures[0] == "Step3p5MTP":
model_num_layers = 1
return model_num_layers
def _resolve_pp_layer_range(*, model: Any, model_num_layers: int) -> _PPLayerRange:
return _PPLayerRange(
start_layer=getattr(model, "start_layer", 0),
end_layer=getattr(model, "end_layer", model_num_layers),
)
def _assert_pp_mtp_compat(
*,
model_has_mtp_layers: bool,
spec_algorithm: SpeculativeAlgorithm,
num_effective_layers: int,
model_num_layers: int,
) -> None:
assert (
(not model_has_mtp_layers)
or (spec_algorithm.is_none())
or (
(not spec_algorithm.is_none())
and (num_effective_layers == model_num_layers)
)
), "PP is not compatible with MTP models."
def adjust_hybrid_swa_layer_ids(
*,
model_config: ModelConfig,
start_layer: int,
end_layer: int,
is_hybrid_swa: bool,
) -> None:
if not is_hybrid_swa:
return
if model_config.is_deepseek_v4_arch:
return
full_attention_layer_ids = [
layer_idx
for layer_idx in range(start_layer, end_layer + 1)
if hasattr(model_config, "full_attention_layer_ids")
and layer_idx in model_config.full_attention_layer_ids
]
swa_attention_layer_ids = [
layer_idx
for layer_idx in range(start_layer, end_layer + 1)
if hasattr(model_config, "swa_attention_layer_ids")
and layer_idx in model_config.swa_attention_layer_ids
]
model_config.swa_attention_layer_ids = swa_attention_layer_ids
model_config.full_attention_layer_ids = full_attention_layer_ids
@@ -485,12 +485,12 @@ class ModelRunnerKVCacheMixin:
mamba_layer_ids = [
i
for i in config.mamba2_cache_params.layers
if self.start_layer <= i < self.end_layer
if self.layer_info.start_layer <= i < self.layer_info.end_layer
]
full_attention_layer_ids = [
i
for i in config.full_attention_layer_ids
if self.start_layer <= i < self.end_layer
if self.layer_info.start_layer <= i < self.layer_info.end_layer
]
bundle = init_unified_mamba_pools(
@@ -499,8 +499,8 @@ class ModelRunnerKVCacheMixin:
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
head_dim=self.model_config.head_dim,
page_size=self.page_size,
start_layer=self.start_layer,
end_layer=self.end_layer,
start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer,
is_draft_worker=self.is_draft_worker,
use_mla_backend=self.use_mla_backend,
mamba_layer_ids=mamba_layer_ids,
@@ -576,12 +576,12 @@ class ModelRunnerKVCacheMixin:
swa_attention_layer_ids = [
i
for i in self.model_config.swa_attention_layer_ids
if self.start_layer <= i < self.end_layer
if self.layer_info.start_layer <= i < self.layer_info.end_layer
]
full_attention_layer_ids = [
i
for i in self.model_config.full_attention_layer_ids
if self.start_layer <= i < self.end_layer
if self.layer_info.start_layer <= i < self.layer_info.end_layer
]
bundle = init_unified_swa_pools(
@@ -594,8 +594,8 @@ class ModelRunnerKVCacheMixin:
swa_head_dim=swa_head_dim,
swa_v_head_dim=swa_v_head_dim,
page_size=self.page_size,
start_layer=self.start_layer,
end_layer=self.end_layer,
start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer,
swa_attention_layer_ids=swa_attention_layer_ids,
full_attention_layer_ids=full_attention_layer_ids,
full_max_total_num_tokens=self.full_max_total_num_tokens,