Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 34 additions & 0 deletions tests/ut/models/test_glm_moe_dsa.py
Original file line number Diff line number Diff line change
@@ -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
10 changes: 7 additions & 3 deletions tests/ut/patch/platform/test_patch_pp_mtp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
94 changes: 94 additions & 0 deletions tests/ut/patch/worker/test_patch_deepseek_v2.py
Original file line number Diff line number Diff line change
@@ -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


Expand Down Expand Up @@ -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
57 changes: 57 additions & 0 deletions tests/ut/patch/worker/test_patch_dspark_pp.py
Original file line number Diff line number Diff line change
@@ -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
Loading
Loading