Skip to content
Merged
30 changes: 2 additions & 28 deletions python/sglang/srt/arg_groups/validation_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,6 @@

from sglang.srt.arg_groups.overrides import (
_hisparse_validation,
model_config_of,
resolved_view,
resolving_view,
run_post_process_pass,
Expand All @@ -25,21 +24,6 @@

logger = logging.getLogger(__name__)

_PP_EAGLE_SUPPORTED_ARCHITECTURES = frozenset(
{
"DeepseekV2ForCausalLM",
"DeepseekV3ForCausalLM",
"DeepseekV32ForCausalLM",
"GlmMoeDsaForCausalLM",
# Qwen3.5 (dense / MoE, text and multimodal); folded in from #39602.
"Qwen3_5ForCausalLM",
"Qwen3_5MoeForCausalLM",
"Qwen3_5ForConditionalGeneration",
"Qwen3_5MoeForConditionalGeneration",
"Qwen4ExpForConditionalGeneration",
}
)


def validate_response_store(server_args: Any) -> None:
cfg = resolving_view(server_args)
Expand All @@ -51,9 +35,7 @@ def validate_response_store(server_args: Any) -> None:
)


def check_pipeline_parallel_compat(
cfg: Any, *, model_architecture: Optional[str] = None
) -> None:
def check_pipeline_parallel_compat(cfg: Any) -> None:
"""Validate features used with pipeline parallelism."""
assert cfg.disable_overlap_schedule, (
"Pipeline parallelism is not compatible with overlap schedule"
Expand Down Expand Up @@ -98,11 +80,6 @@ def check_pipeline_parallel_compat(
"PP + speculative decoding (MTP) is only supported on prefill nodes "
"(disaggregation-mode=prefill)"
)
assert model_architecture in _PP_EAGLE_SUPPORTED_ARCHITECTURES, (
"PP + speculative decoding is only supported for DeepSeek/GLM/Qwen3.5 "
"models whose last pipeline stage supplies the EAGLE draft "
f"embedding; got architecture={model_architecture}"
)
assert cfg.min_free_slots_delay is None, (
"--min-free-slots-delay is not supported with pipeline "
"parallelism: allocatable slots per microbatch are bounded by "
Expand Down Expand Up @@ -135,10 +112,7 @@ def check_server_args(server_args: Any):
)

if cfg.pp_size > 1:
model_architecture = None
if cfg.speculative_algorithm is not None:
model_architecture = model_config_of(server_args).hf_config.architectures[0]
check_pipeline_parallel_compat(cfg, model_architecture=model_architecture)
check_pipeline_parallel_compat(cfg)

assert not (cfg.dp_size > 1 and cfg.nnodes != 1 and not cfg.enable_dp_attention), (
"multi-node data parallel is not supported unless dp attention!"
Expand Down
117 changes: 12 additions & 105 deletions python/sglang/srt/speculative/eagle_worker_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,7 @@
prepare_for_draft_extend,
run_eagle_verify,
)
from sglang.srt.speculative.pp_draft_embedding import resolve_draft_embed_and_head
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import (
draft_pp_context,
Expand Down Expand Up @@ -142,85 +143,6 @@
logger = logging.getLogger(__name__)


# Checkpoint spellings of the input embedding across the model families the
# PP+spec gate admits (GLM/DeepSeek NextN, Bailing MTP, Mistral-style drafts).
_EMBED_TENSOR_NAMES = (
"model.embed_tokens.weight",
"embed.weight",
"model.word_embeddings.weight",
"tok_embeddings.weight",
)


def _find_draft_input_embedding(model) -> "torch.nn.Module":
"""The draft's input embedding, found by type rather than attribute path.

Draft models hang it under different names (embed_tokens, word_embeddings,
embed, tok_embeddings), but it is always the one VocabParallelEmbedding
that is not the ParallelLMHead."""
from sglang.srt.layers.vocab_parallel_embedding import (
ParallelLMHead,
VocabParallelEmbedding,
)

found = [
(name, module)
for name, module in model.named_modules()
if isinstance(module, VocabParallelEmbedding)
and not isinstance(module, ParallelLMHead)
]
if len(found) != 1:
raise ValueError(
"PP+spec needs exactly one input embedding on the draft model, "
f"found {[name for name, _ in found]!r}"
)
return found[0][1]


def _load_checkpoint_tensor(
model_path: str, revision, tensor_names: tuple, load_config
) -> torch.Tensor:
"""Load one tensor from a checkpoint via the standard weight loader."""
from sglang.srt.configs.load_config import LoadFormat
from sglang.srt.model_loader.loader import DefaultModelLoader
from sglang.srt.model_loader.weight_utils import (
pt_weights_iterator,
safetensors_weights_iterator,
)

# Streaming and cache-transport formats have no weight files this helper
# could reopen; the dummy format is already skipped by the caller.
reopenable = (
LoadFormat.AUTO,
LoadFormat.SAFETENSORS,
LoadFormat.FASTSAFETENSORS,
LoadFormat.MISTRAL,
LoadFormat.PT,
LoadFormat.NPCACHE,
)
if load_config.load_format not in reopenable:
raise ValueError(
"PP+spec draft embedding loading cannot re-open weights under "
f"load format {load_config.load_format!r}; use a disk-backed "
"load format or disable SGLANG_ENABLE_PP_SPEC"
)
# The target's own load config keeps --download-dir, ignore patterns and
# the selected format, so hub ids resolve into the same cache the model
# was loaded from instead of a fresh default-location download.
_, weight_files, use_safetensors = DefaultModelLoader(load_config)._prepare_weights(
model_path, revision, fall_back_to_pt=True
)
iterator = (
safetensors_weights_iterator(weight_files)
if use_safetensors
else pt_weights_iterator(weight_files)
)
for name, tensor in iterator:
if name in tensor_names:
return tensor
raise ValueError(f"none of {tensor_names} found in checkpoint at {model_path}")


def _qsa_index_share_requested(hf_config) -> bool:
"""--json-model-override-args writes top-level hf_config attributes, while
checkpoint configs carry the flag on the nested text_config; read both."""
Expand Down Expand Up @@ -403,32 +325,7 @@ def init_token_map(self):
def init_lm_head(self):
from sglang.srt.lora.layers import unwrap_lora_layer

if envs.SGLANG_ENABLE_PP_SPEC.get() and get_parallel().pp_size > 1:
# This branch skips the hot-token-map / EAGLE3 head wiring below.
assert self.hot_token_id is None and not (
self.speculative_algorithm.is_eagle3()
), "PP+spec does not support --speculative-token-map or EAGLE3 drafts yet"
# PP+spec: the target's embedding lives on the first PP stage
# (PPMissingLayer here on the last stage) and NextN/MTP layers
# carry no embedding of their own in the checkpoint, so the
# draft's embedding must be loaded from the checkpoint directly
# — otherwise it stays randomly initialized and accept_length
# collapses to ~1.
embed = _find_draft_input_embedding(self.draft_runner.model).weight
if get_model().load_format != "dummy":
target_runner = self.target_worker.model_runner
loaded_embed = _load_checkpoint_tensor(
model_path=target_runner.model_config.model_path,
revision=target_runner.model_config.revision,
tensor_names=_EMBED_TENSOR_NAMES,
load_config=target_runner.load_config,
)
embed.weight_loader(embed, loaded_embed)
head = self.target_worker.model_runner.model.lm_head.weight
self.draft_runner.model.set_embed_and_head(embed, head)
return

embed, head = self.target_worker.model_runner.model.get_embed_and_head()
embed, head = self._resolve_shared_embed_and_head()
target_lm_head = unwrap_lora_layer(
getattr(self.target_worker.model_runner.model, "lm_head", None)
)
Expand Down Expand Up @@ -470,6 +367,16 @@ def maybe_share_target_lm_head():
self.draft_runner.model.set_embed_and_head(embed, head)
maybe_share_target_lm_head()

def _resolve_shared_embed_and_head(self):
target_runner = self.target_worker.model_runner
return resolve_draft_embed_and_head(
target_model=target_runner.model,
draft_model=self.draft_runner.model,
model_path=target_runner.model_config.model_path,
revision=target_runner.model_config.revision,
load_config=target_runner.load_config,
)

def init_attention_backend(self):
# Create multi-step attn backends and cuda graph runners

Expand Down
10 changes: 9 additions & 1 deletion python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,7 @@
set_frozen_kv_positions,
target_kv_pool_view,
)
from sglang.srt.speculative.pp_draft_embedding import resolve_draft_embed_and_head
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import (
draft_pp_context,
Expand Down Expand Up @@ -148,8 +149,15 @@ def __init__(
context_length=self.target_worker.model_runner.model_config.context_len,
)

embed, head = self.target_worker.model_runner.model.get_embed_and_head()
if hasattr(self.draft_model_runner.model, "set_embed_and_head"):
target_runner = self.target_worker.model_runner
embed, head = resolve_draft_embed_and_head(
target_model=target_runner.model,
draft_model=self.draft_model_runner.model,
model_path=target_runner.model_config.model_path,
revision=target_runner.model_config.revision,
load_config=target_runner.load_config,
)
self.draft_model_runner.model.set_embed_and_head(embed, head)
else:
logger.debug(
Expand Down
10 changes: 9 additions & 1 deletion python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,7 @@
rotate_input_ids,
stash_append_boundary_state_triton,
)
from sglang.srt.speculative.pp_draft_embedding import resolve_draft_embed_and_head
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import (
draft_pp_context,
Expand Down Expand Up @@ -367,9 +368,16 @@ def _fill_boundary_kv_front_and_update_stash(
)

def init_lm_head(self):
embed, head = self.target_worker.model_runner.model.get_embed_and_head()
target_runner = self.target_worker.model_runner
# Share the embedding and lm_head
for i in range(self.speculative_num_steps):
embed, head = resolve_draft_embed_and_head(
target_model=target_runner.model,
draft_model=self.draft_runner_list[i].model,
model_path=target_runner.model_config.model_path,
revision=target_runner.model_config.revision,
load_config=target_runner.load_config,
)
self.draft_runner_list[i].model.set_embed_and_head(embed, head)

def init_attention_backend(self):
Expand Down
Loading
Loading