Files
enginex-ascend-910-vllm/tests/ut/patch/platform/test_patch_pp_mtp.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

58 lines
1.9 KiB
Python

# SPDX-License-Identifier: Apache-2.0
from types import SimpleNamespace
import pytest
from vllm.config.model import ModelConfig
def test_model_config_validates_local_mtp_drafter_as_single_pp_rank(monkeypatch):
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.runner = "draft"
model_config.model_arch_config = SimpleNamespace(
total_num_attention_heads=1,
architectures=["Qwen3_5MTP"],
)
model_config.multimodal_config = None
parallel_config = SimpleNamespace(
tensor_parallel_size=1,
enable_expert_parallel=False,
pipeline_parallel_size=2,
decode_context_parallel_size=1,
)
ModelConfig.verify_with_parallel_config(model_config, parallel_config)
assert parallel_config.pipeline_parallel_size == 2
def test_model_config_keeps_target_model_pp_validation(monkeypatch):
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.runner = "generate"
model_config.model_arch_config = SimpleNamespace(
total_num_attention_heads=1,
architectures=["UnsupportedForPP"],
)
parallel_config = SimpleNamespace(
tensor_parallel_size=1,
enable_expert_parallel=False,
pipeline_parallel_size=2,
decode_context_parallel_size=1,
)
with pytest.raises(NotImplementedError):
ModelConfig.verify_with_parallel_config(model_config, parallel_config)