Files
wehub-resource-sync eec33d25b2
pre-commit / pre-commit (push) Failing after 1s
Build Wheel / build (3.11) (push) Failing after 1s
Build Wheel / build (3.12) (push) Failing after 0s
chore: import upstream snapshot with attribution
2026-07-13 12:29:08 +08:00

313 lines
12 KiB
Python

"""
Tests for Omni config utils. For stability, these tests should largely be
invariant to the specific attributes of vLLM config except in cases where we
explicitly patch values that differ from vLLM.
"""
import argparse
import inspect
from types import SimpleNamespace
from unittest.mock import Mock
import pytest
from omegaconf import OmegaConf
from pydantic import ValidationError
from transformers import PretrainedConfig
from vllm.engine.arg_utils import EngineArgs
from vllm_omni.config.model import OmniModelConfig
from vllm_omni.engine.arg_utils import OmniEngineArgs
from vllm_omni.engine.async_omni_engine import AsyncOmniEngine
from vllm_omni.engine.stage_init_utils import build_engine_args_dict
pytestmark = [pytest.mark.core_model, pytest.mark.cpu]
def test_sync_config_is_omni():
"""Ensure create_model_config gives the right type."""
cfg = OmniEngineArgs().create_model_config()
assert isinstance(cfg, OmniModelConfig)
def test_default_stage_id_is_concrete_int():
"""Ensure `stage_id` stays safe for downstream arithmetic/indexing."""
engine_args = OmniEngineArgs()
assert engine_args.stage_id == 0
assert isinstance(engine_args.stage_id, int)
assert engine_args.log_stats is False
cfg = engine_args.create_model_config()
assert cfg.stage_id == 0
def test_multimodal_kwarg_overrides(mocker):
"""Ensure that overrides in the multimodal config are preserved."""
sig = inspect.signature(OmniEngineArgs)
default_mm_cache = sig.parameters["mm_processor_cache_gb"].default
override_val = default_mm_cache + 1
fake_model_config = SimpleNamespace(
multimodal_config=SimpleNamespace(mm_processor_cache_gb=override_val),
)
def _fake_parent_create_model_config(self):
assert self.mm_processor_cache_gb == override_val
return fake_model_config
mocker.patch.object(EngineArgs, "create_model_config", _fake_parent_create_model_config)
mocker.patch.object(OmniModelConfig, "from_vllm_model_config", side_effect=lambda model_config, **_: model_config)
cfg = OmniEngineArgs(
model="Qwen/Qwen2-VL-2B-Instruct",
mm_processor_cache_gb=override_val,
).create_model_config()
assert cfg.multimodal_config is not None
assert cfg.multimodal_config.mm_processor_cache_gb == override_val
def test_from_vllm_config_validates_invalid_omni_kwargs():
"""Ensure omni-specific field validation catches invalid keys."""
model_config = EngineArgs().create_model_config()
with pytest.raises(ValueError, match="Unexpected omni kwarg"):
OmniModelConfig.from_vllm_model_config(model_config, foo="bar")
def test_from_vllm_config_validates_bad_omni_kwarg_types():
"""Ensure omni-specific field validation catches type errors."""
model_config = EngineArgs().create_model_config()
with pytest.raises(ValidationError):
OmniModelConfig.from_vllm_model_config(model_config, stage_id="not_an_int")
def test_default_all_values_are_initialized():
"""Ensure omni-specific field initializes all fields"""
model_config = EngineArgs().create_model_config()
cfg = OmniModelConfig.from_vllm_model_config(model_config)
# Test a primitive
assert cfg.model_stage == "thinker"
# Test a field initialized with a default factory
assert cfg.stage_connector_config == {
"name": "SharedMemoryConnector",
"extra": {},
}
# Ensure that hf_config is initialized on model_config in the vLLM by ModelConfig's
# __post_init__, and that the hf_config is copied over to the OmniModelConfig;
# we explicitly set this since the field sets init=False
assert isinstance(model_config.hf_config, PretrainedConfig)
assert cfg.hf_config is model_config.hf_config
# Ensure that we can convert it to a string; this will convert
# all attributes, so should raise if we have attributes that are
# not initialized correctly, e.g., due to default factories
str(cfg)
def test_qwen3_tts_codec_frame_rate_patching():
"""Ensure the patch for qwen3 tts is applied correctly when creating the omni config."""
# Create a vLLM ModelConfig
vllm_config = EngineArgs().create_model_config()
# Create a mock talking config with a dummy value for position_id_per_seconds
mock_talker_config = SimpleNamespace()
mock_talker_config.position_id_per_seconds = 12.3
vllm_config.hf_config.talker_config = mock_talker_config
# Ensure creating the config for a Qwen3TTSTalkerForConditionalGenerationARVLLM
# model calls the patch func to apply position_id_per_seconds from the talker
# config to the config's codec_frame_rate_hz
omni_config = OmniModelConfig.from_vllm_model_config(
vllm_config,
model_arch="Qwen3TTSTalkerForConditionalGenerationARVLLM",
)
# Verify codec_frame_rate_hz was patched
assert omni_config.codec_frame_rate_hz == 12.3
def test_from_cli_args_picks_up_stage_configs_path():
"""from_cli_args should pick up stage_configs_path from namespace."""
ns = argparse.Namespace(
model="facebook/opt-125m",
stage_configs_path="/some/path.yaml",
custom_pipeline_args=None,
)
args = OmniEngineArgs.from_cli_args(ns)
assert args.stage_configs_path == "/some/path.yaml"
assert args.custom_pipeline_args is None
def test_qwen3_tts_code2wav_injects_max_position_embeddings(monkeypatch):
"""Ensure Code2Wav mirrors stage max_model_len into nested HF overrides.
Qwen3-TTS Code2Wav is a pure decoder stage whose runtime max_model_len can
legitimately exceed the base checkpoint's default text max length. Recent
vLLM validates these values during ModelConfig creation, so we inject
``talker_config.max_position_embeddings`` before delegating to vLLM.
"""
captured: dict[str, object] = {}
baseline_config = Mock()
def fake_create_model_config(self):
captured["hf_overrides"] = self.hf_overrides
return baseline_config
monkeypatch.setattr(EngineArgs, "create_model_config", fake_create_model_config)
monkeypatch.setattr(
OmniModelConfig,
"from_vllm_model_config",
classmethod(lambda cls, model_config, **omni_kwargs: model_config),
)
OmniEngineArgs(
model_arch="Qwen3TTSCode2Wav",
max_model_len=65536,
).create_model_config()
assert captured["hf_overrides"] == {
"architectures": ["Qwen3TTSCode2Wav"],
"talker_config": {
"max_position_embeddings": 65536,
},
}
def test_stage_specific_text_config_override():
"""Stage swap must refresh hf_text_config, dependent attrs, and model_arch_config."""
vllm_config = EngineArgs().create_model_config()
vllm_config.disable_sliding_window = True
thinker_mac = vllm_config.model_arch_config
talker_num_heads = max(2, thinker_mac.total_num_attention_heads // 2)
talker_num_kv_heads = max(1, talker_num_heads // 8)
talker_head_dim = 128
stage_text_config = SimpleNamespace(
sliding_window=4096,
attention_chunk_size=2048,
max_position_embeddings=4096,
num_attention_heads=talker_num_heads,
num_key_value_heads=talker_num_kv_heads,
head_dim=talker_head_dim,
hidden_size=talker_num_heads * talker_head_dim,
vocab_size=thinker_mac.vocab_size,
num_hidden_layers=4,
)
vllm_config.hf_text_config = SimpleNamespace()
vllm_config.hf_config.thinker_config = SimpleNamespace(get_text_config=lambda: stage_text_config)
omni_config = OmniModelConfig.from_vllm_model_config(
vllm_config,
hf_config_name="thinker_config",
)
assert omni_config.hf_text_config is stage_text_config
assert omni_config.attention_chunk_size == 2048
assert omni_config.max_model_len == 4096
assert omni_config.hf_text_config.sliding_window is None
stage_mac = omni_config.model_arch_config
assert stage_mac is not thinker_mac
assert stage_mac.total_num_attention_heads == talker_num_heads
assert stage_mac.total_num_kv_heads == talker_num_kv_heads
assert stage_mac.head_size == talker_head_dim
parallel_config = SimpleNamespace(
tensor_parallel_size=1,
pipeline_parallel_size=1,
decode_context_parallel_size=1,
)
assert omni_config.get_num_attention_heads(parallel_config) == talker_num_heads
assert omni_config.get_num_kv_heads(parallel_config) == talker_num_kv_heads
assert omni_config.get_head_size() == talker_head_dim
def test_stage_configs_path_field():
"""OmniEngineArgs with stage_configs_path should construct without error."""
args = OmniEngineArgs(stage_configs_path="/some/path.yaml")
assert args.stage_configs_path == "/some/path.yaml"
def test_strip_single_engine_args():
"""_strip_single_engine_args should remove EngineArgs fields but keep omni fields."""
kwargs = {
# Parent EngineArgs fields — stripped unless explicitly allowlisted
"compilation_config": '{"cudagraph_mode": "FULL_AND_PIECEWISE"}',
"tensor_parallel_size": 4,
"gpu_memory_utilization": 0.9,
"model": "some/model",
# Parent field that should be kept (allowlisted)
"worker_extension_cls": "some.Extension",
# OmniEngineArgs-only / non-engine fields — should pass through
"stage_configs_path": "/path/to/yaml",
"custom_pipeline_args": {"pipeline_class": "my.Pipeline"},
"mode": "text-to-image",
"lora_path": "/some/lora",
}
filtered = AsyncOmniEngine._strip_single_engine_args(kwargs)
# Stripped — parent EngineArgs fields
assert "compilation_config" not in filtered
assert filtered["tensor_parallel_size"] == 4
assert "gpu_memory_utilization" not in filtered
assert "model" not in filtered
# Stripped — orchestrator-level OmniEngineArgs field
assert "stage_configs_path" not in filtered
# Kept
assert filtered["worker_extension_cls"] == "some.Extension"
assert filtered["custom_pipeline_args"] == {"pipeline_class": "my.Pipeline"}
assert filtered["mode"] == "text-to-image"
assert filtered["lora_path"] == "/some/lora"
def test_strip_single_engine_args_model_does_not_trigger_warning(mocker):
"""model is always in kwargs (callers set it via from_cli_args/asdict),
so it should not cause the override warning by itself or appear in it."""
mock_warn = mocker.patch("vllm_omni.engine.async_omni_engine.logger.warning")
# Typical caller kwargs: model is always present, no other parent
# EngineArgs fields are explicitly overridden.
AsyncOmniEngine._strip_single_engine_args(
{
"model": "some/model",
"custom_pipeline_args": {"pipeline_class": "my.Pipeline"},
}
)
mock_warn.assert_not_called()
# When there *are* genuinely surprising overrides alongside model,
# the warning should mention them but not model. Keep-listed fields such as
# tensor_parallel_size are intentionally passed through and should not warn.
AsyncOmniEngine._strip_single_engine_args(
{
"model": "some/model",
"compilation_config": '{"cudagraph_mode": "FULL_AND_PIECEWISE"}',
"tensor_parallel_size": 4,
"custom_pipeline_args": {"pipeline_class": "my.Pipeline"},
}
)
mock_warn.assert_called_once()
warned_args = mock_warn.call_args[0][-1] # the formatted arg list
assert "compilation_config" in warned_args
assert "tensor_parallel_size" not in warned_args
assert "model" not in warned_args
# For https://github.com/vllm-project/vllm-omni/issues/3293
def test_tensor_parallel_size_none_is_handled():
"""Ensure the tensor parallel size of None isn't forwarded."""
engine_args = OmegaConf.create({"stage_id": 0, "engine_args": {"tensor_parallel_size": None}})
args = build_engine_args_dict(
engine_args,
model="snu-aidas/Dynin-Omni",
)
assert isinstance(args, dict)
assert "tensor_parallel_size" not in args