From d3ae04ab74d39abd3427f9d665e7a752511a8fac Mon Sep 17 00:00:00 2001 From: LostFox11 Date: Thu, 17 Sep 2026 16:12:35 +0800 Subject: [PATCH] [Feature][MRV2] Support GLM5.2 DSpark with PP Add Ascend Spec+PP support for GLM-5.2 DSpark on MRV2: * Replace the deferred-broadcast PP transport with the vLLM 0.30 sampled-token protocol (participation gate as a pure function of the shared scheduler batch, immediate broadcasts, receive-side generation-counter filtering). This removes the collective-order divergence that deadlocked the engine under KV saturation with async EPLB. Applies to vLLM 0.28/0.29 only; 0.30+ uses the upstream protocol natively. * Route Spec+PP by the paired vLLM version (use_legacy_spec_pp), with the DSpark loader bypass and quant-config inheritance for same-checkpoint drafts. * Aux hidden-state transport across PP ranks with dual-path support: legacy pp_transport buffers (0.28/0.29) and upstream EagleModelMixin slot bookkeeping (0.30+). * DeepseekV2 MLA attention init with Ascend IndexCache and top-k skip-pattern support. Hardware-validated: 97/198 GPQA questions under the original deadlock configuration (c32, u0.80, async EPLB, KV saturation), preemption 78 vs 44,202 without the fix, accuracy parity with upstream 0.30 (91.92% on same questions). Signed-off-by: LostFox11 --- tests/ut/models/test_glm_moe_dsa.py | 34 +++ tests/ut/patch/platform/test_patch_pp_mtp.py | 10 +- .../ut/patch/worker/test_patch_deepseek_v2.py | 94 ++++++++ tests/ut/patch/worker/test_patch_dspark_pp.py | 57 +++++ tests/ut/patch/worker/test_patch_spec_pp.py | 188 +++++++++++++++ tests/ut/worker/test_model_runner_v2.py | 5 +- tests/ut/worker/v2/test_pp_utils.py | 106 +++++++++ vllm_ascend/models/deepseek_mtp.py | 6 + vllm_ascend/patch/platform/patch_pp_mtp.py | 2 +- vllm_ascend/patch/worker/patch_deepseek_v2.py | 100 ++++++-- .../patch/worker/patch_v2/patch_dspark.py | 41 ++-- .../patch/worker/patch_v2/patch_spec_pp.py | 222 ++++++++++++------ vllm_ascend/worker/v2/model_runner.py | 20 +- vllm_ascend/worker/v2/pp_utils.py | 30 ++- 14 files changed, 792 insertions(+), 123 deletions(-) create mode 100644 tests/ut/models/test_glm_moe_dsa.py create mode 100644 tests/ut/patch/worker/test_patch_dspark_pp.py create mode 100644 tests/ut/patch/worker/test_patch_spec_pp.py create mode 100644 tests/ut/worker/v2/test_pp_utils.py diff --git a/tests/ut/models/test_glm_moe_dsa.py b/tests/ut/models/test_glm_moe_dsa.py new file mode 100644 index 000000000000..78910444c8eb --- /dev/null +++ b/tests/ut/models/test_glm_moe_dsa.py @@ -0,0 +1,34 @@ +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import torch.nn as nn +from vllm.model_executor.models.deepseek_v2 import GlmMoeDsaForCausalLM + +from vllm_ascend.models.deepseek_mtp import AscendGlmMoeDsaForCausalLM + + +@pytest.mark.parametrize("use_v2", [False, True]) +@pytest.mark.parametrize("pp_size,local_layers", [(1, 75), (2, 39), (2, 36), (2, 0), (3, 24)]) +def test_eplb_layer_count_matches_local_weights_only_for_v2_pp(use_v2, pp_size, local_layers): + config = SimpleNamespace( + use_v2_model_runner=use_v2, + parallel_config=SimpleNamespace(pipeline_parallel_size=pp_size), + ) + original_count = 75 if local_layers else 0 + + def init(self, *, vllm_config, prefix): + nn.Module.__init__(self) + assert vllm_config is config + assert prefix == "target" + self.num_moe_layers = original_count + self.moe_layers = [object() for _ in range(local_layers)] + + with patch.object(GlmMoeDsaForCausalLM, "__init__", init): + model = AscendGlmMoeDsaForCausalLM(vllm_config=config, prefix="target") + + expected = local_layers if use_v2 and pp_size > 1 else original_count + assert model.num_moe_layers == expected + assert len(model.moe_layers) == local_layers diff --git a/tests/ut/patch/platform/test_patch_pp_mtp.py b/tests/ut/patch/platform/test_patch_pp_mtp.py index b45c3a16b44b..bcd2cbb0d366 100644 --- a/tests/ut/patch/platform/test_patch_pp_mtp.py +++ b/tests/ut/patch/platform/test_patch_pp_mtp.py @@ -26,18 +26,22 @@ from vllm_ascend.worker.model_runner_v1 import NPUModelRunner -def test_model_config_validates_local_mtp_drafter_as_single_pp_rank(monkeypatch): +@pytest.mark.parametrize( + "model_type,architecture", + [("qwen3_5_mtp", "Qwen3_5MTP"), ("qwen3", "DSparkDraftModel"), ("qwen3", "Qwen3DSparkModel")], +) +def test_model_config_validates_local_drafter_as_single_pp_rank(monkeypatch, model_type, architecture): fake_registry = SimpleNamespace( is_pp_supported_model=lambda _architectures, _model_config: False, ) monkeypatch.setattr(ModelConfig, "registry", property(lambda _self: fake_registry)) model_config = ModelConfig.__new__(ModelConfig) - model_config.hf_config = SimpleNamespace(model_type="qwen3_5_mtp") + model_config.hf_config = SimpleNamespace(model_type=model_type) model_config.runner = "draft" model_config.model_arch_config = SimpleNamespace( total_num_attention_heads=1, - architectures=["Qwen3_5MTP"], + architectures=[architecture], ) model_config.multimodal_config = None diff --git a/tests/ut/patch/worker/test_patch_deepseek_v2.py b/tests/ut/patch/worker/test_patch_deepseek_v2.py index 6aeb5d980011..93a9a59ae9bd 100644 --- a/tests/ut/patch/worker/test_patch_deepseek_v2.py +++ b/tests/ut/patch/worker/test_patch_deepseek_v2.py @@ -1,7 +1,13 @@ # SPDX-License-Identifier: Apache-2.0 from types import SimpleNamespace +from unittest.mock import Mock +import pytest +import torch +from vllm.model_executor.models.deepseek_v2 import DeepseekV2Model + +from vllm_ascend.patch.worker import patch_deepseek_v2 from vllm_ascend.patch.worker.patch_deepseek_v2 import _should_skip_indexer_init @@ -34,3 +40,91 @@ def test_mtp_layer_keeps_indexer(): "model.layers.80.self_attn", skip_topk=True, ) + + +@pytest.mark.parametrize("native", [False, True]) +@pytest.mark.parametrize("boundaries", [(0, 78), (0, 42, 78), (0, 20, 40, 59, 78)]) +def test_aux_relay_matches_unpartitioned_forward(monkeypatch, native, boundaries): + if native and not hasattr(DeepseekV2Model, "pack_local_aux_hidden_states"): + pytest.skip("The installed vLLM release has no native aux relay") + aux_layers = (0, 2, 20, 39, 58, 75, 78) + ids = torch.zeros(4, dtype=torch.long) + + def layer(positions, hidden, residual, scaling): + residual = torch.zeros_like(hidden) if residual is None else residual + return hidden + 1, residual + 1 + + def run(split): + incoming = None + for start, end in zip(split, split[1:]): + first, last = start == 0, end == 78 + group = SimpleNamespace(is_first_rank=first, is_last_rank=last, world_size=len(split) - 1) + monkeypatch.setattr(patch_deepseek_v2, "get_pp_group", lambda group=group: group) + model = DeepseekV2Model.__new__(DeepseekV2Model) + torch.nn.Module.__init__(model) + model.config = SimpleNamespace(hidden_size=2) + model.hidden_size = 2 + model.start_layer, model.end_layer = start, end + model.layers = [layer] * 78 + model.aux_hidden_state_layers = aux_layers + model._use_upstream_aux_relay = native + model.embed_input_ids = lambda ids: torch.zeros(len(ids), 2) + model.norm = lambda hidden, residual: (hidden + residual, None) + if native: + import vllm.distributed.parallel_state as parallel_state + from vllm.v1.worker.gpu.pp_utils import PPHandler + + monkeypatch.setattr(parallel_state, "model_parallel_is_initialized", lambda: True) + monkeypatch.setattr(parallel_state, "get_pp_group", lambda group=group: group) + model._set_aux_hidden_state_layers(aux_layers) + output = patch_deepseek_v2._patched_forward(model, ids if first else None, ids, incoming) + if not last: + if native: + # Use the actual upstream relay, without creating streams or process groups. + handler = SimpleNamespace( + aux_hidden_state_relay_keys=[ + f"aux_hidden_states_{i}" for i in range(model._aux_slot_base_cached) + ] + ) + output = PPHandler.relay_aux_hidden_states(handler, incoming, output) + prefix = "aux_hidden_states_" if native else "pp_transport_aux_hidden_states_" + expected_count = sum(idx <= end for idx in aux_layers) + assert set(output.tensors) == {"hidden_states", "residual"} | { + f"{prefix}{idx}" for idx in range(expected_count) + } + incoming = output + return output + + expected_hidden, expected_aux = run((0, 78)) + actual_hidden, actual_aux = run(boundaries) + torch.testing.assert_close(actual_hidden, expected_hidden) + assert len(actual_aux) == len(expected_aux) == len(aux_layers) + for actual, expected in zip(actual_aux, expected_aux): + torch.testing.assert_close(actual, expected) + + +@pytest.mark.parametrize("legacy,v2", [(True, True), (False, True), (False, False)]) +def test_aux_buffer_factory_uses_one_protocol(monkeypatch, legacy, v2): + factory = object() + model = SimpleNamespace(make_empty_intermediate_tensors=factory) + monkeypatch.setattr(patch_deepseek_v2, "_original_deepseek_v2_model_init", lambda *args, **kwargs: None) + monkeypatch.setattr(patch_deepseek_v2.pp_utils, "use_legacy_spec_pp", lambda: legacy) + wrapped = object() + monkeypatch.setattr(patch_deepseek_v2.pp_utils, "make_empty_intermediate_tensors", lambda *args: wrapped) + patch_deepseek_v2._patched_deepseek_v2_model_init(model, vllm_config=SimpleNamespace(use_v2_model_runner=v2)) + assert model._use_upstream_aux_relay is (v2 and not legacy) + assert model.make_empty_intermediate_tensors is (factory if v2 and not legacy else wrapped) + + +@pytest.mark.parametrize("native", [False, True]) +def test_aux_setter_preserves_v1_behavior(native): + if not hasattr(patch_deepseek_v2, "_set_aux_hidden_state_layers"): + pytest.skip("The installed vLLM release uses its original setter") + model = SimpleNamespace(_use_upstream_aux_relay=native, _set_aux_hidden_state_layers=Mock()) + layers = (2, 20, 39, 58, 75) + patch_deepseek_v2._set_aux_hidden_state_layers(SimpleNamespace(model=model), layers) + if native: + model._set_aux_hidden_state_layers.assert_called_once_with(layers) + else: + model._set_aux_hidden_state_layers.assert_not_called() + assert model.aux_hidden_state_layers == layers diff --git a/tests/ut/patch/worker/test_patch_dspark_pp.py b/tests/ut/patch/worker/test_patch_dspark_pp.py new file mode 100644 index 000000000000..77d4947651d2 --- /dev/null +++ b/tests/ut/patch/worker/test_patch_dspark_pp.py @@ -0,0 +1,57 @@ +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace + +import pytest + +from vllm_ascend.patch.worker.patch_v2 import patch_dspark +from vllm_ascend.worker.v2 import pp_utils + + +@pytest.mark.parametrize("version", ["0.28.0", "0.29.0", "0.30.0"]) +@pytest.mark.parametrize("pp_size", [1, 2]) +@pytest.mark.parametrize("fail", [True, False]) +def test_dspark_draft_partition_isolation(monkeypatch, version, pp_size, fail): + legacy = version != "0.30.0" + bypass_pp_guard = legacy and pp_size > 1 + config = SimpleNamespace( + parallel_config=SimpleNamespace(pipeline_parallel_size=pp_size), + model_config=SimpleNamespace(model="target", architecture="GlmMoeDsaForCausalLM"), + speculative_config=SimpleNamespace(method="dspark", draft_model_config=SimpleNamespace(model="draft")), + ) + monkeypatch.setattr(pp_utils, "use_legacy_spec_pp", lambda: legacy) + monkeypatch.setattr(patch_dspark, "use_legacy_spec_pp", lambda: legacy) + get_pp_group = lambda: SimpleNamespace(world_size=pp_size) + monkeypatch.setattr(patch_dspark.dspark_utils, "get_pp_group", get_pp_group, raising=False) + monkeypatch.setattr(pp_utils.vllm_envs, "VLLM_PP_LAYER_PARTITION", "42,36") + should_share = patch_dspark.eagle_utils._should_share + monkeypatch.setattr(patch_dspark.dspark_utils, "_should_share", should_share, raising=False) + + def load(target, received_config): + assert received_config is config + assert config.parallel_config.pipeline_parallel_size == (1 if bypass_pp_guard else pp_size) + expected_partition = None if pp_size > 1 else "42,36" + assert expected_partition == pp_utils.vllm_envs.VLLM_PP_LAYER_PARTITION + assert patch_dspark.dspark_utils.get_pp_group().world_size == (1 if bypass_pp_guard else pp_size) + if bypass_pp_guard: + assert patch_dspark.eagle_utils._should_share is not should_share + assert patch_dspark.dspark_utils._should_share is patch_dspark.eagle_utils._should_share + else: + assert patch_dspark.eagle_utils._should_share is should_share + assert patch_dspark.dspark_utils._should_share is should_share + if fail: + raise RuntimeError("draft load failed") + return target + + monkeypatch.setattr(patch_dspark, "_original_load_dspark_model", load) + target = object() + if fail: + with pytest.raises(RuntimeError, match="draft load failed"): + patch_dspark._load_dspark_model_with_target_quant(target, config) + else: + assert patch_dspark._load_dspark_model_with_target_quant(target, config) is target + assert config.parallel_config.pipeline_parallel_size == pp_size + assert pp_utils.vllm_envs.VLLM_PP_LAYER_PARTITION == "42,36" + assert patch_dspark.eagle_utils._should_share is should_share + assert patch_dspark.dspark_utils.get_pp_group is get_pp_group + assert patch_dspark.dspark_utils._should_share is should_share diff --git a/tests/ut/patch/worker/test_patch_spec_pp.py b/tests/ut/patch/worker/test_patch_spec_pp.py new file mode 100644 index 000000000000..e31c65e903ab --- /dev/null +++ b/tests/ut/patch/worker/test_patch_spec_pp.py @@ -0,0 +1,188 @@ +# SPDX-License-Identifier: Apache-2.0 + +from dataclasses import dataclass +from types import SimpleNamespace + +import numpy as np + +from vllm_ascend.patch.worker.patch_v2.patch_spec_pp import ( + SpecPPPendingRecv, + compute_need_sampled_mask, + install_upstream_spec_pp_protocol, +) + + +@dataclass +class _InputBatch: + num_computed_tokens_np: np.ndarray + num_scheduled_tokens: np.ndarray + prefill_len_np: np.ndarray + # Rank-local bound consulted by the release gate; the vendored 0.30 + # gate must ignore it entirely. + max_seq_len_np: np.ndarray | None = None + + +class TestPureParticipationGate: + def test_matches_release_gate_for_regular_decode(self): + batch = _InputBatch( + num_computed_tokens_np=np.array([56, 0], dtype=np.int32), + num_scheduled_tokens=np.array([8, 25], dtype=np.int32), + prefill_len_np=np.array([25, 50], dtype=np.int32), + ) + np.testing.assert_array_equal( + compute_need_sampled_mask(batch), + np.array([True, False]), + ) + + def test_ignores_rank_local_max_seq_len_bound(self): + """The release gate drops requests whose next token may hit + max_seq_len; two ranks with different views of that bound must + still compute the same mask here.""" + base = dict( + num_computed_tokens_np=np.array([56], dtype=np.int32), + num_scheduled_tokens=np.array([8], dtype=np.int32), + prefill_len_np=np.array([25], dtype=np.int32), + ) + tight = _InputBatch(max_seq_len_np=np.array([57]), **base) + loose = _InputBatch(max_seq_len_np=np.array([98304]), **base) + np.testing.assert_array_equal( + compute_need_sampled_mask(tight), + compute_need_sampled_mask(loose), + ) + assert compute_need_sampled_mask(tight) is not None + + def test_all_prefill_chunks_only_returns_none(self): + batch = _InputBatch( + num_computed_tokens_np=np.array([0, 8], dtype=np.int32), + num_scheduled_tokens=np.array([8, 8], dtype=np.int32), + prefill_len_np=np.array([1024, 512], dtype=np.int32), + ) + assert compute_need_sampled_mask(batch) is None + + +class _FakeTensor: + def __init__(self, shape): + self.shape = shape + + def unbind(self, dim=0): + return _FakeTensor((self.shape[1],)), _FakeTensor((self.shape[1],)) + + def record_stream(self, stream): + pass + + +class _StreamStub: + def wait_stream(self, other): + pass + + def record_event(self): + return object() + + +class _StreamCtx: + def __enter__(self): + return None + + def __exit__(self, *exc): + return False + + +class TestProtocolInstall: + def _handler(self): + return SimpleNamespace( + is_last_rank=False, + device="npu:0", + max_sample_len=8, + last_rank=1, + broadcast_group=object(), + broadcast_stream=_StreamStub(), + main_stream=_StreamStub(), + req_idx_gen_np=np.zeros(4, dtype=np.int64), + queue=[None], + ) + + def _patch_transport(self, monkeypatch, sent): + """Swap the module-global ``torch`` for a stub namespace. Patching + ``torch.distributed`` attributes directly is unreliable once + vllm-ascend has rebinded collectives onto torch_npu.""" + from types import SimpleNamespace as NS + + import vllm_ascend.patch.worker.patch_v2.patch_spec_pp as mod + + def fake_broadcast(tensor, src=None, group=None): + sent.append(tuple(tensor.shape)) + + fake_torch = NS( + distributed=NS(broadcast=fake_broadcast), + cuda=NS(stream=lambda _s: _StreamCtx()), + empty=lambda *a, **k: _FakeTensor((a[0], a[1])), + int64=object(), + int32=object(), + ) + monkeypatch.setattr(mod, "torch", fake_torch) + + def test_installs_and_is_idempotent(self, monkeypatch): + handler = self._handler() + req_states = SimpleNamespace() + install_upstream_spec_pp_protocol(handler, req_states, num_speculative_steps=7) + installed = (handler.receive, handler.broadcast, handler.broadcast_drafts) + install_upstream_spec_pp_protocol(handler, req_states, num_speculative_steps=7) + assert (handler.receive, handler.broadcast, handler.broadcast_drafts) == installed + assert handler.broadcast_draft_tokens is handler.broadcast_drafts + + def test_receive_issues_three_broadcasts_and_reserves_slot(self, monkeypatch): + handler = self._handler() + req_states = SimpleNamespace() + install_upstream_spec_pp_protocol(handler, req_states, num_speculative_steps=7) + + batch = _InputBatch( + num_computed_tokens_np=np.array([56], dtype=np.int32), + num_scheduled_tokens=np.array([8], dtype=np.int32), + prefill_len_np=np.array([25], dtype=np.int32), + ) + batch.num_reqs = 1 # type: ignore[attr-defined] + batch.idx_mapping_np = np.array([2], dtype=np.int64) # type: ignore[attr-defined] + batch.idx_mapping = object() # type: ignore[attr-defined] + + sent: list[tuple[int, ...]] = [] + self._patch_transport(monkeypatch, sent) + gather_all = handler.receive(batch) + + # sampled tokens, combined and drafts: three broadcasts, in order. + assert sent == [(1, 8), (2, 1), (1, 7)] + assert gather_all is True + slot = handler.queue[-1] + assert isinstance(slot, SpecPPPendingRecv) + assert slot.draft_tokens is not None + + def test_receive_without_sample_needs_skips_transport(self, monkeypatch): + handler = self._handler() + req_states = SimpleNamespace() + install_upstream_spec_pp_protocol(handler, req_states, num_speculative_steps=7) + + batch = _InputBatch( + num_computed_tokens_np=np.array([0], dtype=np.int32), + num_scheduled_tokens=np.array([8], dtype=np.int32), + prefill_len_np=np.array([1024], dtype=np.int32), + ) + sent: list[tuple[int, ...]] = [] + self._patch_transport(monkeypatch, sent) + assert handler.receive(batch) is False + assert sent == [] + assert handler.queue[-1] is None + + def test_pending_recv_extends_release_fields_with_drafts(self): + fields = SpecPPPendingRecv.__dataclass_fields__ + release_fields = ( + "event", + "sampled_tokens", + "num_sampled", + "num_rejected", + "idx_mapping", + "idx_mapping_np", + "need_sampled_mask", + "gen_at_receive_np", + ) + for name in release_fields: + assert name in fields + assert fields["draft_tokens"].default is None diff --git a/tests/ut/worker/test_model_runner_v2.py b/tests/ut/worker/test_model_runner_v2.py index c04d403ac5f3..aa79321dc0e7 100644 --- a/tests/ut/worker/test_model_runner_v2.py +++ b/tests/ut/worker/test_model_runner_v2.py @@ -452,7 +452,7 @@ def test_init_spec_pp_full_graph_and_speculator(): patch("vllm_ascend.worker.v2.model_runner.set_cos_and_sin"), patch("vllm_ascend.worker.v2.model_runner.set_mc2_tokens_capacity"), patch("vllm_ascend.worker.v2.model_runner.set_mc2_mask"), - patch("vllm_ascend.patch.worker.patch_v2.patch_spec_pp.install_spec_pp_token_broadcast") as install_pp, + patch("vllm_ascend.patch.worker.patch_v2.patch_spec_pp.install_upstream_spec_pp_protocol") as install_pp, patch("torch.npu.Stream", return_value="stream"), patch("torch.npu.Event", return_value="event"), patch("torch.empty", return_value=torch.zeros(2, dtype=torch.int32)), @@ -465,13 +465,14 @@ def test_init_spec_pp_full_graph_and_speculator(): restore_pp.assert_called_once() assert eplb_cls.call_args.kwargs["load_collection_phase"] == "decode" assert runner.use_aclgraph is True - assert runner.use_spec_pp is True assert runner.use_aux_hidden_state_outputs is True assert runner.speculator is speculator assert speculator.update_stream is runner.update_stream if vllm_version_is("0.28.0"): + assert runner.use_spec_pp is True install_pp.assert_called_once() else: + assert runner.use_spec_pp is False install_pp.assert_not_called() assert runner.update_stream is not None assert runner.decode_query_len == 2 diff --git a/tests/ut/worker/v2/test_pp_utils.py b/tests/ut/worker/v2/test_pp_utils.py new file mode 100644 index 000000000000..4641e1ea38f2 --- /dev/null +++ b/tests/ut/worker/v2/test_pp_utils.py @@ -0,0 +1,106 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM Ascend project + +import os +from types import SimpleNamespace + +import pytest +import vllm.envs as vllm_envs +from vllm.distributed.utils import get_pp_indices + +from vllm_ascend.worker.v2 import pp_utils +from vllm_ascend.worker.v2.pp_utils import SpecPPSupport, bypass_upstream_spec_pp_guard + + +@pytest.fixture(autouse=True) +def _clear_partition_cache(): + """Drop a cached VLLM_PP_LAYER_PARTITION left by other test modules.""" + vllm_envs.__dict__.pop("VLLM_PP_LAYER_PARTITION", None) + yield + vllm_envs.__dict__.pop("VLLM_PP_LAYER_PARTITION", None) + + +@pytest.mark.parametrize( + "version, legacy", + [ + ("0.28.0", True), + ("0.28.1+empty", True), + ("0.29.0", True), + ("0.29.1rc1", True), + ("0.30.0", False), + ("0.31.0.dev12", False), + ], +) +@pytest.mark.parametrize("override", [False, True]) +def test_spec_pp_version_routing(monkeypatch, version, legacy, override): + monkeypatch.setattr(pp_utils.vllm, "__version__", "0.30.0" if override else version) + monkeypatch.setattr(pp_utils.envs, "VLLM_VERSION", version if override else None) + monkeypatch.setattr(pp_utils, "vllm_version_is", lambda version: False) + assert pp_utils.use_legacy_spec_pp() is legacy + + +def test_untagged_release_uses_existing_version_detection(monkeypatch): + monkeypatch.setattr(pp_utils.envs, "VLLM_VERSION", "0.1.dev1+g123.empty") + monkeypatch.setattr(pp_utils, "vllm_version_is", lambda version: version == "0.28.0") + assert pp_utils.use_legacy_spec_pp() + + +@pytest.mark.parametrize("cached", [False, True]) +@pytest.mark.parametrize("partition", [None, "42,36"]) +@pytest.mark.parametrize("fail", [False, True]) +def test_unsharded_draft_preserves_target_partition(monkeypatch, cached, partition, fail): + monkeypatch.setattr(pp_utils, "use_legacy_spec_pp", lambda: True) + was_cached = vllm_envs._is_envs_cache_enabled() + vllm_envs.disable_envs_cache() + if partition is None: + monkeypatch.delenv("VLLM_PP_LAYER_PARTITION", raising=False) + else: + monkeypatch.setenv("VLLM_PP_LAYER_PARTITION", partition) + config = SimpleNamespace(parallel_config=SimpleNamespace(pipeline_parallel_size=2)) + support = SpecPPSupport(bypass_upstream_pp_guard=True) + + def initialize(): + with bypass_upstream_spec_pp_guard(config, support) as bypassed: + assert bypassed + assert config.parallel_config.pipeline_parallel_size == 1 + assert vllm_envs.VLLM_PP_LAYER_PARTITION is None + assert os.environ.get("VLLM_PP_LAYER_PARTITION") == partition + assert get_pp_indices(78, 0, 1) == (0, 78) + with bypass_upstream_spec_pp_guard(config, support): + assert get_pp_indices(78, 0, 1) == (0, 78) + assert vllm_envs.VLLM_PP_LAYER_PARTITION is None + if fail: + raise RuntimeError("draft initialization failed") + + try: + if cached: + vllm_envs.enable_envs_cache() + if fail: + with pytest.raises(RuntimeError, match="draft initialization failed"): + initialize() + else: + initialize() + assert config.parallel_config.pipeline_parallel_size == 2 + assert partition == vllm_envs.VLLM_PP_LAYER_PARTITION + if partition is not None: + assert get_pp_indices(78, 0, 2) == (0, 42) + assert get_pp_indices(78, 1, 2) == (42, 78) + finally: + vllm_envs.disable_envs_cache() + monkeypatch.undo() + if was_cached: + vllm_envs.enable_envs_cache() + + +@pytest.mark.parametrize( + "legacy,support", + [(True, None), (True, SpecPPSupport()), (False, SpecPPSupport(bypass_upstream_pp_guard=True))], +) +def test_pp_guard_noop_preserves_partition(monkeypatch, legacy, support): + monkeypatch.setattr(pp_utils, "use_legacy_spec_pp", lambda: legacy) + monkeypatch.setattr(vllm_envs, "VLLM_PP_LAYER_PARTITION", "42,36") + config = SimpleNamespace(parallel_config=SimpleNamespace(pipeline_parallel_size=2)) + with bypass_upstream_spec_pp_guard(config, support) as bypassed: + assert not bypassed + assert config.parallel_config.pipeline_parallel_size == 2 + assert vllm_envs.VLLM_PP_LAYER_PARTITION == "42,36" diff --git a/vllm_ascend/models/deepseek_mtp.py b/vllm_ascend/models/deepseek_mtp.py index c6db223dbdcb..a5eccc9a80ca 100644 --- a/vllm_ascend/models/deepseek_mtp.py +++ b/vllm_ascend/models/deepseek_mtp.py @@ -72,6 +72,12 @@ def _rewrite_spec_layer_name(self, spec_layer: int, name: str) -> str: class AscendGlmMoeDsaForCausalLM(GlmMoeDsaForCausalLM): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): + super().__init__(vllm_config=vllm_config, prefix=prefix) + if vllm_config.use_v2_model_runner and vllm_config.parallel_config.pipeline_parallel_size > 1: + # EPLB maps and expert weights must describe the same local layers. + self.num_moe_layers = len(self.moe_layers) + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: if vllm_version_is("0.28.0"): loader = AutoWeightsLoader(self, skip_prefixes=["rot."]) diff --git a/vllm_ascend/patch/platform/patch_pp_mtp.py b/vllm_ascend/patch/platform/patch_pp_mtp.py index 4ea0840c554a..3ca52a8636ea 100644 --- a/vllm_ascend/patch/platform/patch_pp_mtp.py +++ b/vllm_ascend/patch/platform/patch_pp_mtp.py @@ -311,7 +311,7 @@ def _patched_verify_with_parallel_config(self, parallel_config): arch.startswith("Eagle") or arch.endswith("Eagle3") for arch in architectures ) is_mtp_drafter = model_type in mtp_model_types - is_dspark_drafter = "DSparkDraftModel" in architectures + is_dspark_drafter = any(arch in ("DSparkDraftModel", "Qwen3DSparkModel") for arch in architectures) if ( getattr(self, "runner", None) == "draft" and (is_eagle_drafter or is_mtp_drafter or is_dspark_drafter) diff --git a/vllm_ascend/patch/worker/patch_deepseek_v2.py b/vllm_ascend/patch/worker/patch_deepseek_v2.py index 07e05a4ea3ea..d48998d95b2a 100644 --- a/vllm_ascend/patch/worker/patch_deepseek_v2.py +++ b/vllm_ascend/patch/worker/patch_deepseek_v2.py @@ -22,6 +22,7 @@ from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.layers.rotary_embedding import get_rope from vllm.model_executor.models.deepseek_v2 import ( + DeepseekV2ForCausalLM, DeepSeekV2FusedQkvAProjLinear, DeepseekV2MLAAttention, DeepseekV2Model, @@ -33,6 +34,7 @@ from vllm.sequence import IntermediateTensors from vllm_ascend.utils import is_mtp_layer +from vllm_ascend.worker.v2 import pp_utils def _should_skip_indexer_init( @@ -298,6 +300,61 @@ def _deepseek_v2_mla_attention_init( DeepseekV2MLAAttention.__init__ = _deepseek_v2_mla_attention_init +# TODO: Retire the legacy aux transport when vLLM 0.28/0.29 are no longer supported. +_original_deepseek_v2_model_init = DeepseekV2Model.__init__ + + +def _patched_deepseek_v2_model_init(self, *args, **kwargs): + _original_deepseek_v2_model_init(self, *args, **kwargs) + # Legacy Spec+PP (0.28/0.29 only): 0.30+ uses the upstream aux relay. + self._use_upstream_aux_relay = kwargs["vllm_config"].use_v2_model_runner and not pp_utils.use_legacy_spec_pp() + if not self._use_upstream_aux_relay: + self.make_empty_intermediate_tensors = pp_utils.make_empty_intermediate_tensors( + self, + self.make_empty_intermediate_tensors, + ) + + +DeepseekV2Model.__init__ = _patched_deepseek_v2_model_init + + +# Legacy Spec+PP (0.28/0.29 only); the upstream branch below is for 0.30+. +if not pp_utils.use_legacy_spec_pp(): + # Release versions do not expose this interface. Reuse only upstream's + # slot bookkeeping; the Ascend forward still owns capture and TP gathering. + from vllm.model_executor.models.interfaces import EagleModelMixin + + for _member in ( + "AUX_HIDDEN_STATE_KEY", + "_aux_slot_base_cached", + "_aux_upstream_total_cached", + "_set_aux_hidden_state_layers", + "_cache_aux_pp_layout", + "pack_local_aux_hidden_states", + "collect_remote_aux_hidden_states", + ): + setattr(DeepseekV2Model, _member, getattr(EagleModelMixin, _member)) + DeepseekV2Model.supports_aux_hidden_states_over_pp = True + + def _set_aux_hidden_state_layers(self, layers: tuple[int, ...]) -> None: + if self.model._use_upstream_aux_relay: + self.model._set_aux_hidden_state_layers(layers) + else: + self.model.aux_hidden_state_layers = layers + + DeepseekV2ForCausalLM.set_aux_hidden_state_layers = _set_aux_hidden_state_layers + + +def _capture_aux_hidden_state(self, aux_hidden_states, layer_id, hidden_states, residual, positions): + if layer_id not in self.aux_hidden_state_layers: + return + aux_hidden_state = hidden_states if residual is None else hidden_states + residual + if aux_hidden_state.shape[0] != positions.shape[0]: + aux_hidden_state = tensor_model_parallel_all_gather(aux_hidden_state, 0) + aux_hidden_state = aux_hidden_state[: positions.shape[0]] + aux_hidden_states.append(aux_hidden_state) + + def _patched_forward( self, input_ids: torch.Tensor | None, @@ -305,7 +362,8 @@ def _patched_forward( intermediate_tensors: IntermediateTensors | None, inputs_embeds: torch.Tensor | None = None, ) -> torch.Tensor | IntermediateTensors: - if get_pp_group().is_first_rank: + pp_group = get_pp_group() + if pp_group.is_first_rank: if inputs_embeds is not None: hidden_states = inputs_embeds else: @@ -313,10 +371,18 @@ def _patched_forward( raise ValueError("Either input_ids or inputs_embeds must be provided to DeepseekV2Model.forward") hidden_states = self.embed_input_ids(input_ids) residual = None + aux_hidden_states: list[torch.Tensor] = [] else: assert intermediate_tensors is not None hidden_states = intermediate_tensors["hidden_states"] residual = intermediate_tensors["residual"] + if self._use_upstream_aux_relay: + aux_hidden_states = self.collect_remote_aux_hidden_states(intermediate_tensors) + else: + aux_hidden_states = pp_utils.get_pp_transport_tensors( + intermediate_tensors, + pp_utils.PPTransportDataType.AUX_HIDDEN_STATES, + ) llama_4_scaling_config = getattr(self.config, "llama_4_scaling", None) llama_4_scaling: torch.Tensor | None @@ -329,21 +395,30 @@ def _patched_forward( else: llama_4_scaling = None - aux_hidden_states = [] + if pp_group.is_first_rank: + _capture_aux_hidden_state(self, aux_hidden_states, 0, hidden_states, residual, positions) for idx, layer in enumerate( islice(self.layers, self.start_layer, self.end_layer), start=self.start_layer, ): - if idx in self.aux_hidden_state_layers: - aux_hidden_state = hidden_states + residual - if aux_hidden_state.shape[0] != positions.shape[0]: - aux_hidden_state = tensor_model_parallel_all_gather(aux_hidden_state, 0) - aux_hidden_state = aux_hidden_state[: positions.shape[0]] - aux_hidden_states.append(aux_hidden_state) hidden_states, residual = layer(positions, hidden_states, residual, llama_4_scaling) - - if not get_pp_group().is_last_rank: - return IntermediateTensors({"hidden_states": hidden_states, "residual": residual}) + # A boundary state belongs to the stage producing it, including end_layer. + _capture_aux_hidden_state(self, aux_hidden_states, idx + 1, hidden_states, residual, positions) + + if not pp_group.is_last_rank: + if self._use_upstream_aux_relay: + return IntermediateTensors( + { + "hidden_states": hidden_states, + "residual": residual, + **self.pack_local_aux_hidden_states(aux_hidden_states), + } + ) + return pp_utils.add_pp_transport_tensors( + IntermediateTensors({"hidden_states": hidden_states, "residual": residual}), + pp_utils.PPTransportDataType.AUX_HIDDEN_STATES, + aux_hidden_states, + ) if hidden_states.shape[0] != positions.shape[0]: combined_states = torch.cat([hidden_states, residual], dim=-1) @@ -353,9 +428,6 @@ def _patched_forward( hidden_states, residual = combined_states.split([hidden_size, hidden_size], dim=-1) residual = residual.contiguous() - if self.end_layer in self.aux_hidden_state_layers: - aux_hidden_states.append(hidden_states + residual) - hidden_states, _ = self.norm(hidden_states, residual) if len(aux_hidden_states) > 0: return hidden_states, aux_hidden_states diff --git a/vllm_ascend/patch/worker/patch_v2/patch_dspark.py b/vllm_ascend/patch/worker/patch_v2/patch_dspark.py index 6827962e84c9..7c3556347732 100644 --- a/vllm_ascend/patch/worker/patch_v2/patch_dspark.py +++ b/vllm_ascend/patch/worker/patch_v2/patch_dspark.py @@ -33,17 +33,21 @@ it upstream to ``load_dspark_model``. """ +from contextlib import AbstractContextManager, nullcontext from types import SimpleNamespace +from typing import cast +from unittest.mock import patch +import vllm.envs as vllm_envs import vllm.model_executor.models.utils as model_utils import vllm.v1.worker.gpu.spec_decode.dspark.speculator as speculator_module import vllm.v1.worker.gpu.spec_decode.dspark.utils as dspark_utils import vllm.v1.worker.gpu.spec_decode.eagle.utils as eagle_utils -from vllm_ascend.utils import vllm_version_is from vllm_ascend.worker.v2.pp_utils import ( bypass_upstream_spec_pp_guard, resolve_spec_pp_support, + use_legacy_spec_pp, ) _original_get_draft_quant_config = model_utils.get_draft_quant_config @@ -59,18 +63,15 @@ def _load_dspark_model_with_target_quant(target_model, vllm_config): draft_model_config = speculative_config.draft_model_config inherits_target_quant = draft_model_config.model == vllm_config.model_config.model spec_pp_support = resolve_spec_pp_support(vllm_config) - bypass_pp_guard = spec_pp_support is not None - original_eagle_should_share = eagle_utils._should_share - if vllm_version_is("0.28.0"): - # Release binds these names at module import; main imports locally. + # Legacy Spec+PP loader bypass (0.28/0.29 only); 0.30+ needs no patching. + bypass_pp_guard = spec_pp_support is not None and use_legacy_spec_pp() + if bypass_pp_guard: + # The legacy loader binds these names at module import. + original_eagle_should_share = eagle_utils._should_share original_get_pp_group = dspark_utils.get_pp_group original_dspark_should_share = dspark_utils._should_share - if inherits_target_quant: - model_utils.get_draft_quant_config = lambda _vllm_config: vllm_config.quant_config - if bypass_pp_guard: - if vllm_version_is("0.28.0"): - single_rank_pp_group = SimpleNamespace(world_size=1) - dspark_utils.get_pp_group = lambda: single_rank_pp_group + single_rank_pp_group = SimpleNamespace(world_size=1) + dspark_utils.get_pp_group = lambda: single_rank_pp_group def should_share(eagle, flag, draft, target): # Non-owning PP ranks expose embed / lm_head as PPMissingLayer @@ -82,20 +83,24 @@ def should_share(eagle, flag, draft, target): return original_eagle_should_share(eagle, flag, draft, target) eagle_utils._should_share = should_share - if vllm_version_is("0.28.0"): - dspark_utils._should_share = should_share + dspark_utils._should_share = should_share + if inherits_target_quant: + model_utils.get_draft_quant_config = lambda _vllm_config: vllm_config.quant_config try: - # get_model also reads the config PP size; keep the draft unsharded. - with bypass_upstream_spec_pp_guard(vllm_config, spec_pp_support): + # Native draft loading already sets PP=1, but still reads the target's + # manual layer partition. Mask that partition on both version paths. + partition_mask = cast(AbstractContextManager[None], nullcontext()) + if spec_pp_support is not None: + partition_mask = patch.object(vllm_envs, "VLLM_PP_LAYER_PARTITION", None) + with partition_mask, bypass_upstream_spec_pp_guard(vllm_config, spec_pp_support): return _original_load_dspark_model(target_model, vllm_config) finally: if inherits_target_quant: model_utils.get_draft_quant_config = _original_get_draft_quant_config if bypass_pp_guard: eagle_utils._should_share = original_eagle_should_share - if vllm_version_is("0.28.0"): - dspark_utils.get_pp_group = original_get_pp_group - dspark_utils._should_share = original_dspark_should_share + dspark_utils.get_pp_group = original_get_pp_group + dspark_utils._should_share = original_dspark_should_share # The speculator binds ``load_dspark_model`` by name at import time, so both diff --git a/vllm_ascend/patch/worker/patch_v2/patch_spec_pp.py b/vllm_ascend/patch/worker/patch_v2/patch_spec_pp.py index 001192848da3..a962c90bd51f 100644 --- a/vllm_ascend/patch/worker/patch_v2/patch_spec_pp.py +++ b/vllm_ascend/patch/worker/patch_v2/patch_spec_pp.py @@ -1,93 +1,171 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved. -"""Speculative decoding support for Model Runner V2 PP.""" +"""0.30 PP sampled-token protocol for release trains (vLLM 0.28/0.29). + +Replaces the deferred-broadcast transport that deadlocked under KV +saturation with async EPLB. Delete once the paired vLLM version ships +the protocol natively. +""" + +from dataclasses import dataclass import numpy as np +import torch from vllm.v1.worker.gpu.buffer_utils import async_copy_to_gpu -_BROADCAST_PATCHED = "_vllm_ascend_spec_pp_broadcast_patched" +_INSTALLED = "_vllm_ascend_upstream_spec_pp_installed" + + +def compute_need_sampled_mask(input_batch): + """Participation gate: pure function of the shared batch (no + rank-local max_seq_len bound). Finished requests are filtered on + the receive side via generation counters.""" + old_computed = input_batch.num_computed_tokens_np + prefill_len = input_batch.prefill_len_np + produces_sample = old_computed + input_batch.num_scheduled_tokens >= prefill_len + return produces_sample if produces_sample.any() else None + +@dataclass +class SpecPPPendingRecv: + """Per-step slot: release fields plus the received draft rows.""" -def install_spec_pp_token_broadcast(pp_handler, req_states) -> None: - """Send accepted and next-draft tokens through the same V2 PP slot.""" - if getattr(pp_handler, _BROADCAST_PATCHED, False): + event: torch.cuda.Event + sampled_tokens: torch.Tensor # [num_reqs, max_sample_len] + num_sampled: torch.Tensor # [num_reqs] + num_rejected: torch.Tensor # [num_reqs] + idx_mapping: torch.Tensor # [num_reqs] + idx_mapping_np: np.ndarray # [num_reqs] + need_sampled_mask: np.ndarray # [num_reqs] + gen_at_receive_np: np.ndarray # [num_reqs] + draft_tokens: torch.Tensor | None = None # [num_reqs, num_speculative_steps] + + +def install_upstream_spec_pp_protocol(pp_handler, req_states, num_speculative_steps) -> None: + """Bind the 0.30 broadcast/receive/consume methods onto the release + PPHandler, adapted to fetch draft rows from ``req_states``.""" + if getattr(pp_handler, _INSTALLED, False): return - max_sample_len = pp_handler.max_sample_len - draft_width = max_sample_len - 1 - token_payload_width = max_sample_len + draft_width - original_get_prev_sampled_outputs = pp_handler.get_prev_sampled_outputs - original_broadcast = pp_handler.broadcast - pending_send = None - - def get_prev_sampled_outputs(): - slot = pp_handler.queue[0] if pp_handler.queue else None - outputs = original_get_prev_sampled_outputs() - if outputs is None: - return None - assert slot is not None - token_payload = outputs["sampled_tokens"] - outputs["sampled_tokens"] = token_payload[:, :max_sample_len] - draft_tokens = token_payload[:, max_sample_len:] + device = pp_handler.device + captured_batch = [None] - # Preserve valid rows on CPU; NPU bool indexing lowers to NonzeroV2. - freed = pp_handler.req_idx_gen_np[slot.idx_mapping_np] != slot.gen_at_receive_np - exclude_mask = freed | ~slot.need_sampled_mask - if exclude_mask.any(): - valid_rows = np.flatnonzero(~exclude_mask) - update_indices = np.stack( - (valid_rows, slot.idx_mapping_np[valid_rows]), - ) - draft_rows, draft_req_indices = async_copy_to_gpu( - update_indices, - device=pp_handler.device, - ).unbind(dim=0) - req_states.draft_tokens.index_copy_( - 0, - draft_req_indices, - draft_tokens.index_select(0, draft_rows), - ) - else: - req_states.draft_tokens.index_copy_( - 0, - outputs["idx_mapping"], - draft_tokens, - ) - return outputs + def receive(input_batch): + assert not pp_handler.is_last_rank + need_sampled_mask = compute_need_sampled_mask(input_batch) + if need_sampled_mask is None: + return False - def broadcast(sampled_token_ids, num_sampled, num_rejected, input_batch): - nonlocal pending_send - assert pp_handler.is_last_rank - if pending_send is not None: - raise RuntimeError("Speculative PP already has a pending sampled-token broadcast.") - pending_send = ( - sampled_token_ids, + gen_at_receive_np = pp_handler.req_idx_gen_np[input_batch.idx_mapping_np] + + num_reqs = input_batch.num_reqs + with torch.cuda.stream(pp_handler.broadcast_stream): + pp_handler.broadcast_stream.wait_stream(pp_handler.main_stream) + sampled_tokens = torch.empty(num_reqs, pp_handler.max_sample_len, dtype=torch.int64, device=device) + combined = torch.empty(2, num_reqs, dtype=torch.int32, device=device) + torch.distributed.broadcast(sampled_tokens, src=pp_handler.last_rank, group=pp_handler.broadcast_group) + torch.distributed.broadcast(combined, src=pp_handler.last_rank, group=pp_handler.broadcast_group) + draft_tokens = None + if num_speculative_steps > 0: + draft_tokens = torch.empty(num_reqs, num_speculative_steps, dtype=torch.int64, device=device) + torch.distributed.broadcast(draft_tokens, src=pp_handler.last_rank, group=pp_handler.broadcast_group) + event = pp_handler.broadcast_stream.record_event() + num_sampled, num_rejected = combined.unbind(dim=0) + sampled_tokens.record_stream(pp_handler.main_stream) + combined.record_stream(pp_handler.main_stream) + if draft_tokens is not None: + draft_tokens.record_stream(pp_handler.main_stream) + pp_handler.queue[-1] = SpecPPPendingRecv( + event, + sampled_tokens, num_sampled, num_rejected, - input_batch, + input_batch.idx_mapping, + input_batch.idx_mapping_np, + need_sampled_mask, + gen_at_receive_np, + draft_tokens, ) + return bool(need_sampled_mask.all()) - def broadcast_draft_tokens(): - nonlocal pending_send - if pending_send is None: + def broadcast(sampled_token_ids, num_sampled, num_rejected, input_batch): + assert pp_handler.is_last_rank + if compute_need_sampled_mask(input_batch) is None: + captured_batch[0] = None return + captured_batch[0] = input_batch - sampled_token_ids, num_sampled, num_rejected, input_batch = pending_send - pending_send = None - num_reqs = input_batch.num_reqs - draft_tokens = req_states.draft_tokens[input_batch.idx_mapping] - token_payload = sampled_token_ids.new_zeros((num_reqs, token_payload_width)) - token_payload[:, : sampled_token_ids.shape[1]].copy_(sampled_token_ids) - token_payload[:, max_sample_len:].copy_(draft_tokens) - original_broadcast( - token_payload, - num_sampled, - num_rejected, - input_batch, + assert sampled_token_ids.dtype == torch.int64 + with torch.cuda.stream(pp_handler.broadcast_stream): + pp_handler.broadcast_stream.wait_stream(pp_handler.main_stream) + send_tokens = torch.nn.functional.pad( + sampled_token_ids, + (0, pp_handler.max_sample_len - sampled_token_ids.shape[-1]), + ) + torch.distributed.broadcast( + send_tokens.contiguous(), + src=pp_handler.last_rank, + group=pp_handler.broadcast_group, + ) + combined = torch.stack((num_sampled, num_rejected), dim=0) + torch.distributed.broadcast(combined, src=pp_handler.last_rank, group=pp_handler.broadcast_group) + for tensor in (sampled_token_ids, num_sampled, num_rejected): + tensor.record_stream(pp_handler.broadcast_stream) + + def broadcast_drafts(draft_tokens=None, input_batch=None): + assert pp_handler.is_last_rank + if input_batch is None: + input_batch = captured_batch[0] + if input_batch is None or compute_need_sampled_mask(input_batch) is None: + return + if draft_tokens is None: + draft_tokens = req_states.draft_tokens[input_batch.idx_mapping] + with torch.cuda.stream(pp_handler.broadcast_stream): + pp_handler.broadcast_stream.wait_stream(pp_handler.main_stream) + send = draft_tokens.contiguous() + input_batch.idx_mapping.record_stream(pp_handler.broadcast_stream) + torch.distributed.broadcast(send, src=pp_handler.last_rank, group=pp_handler.broadcast_group) + + def get_prev_sampled_outputs(draft_tokens_to_update=None): + if not pp_handler.queue: + return None + slot = pp_handler.queue.popleft() + pp_handler.queue.append(None) + if slot is None: + return None + + freed = pp_handler.req_idx_gen_np[slot.idx_mapping_np] != slot.gen_at_receive_np + exclude_mask = freed | ~slot.need_sampled_mask + idx_mapping = slot.idx_mapping + if exclude_mask.any(): + if exclude_mask.all(): + return None + idx_mapping_np = np.where(exclude_mask, -1, slot.idx_mapping_np) + idx_mapping = async_copy_to_gpu(idx_mapping_np, device=device) + + pp_handler.main_stream.wait_event(slot.event) + if draft_tokens_to_update is None: + draft_tokens_to_update = req_states.draft_tokens + if slot.draft_tokens is not None and draft_tokens_to_update is not None: + draft_tokens = slot.draft_tokens + draft_idx_mapping = slot.idx_mapping + if exclude_mask.any(): + keep = ~exclude_mask + keep_t = torch.as_tensor(keep, device=device) + draft_tokens = draft_tokens[keep_t] + draft_idx_mapping = async_copy_to_gpu(slot.idx_mapping_np[keep], device=device) + draft_tokens_to_update[draft_idx_mapping] = draft_tokens + + return dict( + sampled_tokens=slot.sampled_tokens, + num_sampled=slot.num_sampled, + num_rejected=slot.num_rejected, + idx_mapping=idx_mapping, ) - pp_handler.max_sample_len = token_payload_width - pp_handler.get_prev_sampled_outputs = get_prev_sampled_outputs + pp_handler.receive = receive pp_handler.broadcast = broadcast - pp_handler.broadcast_draft_tokens = broadcast_draft_tokens - setattr(pp_handler, _BROADCAST_PATCHED, True) + pp_handler.broadcast_drafts = broadcast_drafts + pp_handler.broadcast_draft_tokens = broadcast_drafts + pp_handler.get_prev_sampled_outputs = get_prev_sampled_outputs + setattr(pp_handler, _INSTALLED, True) diff --git a/vllm_ascend/worker/v2/model_runner.py b/vllm_ascend/worker/v2/model_runner.py index 988710204ce8..f9347c39561a 100644 --- a/vllm_ascend/worker/v2/model_runner.py +++ b/vllm_ascend/worker/v2/model_runner.py @@ -71,6 +71,7 @@ bypass_upstream_spec_pp_guard, resolve_spec_pp_support, restore_pp_after_upstream_init, + use_legacy_spec_pp, ) from vllm_ascend.worker.v2.spec_decode import init_speculator from vllm_ascend.worker.v2.spec_decode.eagle.speculator import AscendEagleSpeculator @@ -103,15 +104,16 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): set_potential_max_tokens(vllm_config) parallel_config = vllm_config.parallel_config - # Eagle3/DSpark drafters are rank-local. Hide PP from the upstream - # initializer, then rebuild the skipped PP state. + # Only release versions need PP hidden during upstream initialization. spec_pp_support = resolve_spec_pp_support(vllm_config) with torch_cuda_wrapper(): with bypass_upstream_spec_pp_guard(vllm_config, spec_pp_support) as pp_disabled: super().__init__(vllm_config, device) if pp_disabled: restore_pp_after_upstream_init(self, vllm_config) - self.use_spec_pp = spec_pp_support is not None + # Native PP owns token broadcast/writeback; only releases use our packing. + # Legacy Spec+PP transport (0.28/0.29 only); deleted when 0.30+ is the floor. + self.use_spec_pp = spec_pp_support is not None and use_legacy_spec_pp() # These draft heads consume target aux states collected across PP ranks. if spec_pp_support is not None and spec_pp_support.needs_aux_hidden_states: self.use_aux_hidden_state_outputs = True @@ -145,7 +147,7 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): # init_speculator will return AscendEagleSpeculator when eagle is used. # so here we just call init_speculator to reinitialize speculator. self.speculator: AscendEagleSpeculator | None = None - if self.speculative_config is not None and (not self.use_spec_pp or self.is_last_pp_rank): + if self.speculative_config is not None and self.is_last_pp_rank: self.speculator = init_speculator(self.vllm_config, self.device) # Shared update_stream: main model (ModelAclGraphManager) and draft # (Eagle/DFlash/DSpark AclGraphManager) all use this same stream. @@ -161,13 +163,13 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): vocab_size=self.vocab_size, device=self.device, ) - if self.use_spec_pp and vllm_version_is("0.28.0"): + if self.use_spec_pp: from vllm_ascend.patch.worker.patch_v2.patch_spec_pp import ( - install_spec_pp_token_broadcast, + install_upstream_spec_pp_protocol, ) assert self.pp_handler is not None - install_spec_pp_token_broadcast(self.pp_handler, self.req_states) + install_upstream_spec_pp_protocol(self.pp_handler, self.req_states, self.num_speculative_steps) # AscendInputBuffers has extra `seq_lens_cpu` attribute. # so reinitialize input_buffers here. self.input_buffers: AscendInputBuffers = AscendInputBuffers( @@ -242,9 +244,9 @@ def sample_tokens(self, grammar_output): self._restore_replicated_draft_target_states() output = super().sample_tokens(grammar_output) - if vllm_version_is("0.28.0") and self.use_spec_pp and self.is_last_pp_rank: + if self.use_spec_pp and self.is_last_pp_rank: assert self.pp_handler is not None - self.pp_handler.broadcast_draft_tokens() + self.pp_handler.broadcast_drafts() return output def initialize_kv_cache( diff --git a/vllm_ascend/worker/v2/pp_utils.py b/vllm_ascend/worker/v2/pp_utils.py index 933d1cc636af..7c806a2e46fa 100644 --- a/vllm_ascend/worker/v2/pp_utils.py +++ b/vllm_ascend/worker/v2/pp_utils.py @@ -8,11 +8,18 @@ from enum import Enum from types import MappingProxyType from typing import TYPE_CHECKING, Protocol +from unittest.mock import patch import torch +import vllm +import vllm.envs as vllm_envs +from packaging.version import Version from vllm.config import VllmConfig from vllm.sequence import IntermediateTensors +from vllm_ascend import envs +from vllm_ascend.utils import vllm_version_is + if TYPE_CHECKING: from transformers import PretrainedConfig from vllm.v1.worker.gpu.model_runner import GPUModelRunner @@ -20,6 +27,19 @@ _PP_TRANSPORT_PREFIX = "pp_transport" +def use_legacy_spec_pp() -> bool: + """True when the paired vLLM needs Ascend's legacy Spec+PP transport. + + This is the Spec+PP workaround for vLLM 0.28/0.29 only; vLLM 0.30+ + ships the upstream PP sampled-token protocol natively and skips + every call site gated on this flag. The entire legacy path (this + flag, patch_spec_pp.py, the loader bypasses) will be deleted once + the release trains move to 0.30+.""" + version = Version(envs.VLLM_VERSION or vllm.__version__) + # Also recognize untagged 0.28 CPU builds via the existing version helper. + return version.release[:2] in ((0, 28), (0, 29)) or vllm_version_is("0.28.0") + + class _PPAuxHiddenStateModel(Protocol): config: "PretrainedConfig" start_layer: int @@ -51,7 +71,7 @@ class SpecPPSupport: unsupported_feature="EAGLE3 with pipeline parallelism", ), "dspark": SpecPPSupport( - architectures=frozenset({"DeepseekV4ForCausalLM"}), + architectures=frozenset({"DeepseekV4ForCausalLM", "GlmMoeDsaForCausalLM"}), needs_aux_hidden_states=True, bypass_upstream_pp_guard=True, ), @@ -82,9 +102,9 @@ def bypass_upstream_spec_pp_guard( vllm_config: VllmConfig, support: SpecPPSupport | None, ) -> Iterator[bool]: - """Initialize the upstream runner as PP=1 to bypass its Spec+PP guard.""" + """Bypass the legacy Spec+PP guard, leaving native PP initialization intact.""" bypass_guard = support.bypass_upstream_pp_guard if support is not None else False - if not bypass_guard: + if not bypass_guard or not use_legacy_spec_pp(): yield False return @@ -92,7 +112,9 @@ def bypass_upstream_spec_pp_guard( original_pp_size = parallel_config.pipeline_parallel_size parallel_config.pipeline_parallel_size = 1 try: - yield True + # The unsharded draft must not inherit the target's manual PP split. + with patch.object(vllm_envs, "VLLM_PP_LAYER_PARTITION", None): + yield True finally: parallel_config.pipeline_parallel_size = original_pp_size