diff --git a/python/sglang/srt/arg_groups/validation_hook.py b/python/sglang/srt/arg_groups/validation_hook.py index 352b5f4c9..335cb0222 100644 --- a/python/sglang/srt/arg_groups/validation_hook.py +++ b/python/sglang/srt/arg_groups/validation_hook.py @@ -10,6 +10,7 @@ from typing import Any, Dict, List, Optional from sglang.srt.arg_groups.overrides import ( _hisparse_validation, + model_config_of, resolved_view, resolving_view, run_post_process_pass, @@ -24,6 +25,20 @@ from sglang.srt.utils.runai_utils import is_runai_obj_uri 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", + } +) + def validate_response_store(server_args: Any) -> None: cfg = resolving_view(server_args) @@ -35,6 +50,58 @@ def validate_response_store(server_args: Any) -> None: ) +def check_pipeline_parallel_compat( + cfg: Any, *, model_architecture: Optional[str] = None +) -> None: + """Validate features used with pipeline parallelism.""" + assert cfg.disable_overlap_schedule, ( + "Pipeline parallelism is not compatible with overlap schedule" + ) + if cfg.speculative_algorithm is not None: + assert ( + cfg.speculative_algorithm.upper() == "EAGLE" + and not cfg.enable_multi_layer_eagle + ), ( + "Pipeline parallelism currently only supports EAGLE " + "(non-multi-layer) speculative decoding" + ) + if envs.SGLANG_ENABLE_PP_SPEC.get(): + # The aggregate relay carries an EAGLE-shaped tree and only + # EAGLEWorkerV2 tail-drafts. PD prefill relays topk_p / + # topk_index / hidden states through RelayPayload; the gated + # flow replaces that relay with its own and does not carry + # those fields. + assert cfg.disaggregation_mode == "null", ( + "SGLANG_ENABLE_PP_SPEC is not compatible with --disaggregation-mode" + ) + # The PP relay slices spec results with the configured + # num_draft_tokens; adaptive spec changes it at runtime. + assert not cfg.speculative_adaptive, ( + "SGLANG_ENABLE_PP_SPEC is not compatible with --speculative-adaptive" + ) + # Every stage rebuilds the same verify input from the relayed + # per-request state, so all stages must see the same batch. + # DP attention partitions it per DP rank. + assert not cfg.enable_dp_attention, ( + "SGLANG_ENABLE_PP_SPEC is not compatible with --enable-dp-attention" + ) + else: + assert cfg.disaggregation_mode == "prefill", ( + "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 " + "pp-max-micro-batch-size, so the threshold may never be reached" + ) + + def check_server_args(server_args: Any): from sglang.srt.arg_groups.lora_hook import check_lora_server_args @@ -60,63 +127,10 @@ def check_server_args(server_args: Any): ) if cfg.pp_size > 1: - if get_platform().is_npu: - # NPU: allow PP + EAGLE speculative decoding - assert cfg.disable_overlap_schedule, ( - "Pipeline parallelism is not compatible with overlap schedule" - ) - if cfg.speculative_algorithm is not None: - assert ( - cfg.speculative_algorithm.upper() == "EAGLE" - and not cfg.enable_multi_layer_eagle - ), ( - "Pipeline parallelism currently only supports EAGLE " - "(non-multi-layer) speculative decoding" - ) - assert cfg.disaggregation_mode == "prefill", ( - "NPU PP + speculative decoding (MTP) is only supported " - "on prefill nodes (disaggregation-mode=prefill)" - ) - elif envs.SGLANG_ENABLE_PP_SPEC.get(): - assert cfg.disable_overlap_schedule, ( - "SGLANG_ENABLE_PP_SPEC requires --disable-overlap-schedule" - ) - # The relay carries an EAGLE-shaped tree and only EAGLEWorkerV2 - # tail-drafts; every other algorithm would be mis-rebuilt. - assert ( - cfg.speculative_algorithm == "EAGLE" - and not cfg.enable_multi_layer_eagle - ), ( - "SGLANG_ENABLE_PP_SPEC supports single-layer EAGLE/MTP only, " - f"got {cfg.speculative_algorithm}" - ) - # PD prefill relays topk_p / topk_index / hidden states through - # RelayPayload; the gated flow replaces that relay with its own - # and does not carry those fields. - assert cfg.disaggregation_mode == "null", ( - "SGLANG_ENABLE_PP_SPEC is not compatible with --disaggregation-mode" - ) - # The PP relay slices spec results with the configured - # num_draft_tokens; adaptive spec changes it at runtime. - assert not cfg.speculative_adaptive, ( - "SGLANG_ENABLE_PP_SPEC is not compatible with --speculative-adaptive" - ) - # Every stage rebuilds the same verify input from the relayed - # per-request state, so all stages must see the same batch. - # DP attention partitions it per DP rank. - assert not cfg.enable_dp_attention, ( - "SGLANG_ENABLE_PP_SPEC is not compatible with --enable-dp-attention" - ) - else: - # Non-NPU: PP + speculative decoding is not supported - assert cfg.disable_overlap_schedule and cfg.speculative_algorithm is None, ( - "Pipeline parallelism is not compatible with overlap schedule, speculative decoding" - ) - 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 " - "pp-max-micro-batch-size, so the threshold may never be reached" - ) + 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) 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/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index d6935939f..b4196f657 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -1265,12 +1265,18 @@ class CommonKVManager(BaseKVManager): ) # Regular MLA PP slicing - start_layer = self.kv_args.prefill_start_layer - end_layer = start_layer + len(src_kv_ptrs) # Decode pp size should be equal to prefill pp size or 1 + start_layer, end_layer = self._mla_kv_entry_span_with_pp(len(src_kv_ptrs)) sliced_dst_kv_ptrs = dst_kv_ptrs[start_layer:end_layer] return src_kv_ptrs, sliced_dst_kv_ptrs, len(src_kv_ptrs) + def _mla_kv_entry_span_with_pp(self, n_src: int) -> Tuple[int, int]: + # A plain MLA pool registers one region per layer ascending, addressed as + # layer_id - start_layer, so this stage occupies [start, start + n_src) of a + # peer that registered the whole model. Pointer view: get_mla_kv_ptrs_with_pp. + start_layer = self.kv_args.prefill_start_layer + return start_layer, start_layer + n_src + def _mla_slice_ptrs_for_pp( self, src_kv_ptrs: List[int], diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index d4e2b3d1b..9a1094bba 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -1026,6 +1026,59 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): ) peer_info.kv_xfer_segments = prepared_segments + def _build_transfer_dst_indices( + self, *, peer_info: KVArgsRegisterInfo, n_src: int, n_dst: int + ) -> List[int]: + """Map source entries to destination entries for this transfer. + + Heterogeneous PP over a plain MLA pool uses the source stage's layer span. + All other layouts use explicit layer IDs, or positional pairing for non-PP. + """ + use_pp_mla_offsets = ( + self.pp_size > 1 + and n_src != n_dst + and not self.kv_args.kv_layer_ids + and not peer_info.dst_kv_layer_ids + and self.is_mla_backend + and not self.is_hybrid_mla_backend + and not self.kv_args.mla_compression_ratios + ) + if not use_pp_mla_offsets: + pairs = build_transfer_entry_pairs( + self.kv_args.kv_layer_ids, + peer_info.dst_kv_layer_ids, + n_src, + n_dst, + allow_positional_fallback=self.pp_size == 1, + ) + return [j for _, j in pairs] + + start, end = self._mla_kv_entry_span_with_pp(n_src) + # Bootstrap admits a peer running our pp or 1, so a peer that does not cover the + # span is a matched-pp stage above 0, whose entries start at its own index 0. + if end > n_dst: + pairs = build_transfer_entry_pairs( + self.kv_args.kv_layer_ids, + peer_info.dst_kv_layer_ids, + n_src, + n_dst, + allow_positional_fallback=False, + ) + return [j for _, j in pairs] + + indices = list(range(start, end)) + src_item_lens = list(self.kv_args.kv_item_lens) + dst_item_lens = [peer_info.dst_kv_item_lens[j] for j in indices] + if src_item_lens != dst_item_lens: + # Disagreeing cell sizes mean the peers did not build the same KV geometry; + # writing anyway would silently corrupt the peer's pool. + raise RuntimeError( + "PP-heterogeneous MLA transfer: decode KV cell geometry differs from " + f"prefill over layers [{start}, {end}): prefill item_lens=" + f"{src_item_lens}, decode item_lens={dst_item_lens}" + ) + return indices + def _prepare_payload_xfer(self, peer_info: KVArgsRegisterInfo): # If prefill does not run speculative decoding (the usual case), # decode with speculative decoding will have more kv items. @@ -1110,14 +1163,9 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager): else self._num_slots_src ) - pairs = build_transfer_entry_pairs( - self.kv_args.kv_layer_ids, - peer_info.dst_kv_layer_ids, - n_src, - n_dst, - allow_positional_fallback=self.pp_size == 1, + dst_indices = self._build_transfer_dst_indices( + peer_info=peer_info, n_src=n_src, n_dst=n_dst ) - dst_indices = [j for _, j in pairs] dst_kv_ptrs = [peer_info.dst_kv_ptrs[j] for j in dst_indices] dst_kv_item_lens = [peer_info.dst_kv_item_lens[j] for j in dst_indices] dst_kv_data_lens = [ diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index c293589ba..50918355d 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -890,6 +890,11 @@ class SchedulerPPMixin: tensor_dict["draft_topk_p"] = draft_input.topk_p.contiguous() tensor_dict["draft_topk_index"] = draft_input.topk_index.contiguous() tensor_dict["draft_hidden_states"] = draft_input.hidden_states.contiguous() + # Preserve the DSA IndexShare seed when rebuilding spec_info on each rank. + if draft_input.dsa_topk_indices is not None: + tensor_dict["draft_dsa_topk_indices"] = ( + draft_input.dsa_topk_indices.contiguous() + ) has_sampling_mask_output = ( result.logits_output is not None @@ -1130,6 +1135,7 @@ class SchedulerPPMixin: bonus_tokens=next_token_ids, num_tokens_per_req=1, num_tokens_for_logprob_per_req=1, + dsa_topk_indices=pp_outputs.tensors.get("draft_dsa_topk_indices"), ) batch.spec_info = next_draft_input diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index fe865d0b8..c3b51276f 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -683,7 +683,6 @@ class ModelRunner: model=self.model, model_config=self.model_config, is_draft_worker=self.is_draft_worker, - spec_algorithm=self.spec_algorithm, ) adjust_hybrid_swa_layer_ids( model_config=self.model_config, diff --git a/python/sglang/srt/model_executor/model_runner_components/layer_setup.py b/python/sglang/srt/model_executor/model_runner_components/layer_setup.py index 52f37cb6f..dae8990fe 100644 --- a/python/sglang/srt/model_executor/model_runner_components/layer_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/layer_setup.py @@ -5,12 +5,8 @@ from typing import TYPE_CHECKING, Any, NamedTuple import msgspec from torch import nn -from sglang.srt.environ import envs -from sglang.srt.utils import is_npu - if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig - from sglang.srt.speculative.spec_info import SpeculativeAlgorithm class AttentionAndMoeLayers(NamedTuple): @@ -141,7 +137,6 @@ 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 @@ -149,8 +144,6 @@ def resolve_layer_indices( 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 @@ -159,14 +152,6 @@ def resolve_layer_indices( if loop_num > 1: num_effective_layers = num_effective_layers * loop_num - if not is_npu(): - _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, @@ -213,25 +198,6 @@ def _resolve_pp_layer_range(*, model: Any, model_num_layers: int) -> _PPLayerRan ) -def _assert_pp_mtp_compat( - *, - model_has_mtp_layers: bool, - spec_algorithm: SpeculativeAlgorithm, - num_effective_layers: int, - model_num_layers: int, -) -> None: - if envs.SGLANG_ENABLE_PP_SPEC.get(): - return - 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, diff --git a/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py b/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py index c526df775..774f325b7 100644 --- a/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py +++ b/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py @@ -403,8 +403,8 @@ class DeepseekV2WeightLoaderMixin: # Skip loading extra bias for GPTQ models. if name.endswith(".bias") and name not in params_dict: continue - # Skip loading embed_tokens if not first rank in pipeline parallelism - if ".embed_tokens." in name and not self.pp_group.is_first_rank: + # The last PP stage also owns an embedding for its NextN draft. + if ".embed_tokens." in name and name not in params_dict: continue # Skip loading norm if not last rank in pipeline parallelism if ".norm." in name and not self.pp_group.is_last_rank: diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 7c37b22e3..c2356c907 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -2793,6 +2793,13 @@ class DeepseekV2DecoderLayer(nn.Module): return output +def pp_stage_needs_embedding(pp_group, speculative_algorithm) -> bool: + """The first stage embeds inputs; the last supplies the EAGLE draft embedding.""" + return pp_group.is_first_rank or ( + pp_group.is_last_rank and speculative_algorithm is not None + ) + + class DeepseekV2Model(nn.Module): fall_back_to_pt_during_load = False @@ -2810,7 +2817,7 @@ class DeepseekV2Model(nn.Module): self.first_k_dense_replace = config.first_k_dense_replace self.pp_group = get_parallel().pp_group - if self.pp_group.is_first_rank or (_is_npu and self.pp_group.is_last_rank): + if pp_stage_needs_embedding(self.pp_group, get_spec().speculative_algorithm): self.embed_tokens = VocabParallelEmbedding( config.vocab_size, config.hidden_size, diff --git a/test/registered/unit/disaggregation/test_pp_mla_kv_transfer.py b/test/registered/unit/disaggregation/test_pp_mla_kv_transfer.py new file mode 100644 index 000000000..0c6034705 --- /dev/null +++ b/test/registered/unit/disaggregation/test_pp_mla_kv_transfer.py @@ -0,0 +1,258 @@ +"""Unit tests for NIXL KV transfer from a prefill with pp_size > 1 to a decode +peer that is not pipelined, on plain-MLA models (MLATokenToKVPool).""" + +import unittest +from types import SimpleNamespace + +from sglang.srt.disaggregation.common.conn import CommonKVManager +from sglang.srt.disaggregation.nixl.conn import NixlKVManager +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") + +NUM_LAYERS = 78 +ITEM_LEN = 576 + + +def _pp_layer_split(total: int, pp: int) -> list: + """(start_layer, num_layers) per PP rank, remainder to the last ranks.""" + base, rem = divmod(total, pp) + out, start = [], 0 + for rank in range(pp): + num = base + (1 if rank >= pp - rem else 0) + out.append((start, num)) + start += num + return out + + +class _StubKVManager: + """Carries the real span primitive, so the cases exercise the shared mapping + rather than a reimplementation of it.""" + + _mla_kv_entry_span_with_pp = CommonKVManager._mla_kv_entry_span_with_pp + get_mla_kv_ptrs_with_pp = CommonKVManager.get_mla_kv_ptrs_with_pp + + def __init__( + self, + *, + pp_size: int, + prefill_start_layer: int, + n_src: int, + is_mla_backend: bool = True, + is_hybrid_mla_backend: bool = False, + kv_layer_ids: list = (), + mla_compression_ratios=None, + item_len: int = ITEM_LEN, + ): + self.pp_size = pp_size + self.is_mla_backend = is_mla_backend + self.is_hybrid_mla_backend = is_hybrid_mla_backend + self.kv_args = SimpleNamespace( + prefill_start_layer=prefill_start_layer, + kv_layer_ids=list(kv_layer_ids), + kv_item_lens=[item_len] * n_src, + mla_compression_ratios=mla_compression_ratios, + ) + + +def _peer(*, n_dst: int, dst_kv_layer_ids: list = (), item_len: int = ITEM_LEN): + return SimpleNamespace( + dst_kv_layer_ids=list(dst_kv_layer_ids), + dst_kv_item_lens=[item_len] * n_dst, + ) + + +def _resolve(manager, peer, n_src, n_dst): + return NixlKVManager._build_transfer_dst_indices( + manager, peer_info=peer, n_src=n_src, n_dst=n_dst + ) + + +class TestPpPrefillToUnpipelinedDecode(CustomTestCase): + """Bug regression: a pp>1 prefill paired with a pp=1 decode never transferred + KV on NIXL, and failed after /health was green so the request hung rather + than the launch failing. Each stage must write into its own layer range of + the peer's pool.""" + + def test_stages_tile_the_decode_pool_exactly_once(self): + covered = [] + for start, num in _pp_layer_split(NUM_LAYERS, 4): + with self.subTest(start_layer=start): + indices = _resolve( + _StubKVManager(pp_size=4, prefill_start_layer=start, n_src=num), + _peer(n_dst=NUM_LAYERS), + num, + NUM_LAYERS, + ) + self.assertEqual(indices, list(range(start, start + num))) + covered += indices + self.assertEqual(sorted(covered), list(range(NUM_LAYERS))) + self.assertEqual(len(covered), len(set(covered))) + + def test_decode_side_draft_regions_are_never_targeted(self): + """Decode-only speculative decoding appends draft buffers after the + target's, so the span must stay inside the target layer range.""" + start, num = _pp_layer_split(NUM_LAYERS, 4)[3] + indices = _resolve( + _StubKVManager(pp_size=4, prefill_start_layer=start, n_src=num), + _peer(n_dst=NUM_LAYERS + 1), + num, + NUM_LAYERS + 1, + ) + self.assertEqual(indices, list(range(start, start + num))) + self.assertNotIn(NUM_LAYERS, indices) + + def test_cell_geometry_disagreement_raises(self): + """A per-layer item_len mismatch means the peers did not build the same + KV geometry, which must fail loudly instead of corrupting the pool.""" + start, num = _pp_layer_split(NUM_LAYERS, 4)[3] + with self.assertRaisesRegex(RuntimeError, "geometry differs"): + _resolve( + _StubKVManager(pp_size=4, prefill_start_layer=start, n_src=num), + _peer(n_dst=NUM_LAYERS, item_len=ITEM_LEN + 80), + num, + NUM_LAYERS, + ) + + +class TestExistingPairingIsUnchanged(CustomTestCase): + """All layouts outside heterogeneous plain MLA retain existing pairing.""" + + def test_unpipelined_prefill_uses_positional_pairing(self): + for n_dst in (NUM_LAYERS, NUM_LAYERS + 1): + with self.subTest(n_dst=n_dst): + self.assertEqual( + _resolve( + _StubKVManager( + pp_size=1, prefill_start_layer=0, n_src=NUM_LAYERS + ), + _peer(n_dst=n_dst), + NUM_LAYERS, + n_dst, + ), + list(range(NUM_LAYERS)), + ) + + def test_matched_pp_uses_identity_pairing(self): + for start, num in _pp_layer_split(NUM_LAYERS, 4): + with self.subTest(start_layer=start): + self.assertEqual( + _resolve( + _StubKVManager(pp_size=4, prefill_start_layer=start, n_src=num), + _peer(n_dst=num), + num, + num, + ), + list(range(num)), + ) + + def test_one_sided_layer_ids_are_rejected(self): + start, num = _pp_layer_split(NUM_LAYERS, 4)[2] + with self.assertRaisesRegex(RuntimeError, "both PD peers or neither"): + _resolve( + _StubKVManager( + pp_size=4, + prefill_start_layer=start, + n_src=num, + kv_layer_ids=range(start, start + num), + ), + _peer(n_dst=NUM_LAYERS), + num, + NUM_LAYERS, + ) + with self.assertRaisesRegex(RuntimeError, "both PD peers or neither"): + _resolve( + _StubKVManager(pp_size=4, prefill_start_layer=start, n_src=num), + _peer(n_dst=NUM_LAYERS, dst_kv_layer_ids=range(NUM_LAYERS)), + num, + NUM_LAYERS, + ) + + def test_heterogeneous_non_plain_mla_requires_layer_ids(self): + start, num = _pp_layer_split(NUM_LAYERS, 4)[2] + for label, kwargs, n_dst in ( + ("mha", {"is_mla_backend": False}, 2 * NUM_LAYERS), + ("hybrid_mla", {"is_hybrid_mla_backend": True}, NUM_LAYERS), + ( + "compressed_mla", + {"mla_compression_ratios": [4] * NUM_LAYERS}, + 2 * NUM_LAYERS, + ), + ): + with self.subTest(pool=label): + with self.assertRaisesRegex( + RuntimeError, "PP-heterogeneous transfer requires layer ids" + ): + _resolve( + _StubKVManager( + pp_size=4, + prefill_start_layer=start, + n_src=num, + **kwargs, + ), + _peer(n_dst=n_dst), + num, + n_dst, + ) + + +class TestMatchedPpWithDraftRegions(CustomTestCase): + """Derived property: bootstrap admits a peer at our pp or at 1, so under + matched pp a decode-side draft buffer makes n_src != n_dst and reaches the + span. Stage 0 degenerates to the identity and stays correct; every stage + above it must be rejected by the bound rather than writing at an offset the + peer does not have.""" + + def test_rank0_identity_span_stays_correct(self): + start, num = _pp_layer_split(NUM_LAYERS, 4)[0] + self.assertEqual(start, 0) + self.assertEqual( + _resolve( + _StubKVManager(pp_size=4, prefill_start_layer=start, n_src=num), + _peer(n_dst=num + 1), + num, + num + 1, + ), + list(range(num)), + ) + + def test_later_ranks_are_rejected(self): + for start, num in _pp_layer_split(NUM_LAYERS, 4)[1:]: + with self.subTest(start_layer=start): + with self.assertRaisesRegex( + RuntimeError, "PP-heterogeneous transfer requires layer ids" + ): + _resolve( + _StubKVManager(pp_size=4, prefill_start_layer=start, n_src=num), + _peer(n_dst=num + 1), + num, + num + 1, + ) + + +class TestPointerAndIndexViewsAgree(CustomTestCase): + """Derived property: get_mla_kv_ptrs_with_pp (pointer view) and + _build_transfer_dst_indices (index view) read the same span, so a change to + one must not silently diverge from the other.""" + + def test_views_select_the_same_destination_entries(self): + dst_ptrs = [9000 + i for i in range(NUM_LAYERS)] + for start, num in _pp_layer_split(NUM_LAYERS, 4): + with self.subTest(start_layer=start): + manager = _StubKVManager( + pp_size=4, prefill_start_layer=start, n_src=num + ) + src_ptrs = [1000 + i for i in range(num)] + + _, sliced_dst, count = manager.get_mla_kv_ptrs_with_pp( + src_ptrs, dst_ptrs + ) + indices = _resolve(manager, _peer(n_dst=NUM_LAYERS), num, NUM_LAYERS) + + self.assertEqual(count, num) + self.assertEqual(sliced_dst, [dst_ptrs[j] for j in indices]) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/managers/test_auxiliary_output.py b/test/registered/unit/managers/test_auxiliary_output.py index 049d1343b..4fed959b8 100644 --- a/test/registered/unit/managers/test_auxiliary_output.py +++ b/test/registered/unit/managers/test_auxiliary_output.py @@ -25,6 +25,7 @@ from sglang.srt.model_executor.forward_batch_info import PPProxyTensors from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.runtime_context import publish, reset_context from sglang.srt.server_args import ServerArgs +from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import published_topology @@ -689,6 +690,50 @@ def test_pipeline_parallel_auxiliary_output_round_trip(): receiver.future_map.stash.assert_called_once() +@pytest.mark.parametrize("dsa_topk_indices", [None, torch.tensor([[2, 5, 7]])]) +def test_pipeline_parallel_dsa_seed_round_trip(dsa_topk_indices): + draft_input = EagleDraftInput( + topk_p=torch.tensor([[0.8, 0.2]]), + topk_index=torch.tensor([[11, 13]]), + hidden_states=torch.tensor([[1.0, 2.0]]), + dsa_topk_indices=dsa_topk_indices, + ) + result = GenerationBatchResult( + logits_output=None, + next_token_ids=torch.tensor([17]), + next_draft_input=draft_input, + ) + batch = SimpleNamespace( + return_logprob=False, + req_pool_indices=torch.tensor([3]), + input_ids=torch.tensor([5]), + spec_info=None, + ) + + tensors = Scheduler._pp_prepare_tensor_dict( + object.__new__(Scheduler), result, batch + ) + if dsa_topk_indices is None: + assert "draft_dsa_topk_indices" not in tensors + else: + assert torch.equal(tensors["draft_dsa_topk_indices"], dsa_topk_indices) + + receiver = object.__new__(Scheduler) + receiver.pp_group = SimpleNamespace(is_first_rank=False) + receiver.future_map = SimpleNamespace(stash=Mock()) + Scheduler._pp_prep_batch_result( + receiver, + batch, + PPBatchMetadata(can_run_cuda_graph=True), + PPProxyTensors(tensors), + ) + + if dsa_topk_indices is None: + assert batch.spec_info.dsa_topk_indices is None + else: + assert torch.equal(batch.spec_info.dsa_topk_indices, dsa_topk_indices) + + def test_pipeline_parallel_auxiliary_output_stays_packed_before_first_rank(): device_output = DeviceOutput(torch.tensor([1.0])) result = GenerationBatchResult( diff --git a/test/registered/unit/models/test_deepseek_pp_mtp_embedding.py b/test/registered/unit/models/test_deepseek_pp_mtp_embedding.py new file mode 100644 index 000000000..cfb583670 --- /dev/null +++ b/test/registered/unit/models/test_deepseek_pp_mtp_embedding.py @@ -0,0 +1,58 @@ +"""Pipeline-stage embedding ownership for DeepSeek/GLM NextN.""" + +import unittest +from types import SimpleNamespace + +import torch + +from sglang.srt.layers.utils import PPMissingLayer +from sglang.srt.models.deepseek_v2 import pp_stage_needs_embedding +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +def _pp_group(*, rank: int, size: int) -> SimpleNamespace: + return SimpleNamespace(is_first_rank=rank == 0, is_last_rank=rank == size - 1) + + +class TestDeepseekPPStageEmbedding(CustomTestCase): + def test_without_speculative_decoding_only_the_first_stage_embeds(self): + for rank in range(4): + with self.subTest(rank=rank): + self.assertEqual( + pp_stage_needs_embedding(_pp_group(rank=rank, size=4), None), + rank == 0, + ) + + def test_speculative_decoding_adds_the_last_stage(self): + wants = [ + pp_stage_needs_embedding(_pp_group(rank=rank, size=4), "EAGLE") + for rank in range(4) + ] + self.assertEqual(wants, [True, False, False, True]) + + def test_middle_stages_never_embed(self): + self.assertFalse(pp_stage_needs_embedding(_pp_group(rank=2, size=4), "EAGLE")) + + def test_without_pipeline_parallelism_the_single_stage_embeds_either_way(self): + for spec in (None, "EAGLE"): + with self.subTest(spec=spec): + self.assertTrue( + pp_stage_needs_embedding(_pp_group(rank=0, size=1), spec) + ) + + def test_pp_missing_layer_registers_no_parameters(self): + placeholder = PPMissingLayer() + self.assertEqual(dict(placeholder.named_parameters()), {}) + + def test_a_real_embedding_registers_weight_under_its_prefix(self): + model = torch.nn.Module() + model.embed_tokens = torch.nn.Embedding(8, 4) + params = dict(model.named_parameters()) + self.assertIn("embed_tokens.weight", params) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 146870a94..60da7dcd5 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -13,7 +13,12 @@ import msgspec import msgspec.structs import sglang.srt.server_args as server_args_module -from sglang.srt.arg_groups import parallel_hook, pd_disaggregation_hook, serving_hook +from sglang.srt.arg_groups import ( + parallel_hook, + pd_disaggregation_hook, + serving_hook, + validation_hook, +) from sglang.srt.arg_groups.attention_hook import ( handle_attention_backend_compatibility, handle_deterministic_inference, @@ -68,6 +73,7 @@ from sglang.srt.arg_groups.serving_hook import ( ) from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding from sglang.srt.arg_groups.validation_hook import ( + check_pipeline_parallel_compat, check_two_batch_overlap, ) from sglang.srt.entrypoints.sidecar import ( @@ -2364,6 +2370,115 @@ class TestCudaGraphConfigDataclassAccess(CustomTestCase): self.assertEqual(config.compiler, "eager") +class TestPipelineParallelCompat(CustomTestCase): + """Features supported with `pipeline-parallel-size > 1`.""" + + _SUPPORTED_ARCH = "GlmMoeDsaForCausalLM" + + @staticmethod + def _cfg(**overrides): + cfg = dict( + disable_overlap_schedule=True, + speculative_algorithm=None, + enable_multi_layer_eagle=False, + disaggregation_mode="prefill", + min_free_slots_delay=None, + ) + cfg.update(overrides) + return SimpleNamespace(**cfg) + + def test_overlap_schedule_must_be_off(self): + with self.assertRaisesRegex(AssertionError, "overlap schedule"): + check_pipeline_parallel_compat(self._cfg(disable_overlap_schedule=False)) + + def test_no_speculative_decoding_is_fine(self): + check_pipeline_parallel_compat(self._cfg()) + + def test_eagle_is_allowed_on_prefill(self): + check_pipeline_parallel_compat( + self._cfg(speculative_algorithm="EAGLE"), + model_architecture=self._SUPPORTED_ARCH, + ) + + def test_eagle_is_rejected_outside_prefill(self): + for mode in ("decode", "null"): + with self.subTest(disaggregation_mode=mode): + with self.assertRaisesRegex(AssertionError, "prefill nodes"): + 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", + ): + 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", + disaggregation_mode="null", + speculative_adaptive=False, + enable_dp_attention=False, + ) + with patch.object( + validation_hook.envs.SGLANG_ENABLE_PP_SPEC, "get", return_value=True + ): + check_pipeline_parallel_compat(cfg, model_architecture="LlamaForCausalLM") + with self.assertRaisesRegex(AssertionError, "SGLANG_ENABLE_PP_SPEC"): + check_pipeline_parallel_compat( + self._cfg(speculative_algorithm="EAGLE"), + model_architecture=self._SUPPORTED_ARCH, + ) + + 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, + ) + + 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, + ) + + 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, + ) + + def test_min_free_slots_delay_is_rejected(self): + with self.assertRaisesRegex(AssertionError, "min-free-slots-delay"): + check_pipeline_parallel_compat(self._cfg(min_free_slots_delay=4)) + + class TestCudaGraphPrefillMaxContextResolution(CustomTestCase): @staticmethod def _make_args(max_context_size, model_context_len=4096, page_size=64):