""" 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