diff --git a/python/sglang/srt/arg_groups/validation_hook.py b/python/sglang/srt/arg_groups/validation_hook.py index 57cd08d73499..5ecf25a0bd33 100644 --- a/python/sglang/srt/arg_groups/validation_hook.py +++ b/python/sglang/srt/arg_groups/validation_hook.py @@ -10,7 +10,6 @@ from sglang.srt.arg_groups.overrides import ( _hisparse_validation, - model_config_of, resolved_view, resolving_view, run_post_process_pass, @@ -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) @@ -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" @@ -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 " @@ -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!" diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 8fe05442c3f3..c8d58e5f5e11 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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, @@ -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.""" @@ -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) ) @@ -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 diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py index bf31d835cc16..fa73bf20712e 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py @@ -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, @@ -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( diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index 1a9d74be5044..753c1f082a48 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -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, @@ -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): diff --git a/python/sglang/srt/speculative/pp_draft_embedding.py b/python/sglang/srt/speculative/pp_draft_embedding.py new file mode 100644 index 000000000000..f13a5125512a --- /dev/null +++ b/python/sglang/srt/speculative/pp_draft_embedding.py @@ -0,0 +1,241 @@ +"""Model-agnostic draft embedding for speculative decoding under pipeline parallelism. + +The draft runs only on the last stage while the target embedding lives only on the +first stage, so the draft cannot share ``target.get_embed_and_head()``. This module +reads the target embed/head PP-safely and loads the checkpoint's input embedding +into the draft's own embedding parameter, touching only the shard that holds it. +""" + +from __future__ import annotations + +import json +import logging +import os +import re +from typing import Iterable, List, Optional, Tuple + +import torch +from torch import nn + +from sglang.srt.configs.load_config import LoadConfig, LoadFormat +from sglang.srt.layers.utils.common import PPMissingLayer +from sglang.srt.model_loader.weight_utils import default_weight_loader + +logger = logging.getLogger(__name__) + +SAFETENSORS_INDEX_NAME = "model.safetensors.index.json" + +# Module attribute names under which model families hang the input embedding. +_EMBED_ATTR_NAMES: Tuple[str, ...] = ( + "embed_tokens", + "word_embeddings", + "tok_embeddings", + "embed", +) +# Checkpoint spellings of the input embedding, most common first. +EMBED_KEY_CANDIDATES: Tuple[str, ...] = ( + "model.embed_tokens.weight", + "model.language_model.embed_tokens.weight", + "language_model.model.embed_tokens.weight", + "model.word_embeddings.weight", + "tok_embeddings.weight", + "embed_tokens.weight", + "embed.weight", +) +_EMBED_KEY_SUFFIXES: Tuple[str, ...] = tuple(f"{n}.weight" for n in _EMBED_ATTR_NAMES) +# MTP / NextN layers carry their own embedding under ``layers..``; never pick it. +_LAYER_KEY_RE = re.compile(r"(^|\.)layers\.\d+\.") +_DRAFT_SUBMODULE_MARKERS: Tuple[str, ...] = ("mtp", "nextn", "eagle", "draft") +# Streaming and cache-transport formats have no weight files to re-open. +_REOPENABLE_LOAD_FORMATS = ( + LoadFormat.AUTO, + LoadFormat.SAFETENSORS, + LoadFormat.FASTSAFETENSORS, + LoadFormat.MISTRAL, + LoadFormat.PT, + LoadFormat.NPCACHE, +) + + +def _target_input_embedding_is_missing(target_model: nn.Module) -> bool: + """True when the target's input embedding on this stage is a ``PPMissingLayer``.""" + for name, module in target_model.named_modules(): + if name.rsplit(".", 1)[-1] in _EMBED_ATTR_NAMES and isinstance( + module, PPMissingLayer + ): + return True + return False + + +def _weight_or_none(module: Optional[nn.Module]) -> Optional[torch.Tensor]: + if module is None or isinstance(module, PPMissingLayer): + return None + from sglang.srt.lora.layers import unwrap_lora_layer + + return unwrap_lora_layer(module).weight + + +def resolve_target_embed_and_head( + target_model: nn.Module, +) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]: + """PP-safe getter: a PPMissingLayer embedding maps to ``embed=None``; any + other AttributeError propagates.""" + try: + return target_model.get_embed_and_head() + except AttributeError: + if not _target_input_embedding_is_missing(target_model): + raise + return None, _weight_or_none(target_model.lm_head) + + +def find_draft_embedding_param( + draft_model: nn.Module, +) -> Optional[Tuple[str, nn.Parameter]]: + """The draft's own input embedding: its single ``VocabParallelEmbedding``.""" + from sglang.srt.layers.vocab_parallel_embedding import ( + ParallelLMHead, + VocabParallelEmbedding, + ) + + found = [ + (f"{name}.weight", module.weight) + for name, module in draft_model.named_modules() + if isinstance(module, VocabParallelEmbedding) + and not isinstance(module, ParallelLMHead) + ] + return found[0] if len(found) == 1 else None + + +def _is_input_embedding_key(key: str) -> bool: + if not key.endswith(_EMBED_KEY_SUFFIXES) or _LAYER_KEY_RE.search(key): + return False + lowered = key.lower() + return not any(marker in lowered for marker in _DRAFT_SUBMODULE_MARKERS) + + +def _pick_embedding_key(keys: Iterable[str]) -> Optional[str]: + keys = set(keys) + for candidate in EMBED_KEY_CANDIDATES: + if candidate in keys: + return candidate + fallback = sorted((k for k in keys if _is_input_embedding_key(k)), key=len) + return fallback[0] if fallback else None + + +def _safetensors_keys(path: str) -> List[str]: + import safetensors + + with safetensors.safe_open(path, framework="pt", device="cpu") as f: + return list(f.keys()) + + +def _read_safetensors_tensor(path: str, key: str) -> torch.Tensor: + import safetensors + + with safetensors.safe_open(path, framework="pt", device="cpu") as f: + return f.get_tensor(key) + + +def prepare_checkpoint_files( + model_path: str, *, revision: Optional[str], load_config: LoadConfig +) -> Tuple[str, List[str], bool]: + """``(folder, weight_files, use_safetensors)`` via the standard model loader, + so ModelScope resolution, ``--download-dir`` and ``load_format`` are honored.""" + from sglang.srt.model_loader.loader import DefaultModelLoader + + if load_config.load_format not in _REOPENABLE_LOAD_FORMATS: + raise ValueError( + "Pipeline-parallel speculative decoding needs to re-open the target " + f"checkpoint for the draft embedding, which load format " + f"{load_config.load_format!r} does not allow; use a disk-backed format." + ) + return DefaultModelLoader(load_config)._prepare_weights( + model_path, revision, fall_back_to_pt=True + ) + + +def load_embedding_tensor( + folder: str, weight_files: List[str], *, use_safetensors: bool +) -> Tuple[str, torch.Tensor]: + """``(key, tensor)`` of the checkpoint's input embedding; with an index only + the owning shard is read, otherwise the key is picked over all shard headers.""" + if use_safetensors: + index_file = os.path.join(folder, SAFETENSORS_INDEX_NAME) + if os.path.exists(index_file): + with open(index_file) as f: + weight_map = json.load(f).get("weight_map", {}) or {} + key = _pick_embedding_key(weight_map.keys()) + if key is not None: + shard = os.path.join(folder, weight_map[key]) + return key, _read_safetensors_tensor(shard, key) + keys_by_file = {path: _safetensors_keys(path) for path in weight_files} + key = _pick_embedding_key(k for keys in keys_by_file.values() for k in keys) + if key is not None: + shard = next(p for p, keys in keys_by_file.items() if key in keys) + return key, _read_safetensors_tensor(shard, key) + else: + from sglang.srt.model_loader.weight_utils import pt_weights_iterator + + for name, tensor in pt_weights_iterator(weight_files): + if name in EMBED_KEY_CANDIDATES or _is_input_embedding_key(name): + return name, tensor + raise ValueError( + f"No input embedding found in checkpoint under {folder}; looked for " + f"{EMBED_KEY_CANDIDATES} or '*.<{'|'.join(_EMBED_ATTR_NAMES)}>.weight'." + ) + + +def load_draft_embedding_from_checkpoint( + draft_model: nn.Module, + model_path: str, + *, + revision: Optional[str], + load_config: LoadConfig, +) -> nn.Parameter: + """Load the checkpoint's input embedding into the draft's own parameter and + return it; the parameter's ``weight_loader`` applies the TP shard.""" + found = find_draft_embedding_param(draft_model) + if found is None: + raise ValueError( + f"Draft model {draft_model.__class__.__name__} has no single " + "VocabParallelEmbedding, and the target cannot share its embedding from " + "this pipeline stage (https://github.com/sgl-project/sglang/issues/39634)." + ) + param_name, param = found + if load_config.load_format == LoadFormat.DUMMY: + return param + + folder, weight_files, use_safetensors = prepare_checkpoint_files( + model_path, revision=revision, load_config=load_config + ) + key, loaded = load_embedding_tensor( + folder, weight_files, use_safetensors=use_safetensors + ) + weight_loader = getattr(param, "weight_loader", default_weight_loader) + weight_loader(param, loaded) + logger.info( + "Loaded draft embedding %s from checkpoint key %s %s for pipeline-parallel " + "speculative decoding.", + param_name, + key, + tuple(loaded.shape), + ) + return param + + +def resolve_draft_embed_and_head( + *, + target_model: nn.Module, + draft_model: nn.Module, + model_path: str, + revision: Optional[str], + load_config: LoadConfig, +) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]: + """Embed/head to bind into a draft; a stage without the target embedding + loads the draft's own from the checkpoint.""" + embed, head = resolve_target_embed_and_head(target_model) + if embed is None: + embed = load_draft_embedding_from_checkpoint( + draft_model, model_path, revision=revision, load_config=load_config + ) + return embed, head diff --git a/test/registered/e2e/pp/test_pp_spec_embed_scan.py b/test/registered/e2e/pp/test_pp_spec_embed_scan.py index f42b5f85fac5..5dcdfd312b9f 100644 --- a/test/registered/e2e/pp/test_pp_spec_embed_scan.py +++ b/test/registered/e2e/pp/test_pp_spec_embed_scan.py @@ -88,16 +88,21 @@ def setUpClass(cls): def test_bailing_nextn_embedding_is_found(self): from sglang.srt.models.bailing_moe_nextn import BailingMoeForCausalLMNextN - from sglang.srt.speculative.eagle_worker_v2 import _find_draft_input_embedding + from sglang.srt.speculative.pp_draft_embedding import ( + find_draft_embedding_param, + ) config = PretrainedConfig.from_dict(dict(_BAILING_CONFIG)) with torch.device("cuda"): model = BailingMoeForCausalLMNextN(config) - self.assertIs(_find_draft_input_embedding(model), model.model.word_embeddings) + _, param = find_draft_embedding_param(model) + self.assertIs(param, model.model.word_embeddings.weight) def test_mistral_eagle_embedding_is_found(self): from sglang.srt.models.mistral_eagle import MistralForCausalLMEagle - from sglang.srt.speculative.eagle_worker_v2 import _find_draft_input_embedding + from sglang.srt.speculative.pp_draft_embedding import ( + find_draft_embedding_param, + ) config = MistralConfig( vocab_size=1024, @@ -109,16 +114,17 @@ def test_mistral_eagle_embedding_is_found(self): ) with torch.device("cuda"): model = MistralForCausalLMEagle(config) - self.assertIs(_find_draft_input_embedding(model), model.model.embed_tokens) + _, param = find_draft_embedding_param(model) + self.assertIs(param, model.model.embed_tokens.weight) def test_name_table_matches_published_checkpoints(self): - from sglang.srt.speculative.eagle_worker_v2 import _EMBED_TENSOR_NAMES + from sglang.srt.speculative.pp_draft_embedding import EMBED_KEY_CANDIDATES # External-source literals: the embedding tensor's spelling in each # family's published target checkpoint (inclusionAI/Ling-*-2.0 # model.safetensors.index.json; Mistral-Large-3 consolidated index). - self.assertIn("model.word_embeddings.weight", _EMBED_TENSOR_NAMES) - self.assertIn("tok_embeddings.weight", _EMBED_TENSOR_NAMES) + self.assertIn("model.word_embeddings.weight", EMBED_KEY_CANDIDATES) + self.assertIn("tok_embeddings.weight", EMBED_KEY_CANDIDATES) if __name__ == "__main__": diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 1163a411d86c..2426a8193cc9 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -2480,8 +2480,6 @@ def test_tc_piecewise_build_config_reads_phase_config_dataclass( class TestPipelineParallelCompat(CustomTestCase): """Features supported with `pipeline-parallel-size > 1`.""" - _SUPPORTED_ARCH = "GlmMoeDsaForCausalLM" - @staticmethod def _cfg(**overrides): cfg = dict( @@ -2524,10 +2522,7 @@ def test_dspark_rejects_eagle_pp_relay(self): ) def test_eagle_is_allowed_on_prefill(self): - check_pipeline_parallel_compat( - self._cfg(speculative_algorithm="EAGLE"), - model_architecture=self._SUPPORTED_ARCH, - ) + check_pipeline_parallel_compat(self._cfg(speculative_algorithm="EAGLE")) def test_eagle_is_rejected_outside_prefill(self): for mode in ("decode", "null"): @@ -2536,35 +2531,9 @@ def test_eagle_is_rejected_outside_prefill(self): check_pipeline_parallel_compat( self._cfg( speculative_algorithm="EAGLE", disaggregation_mode=mode - ), - model_architecture=self._SUPPORTED_ARCH, + ) ) - def test_eagle_is_rejected_for_unsupported_model(self): - with self.assertRaisesRegex(AssertionError, "DeepSeek/GLM/Qwen3.5 models"): - check_pipeline_parallel_compat( - self._cfg(speculative_algorithm="EAGLE"), - model_architecture="LlamaForCausalLM", - ) - - def test_supported_architectures(self): - for architecture in ( - "DeepseekV2ForCausalLM", - "DeepseekV3ForCausalLM", - "DeepseekV32ForCausalLM", - "GlmMoeDsaForCausalLM", - "Qwen3_5ForCausalLM", - "Qwen3_5MoeForCausalLM", - "Qwen3_5ForConditionalGeneration", - "Qwen3_5MoeForConditionalGeneration", - "Qwen4ExpForConditionalGeneration", - ): - with self.subTest(architecture=architecture): - check_pipeline_parallel_compat( - self._cfg(speculative_algorithm="EAGLE"), - model_architecture=architecture, - ) - def test_pp_spec_env_gate_allows_aggregate_and_rejects_pd(self): cfg = self._cfg( speculative_algorithm="EAGLE", @@ -2575,33 +2544,23 @@ def test_pp_spec_env_gate_allows_aggregate_and_rejects_pd(self): with patch.object( validation_hook.envs.SGLANG_ENABLE_PP_SPEC, "get", return_value=True ): - check_pipeline_parallel_compat(cfg, model_architecture="LlamaForCausalLM") + check_pipeline_parallel_compat(cfg) with self.assertRaisesRegex(AssertionError, "SGLANG_ENABLE_PP_SPEC"): - check_pipeline_parallel_compat( - self._cfg(speculative_algorithm="EAGLE"), - model_architecture=self._SUPPORTED_ARCH, - ) + check_pipeline_parallel_compat(self._cfg(speculative_algorithm="EAGLE")) def test_nextn_resolves_to_eagle_and_is_allowed(self): """`--speculative-algorithm NEXTN` has collapsed to EAGLE by the time the validation hook runs, so the check only ever sees the resolved name.""" - check_pipeline_parallel_compat( - self._cfg(speculative_algorithm="eagle"), - model_architecture=self._SUPPORTED_ARCH, - ) + check_pipeline_parallel_compat(self._cfg(speculative_algorithm="eagle")) def test_non_eagle_speculative_algorithms_are_rejected(self): with self.assertRaisesRegex(AssertionError, "only supports EAGLE"): - check_pipeline_parallel_compat( - self._cfg(speculative_algorithm="EAGLE3"), - model_architecture=self._SUPPORTED_ARCH, - ) + check_pipeline_parallel_compat(self._cfg(speculative_algorithm="EAGLE3")) def test_multi_layer_eagle_is_rejected(self): with self.assertRaisesRegex(AssertionError, "only supports EAGLE"): check_pipeline_parallel_compat( - self._cfg(speculative_algorithm="EAGLE", enable_multi_layer_eagle=True), - model_architecture=self._SUPPORTED_ARCH, + self._cfg(speculative_algorithm="EAGLE", enable_multi_layer_eagle=True) ) def test_min_free_slots_delay_is_rejected(self): diff --git a/test/registered/unit/spec/test_pp_draft_embedding.py b/test/registered/unit/spec/test_pp_draft_embedding.py new file mode 100644 index 000000000000..2ab8271f68d6 --- /dev/null +++ b/test/registered/unit/spec/test_pp_draft_embedding.py @@ -0,0 +1,198 @@ +"""Unit tests for srt/speculative/pp_draft_embedding.""" + +import json +import os +import tempfile +import unittest + +import torch +from safetensors.torch import save_file +from torch import nn + +from sglang.srt.configs.load_config import LoadConfig, LoadFormat +from sglang.srt.layers.utils.common import PPMissingLayer +from sglang.srt.runtime_context import get_context +from sglang.srt.speculative.pp_draft_embedding import ( + load_draft_embedding_from_checkpoint, + resolve_target_embed_and_head, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import ( + CustomTestCase, + enter_scope, + maybe_stub_sgl_kernel, + published_topology, +) + +maybe_stub_sgl_kernel() + +from sglang.srt.layers.vocab_parallel_embedding import ( # noqa: E402 + VocabParallelEmbedding, +) + +register_cpu_ci(est_time=8, suite="base-a-test-cpu") + +VOCAB, HIDDEN = 16, 8 +AUTO = LoadConfig(load_format=LoadFormat.AUTO) +MAIN_KEY = "model.embed_tokens.weight" +# A DeepSeek/GLM-style MTP layer embedding: no mtp/nextn marker in the key. +MTP_LAYER_KEY = "model.layers.61.embed_tokens.weight" + + +class _Inner(nn.Module): + def __init__(self, embed: nn.Module): + super().__init__() + self.embed_tokens = embed + + +class _Target(nn.Module): + """Target with the common getter shape: ``self.model.embed_tokens.weight``.""" + + def __init__(self, *, owns_embedding: bool): + super().__init__() + self.model = _Inner( + nn.Embedding(VOCAB, HIDDEN) if owns_embedding else PPMissingLayer() + ) + self.lm_head = nn.Linear(HIDDEN, VOCAB, bias=False) + + def get_embed_and_head(self): + return self.model.embed_tokens.weight, self.lm_head.weight + + +class _BuggyGetterTarget(_Target): + def get_embed_and_head(self): + return self.model.embed_tokens.weight, self.lm_haed.weight # typo on purpose + + +class _Draft(nn.Module): + def __init__(self, *, with_embedding: bool = True): + super().__init__() + self.model = _Inner( + VocabParallelEmbedding(VOCAB, HIDDEN) if with_embedding else nn.Identity() + ) + + +def _write_sharded_checkpoint(root: str, embed: torch.Tensor) -> None: + """Two shards + index; the MTP-layer decoy sorts before the real embedding.""" + decoy = torch.full_like(embed, -1.0) + shard_a, shard_b = ( + "model-00001-of-00002.safetensors", + "model-00002-of-00002.safetensors", + ) + save_file({MTP_LAYER_KEY: decoy}, os.path.join(root, shard_a)) + save_file({MAIN_KEY: embed, "lm_head.weight": decoy}, os.path.join(root, shard_b)) + weight_map = {MTP_LAYER_KEY: shard_a, MAIN_KEY: shard_b, "lm_head.weight": shard_b} + with open(os.path.join(root, "model.safetensors.index.json"), "w") as f: + json.dump({"weight_map": weight_map}, f) + + +class TestResolveTargetEmbedAndHead(CustomTestCase): + def test_non_first_stage_returns_head_without_embed(self): + """Regression: a PPMissingLayer embed_tokens used to raise AttributeError and + abort draft init on the last stage; the head must still be returned.""" + target = _Target(owns_embedding=False) + embed, head = resolve_target_embed_and_head(target) + self.assertIsNone(embed) + self.assertIs(head, target.lm_head.weight) + + def test_unrelated_attribute_error_is_not_swallowed(self): + """A getter bug on a stage that owns its embedding must propagate rather + than silently redirect to checkpoint loading.""" + target = _BuggyGetterTarget(owns_embedding=True) + with self.assertRaises(AttributeError): + resolve_target_embed_and_head(target) + + +class TestLoadDraftEmbeddingFromCheckpoint(CustomTestCase): + @classmethod + def setUpClass(cls): + # The model loader reads the published context once (checksum check). + cls._override = get_context().override_server_args() + cls._override.install() + + @classmethod + def tearDownClass(cls): + cls._override.restore() + + def setUp(self): + enter_scope(self, published_topology("test", tp_size=1, pp_size=1)) + + def _load(self, draft, root, load_config=AUTO): + return load_draft_embedding_from_checkpoint( + draft, root, revision=None, load_config=load_config + ) + + def _assert_loaded(self, param, expected): + # VocabParallelEmbedding pads the vocab; only the real rows are loaded. + torch.testing.assert_close(param.detach()[:VOCAB], expected) + + def test_index_picks_input_embedding_over_mtp_layer_key(self): + expected = torch.randn(VOCAB, HIDDEN) + draft = _Draft() + with tempfile.TemporaryDirectory() as root: + _write_sharded_checkpoint(root, expected) + param = self._load(draft, root) + self.assertIs(param, draft.model.embed_tokens.weight) + self._assert_loaded(param, expected) + + def test_no_index_picks_over_whole_checkpoint_not_first_shard(self): + """Regression: the first shard holding *an* embedding-like key was chosen; + an MTP-layer embedding in an earlier shard must lose to the real one.""" + expected = torch.randn(VOCAB, HIDDEN) + with tempfile.TemporaryDirectory() as root: + save_file( + {MTP_LAYER_KEY: torch.zeros(VOCAB, HIDDEN)}, + os.path.join(root, "a.safetensors"), + ) + save_file({MAIN_KEY: expected}, os.path.join(root, "b.safetensors")) + param = self._load(_Draft(), root) + self._assert_loaded(param, expected) + + def test_bin_only_checkpoint_follows_loader_fallback(self): + """Regression: a checkpoint with only pytorch_model.bin raised + FileNotFoundError because only safetensors were searched.""" + expected = torch.randn(VOCAB, HIDDEN) + with tempfile.TemporaryDirectory() as root: + torch.save({MAIN_KEY: expected}, os.path.join(root, "pytorch_model.bin")) + param = self._load(_Draft(), root) + self._assert_loaded(param, expected) + + def test_weight_loader_receives_full_vocab_tensor(self): + """The TP shard is the parameter's weight_loader's job; the helper must hand + it the whole checkpoint tensor, not a pre-sliced one.""" + expected = torch.randn(VOCAB, HIDDEN) + draft = _Draft() + seen = {} + + def sharding_loader(p, loaded): + seen["shape"] = tuple(loaded.shape) + p.data[: loaded.shape[0]].copy_(loaded) + + draft.model.embed_tokens.weight.weight_loader = sharding_loader + with tempfile.TemporaryDirectory() as root: + _write_sharded_checkpoint(root, expected) + self._load(draft, root) + self.assertEqual(seen["shape"], (VOCAB, HIDDEN)) + + def test_load_format_gates_disk_access(self): + """dummy returns the parameter untouched; a streaming format has no weight + files to re-open and must fail instead of leaving the embedding random.""" + draft = _Draft() + param = self._load( + draft, "/nonexistent", LoadConfig(load_format=LoadFormat.DUMMY) + ) + self.assertIs(param, draft.model.embed_tokens.weight) + with self.assertRaises(ValueError): + self._load( + _Draft(), + "/nonexistent", + LoadConfig(load_format=LoadFormat.REMOTE_INSTANCE), + ) + + def test_draft_without_embedding_fails_loudly(self): + with self.assertRaises(ValueError): + self._load(_Draft(with_embedding=False), "/nonexistent") + + +if __name__ == "__main__": + unittest.main()