Skip to content
Open
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
38 changes: 38 additions & 0 deletions tests/config/test_omni_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,6 +160,43 @@ def test_non_duplex_deploy_keeps_model_session_capacity_at_one(tmp_path: Path) -
assert [stage.model_config.session_mode for stage in omni_config.stage_configs] == ["turn", "turn"]


def test_model_local_cudagraph_config_is_generation_stage_local(tmp_path: Path) -> None:
deploy_path = tmp_path / "qwen3-tts-vocoder-graph.yaml"
deploy_path.write_text(
"""
async_chunk: false
stages:
- stage_id: 1
model_local_cudagraph:
qwen3_tts.stateless:
capture_bucket_sizes: [150, 325]
""",
encoding="utf-8",
)

omni_config = _from_pipeline_key("qwen3_tts", deploy_config_path=str(deploy_path))

assert omni_config.stage_by_id(0).model_config.model_local_cudagraph is None
assert omni_config.stage_by_id(1).model_config.model_local_cudagraph == {
"qwen3_tts.stateless": {"capture_bucket_sizes": [150, 325]},
}


def test_model_local_cudagraph_config_rejects_ar_stage(tmp_path: Path) -> None:
deploy_path = tmp_path / "qwen3-tts-ar-vocoder-graph.yaml"
deploy_path.write_text(
"""
stages:
- stage_id: 0
model_local_cudagraph: {}
""",
encoding="utf-8",
)

with pytest.raises(ValueError, match="only for LLM_GENERATION"):
_from_pipeline_key("qwen3_tts", deploy_config_path=str(deploy_path))


def test_nested_stage_override_deep_merges_structured_model_config() -> None:
config = _from_pipeline_key(
"cosmos3_policy",
Expand Down Expand Up @@ -982,6 +1019,7 @@ def test_sub_config_fields_match_structured_scopes():
"model_subdir",
"tokenizer_subdir",
"requires_full_payload_input",
"model_local_cudagraph",
"served_model_name",
"allowed_local_media_path",
"allowed_media_domains",
Expand Down
16 changes: 16 additions & 0 deletions tests/engine/test_stage_engine_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,7 @@ def _effective_backend_values(config_cls: type, engine_args: dict) -> dict[str,
"subtalker_sampling_params",
"task_type",
"tokenizer_subdir",
"model_local_cudagraph",
}
)

Expand Down Expand Up @@ -884,3 +885,18 @@ def test_typed_engine_args_match_current_registry_backend_semantics(model_type,
assert typed_effective_args == legacy_effective_args, (
f"{model_type} stage {stage_id} changed effective backend arguments"
)


@pytest.mark.parametrize("config", [None, {"decode": {}}])
def test_create_model_config_projects_model_local_cudagraph(monkeypatch, config):
from vllm_omni.config.model import OmniModelConfig

monkeypatch.setattr(OmniEngineArgs, "_ensure_omni_models_registered", lambda self: None)
monkeypatch.setattr(EngineArgs, "create_model_config", lambda self: types.SimpleNamespace(hf_config=None))
monkeypatch.setattr(OmniModelConfig, "_maybe_override_text_config", lambda self: None)
args = OmniEngineArgs(
model="unused", tokenizer="unused", worker_type="generation", worker_cls="unused", model_local_cudagraph=config
)
result = args.create_model_config()
assert isinstance(result, OmniModelConfig)
assert result.model_local_cudagraph == config
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

from __future__ import annotations

import pytest

from vllm_omni.model_executor.models.interfaces.model_local_cudagraph import (
SupportsModelLocalCUDAGraph,
supports_model_local_cudagraph,
)

pytestmark = [pytest.mark.core_model, pytest.mark.cpu]


def test_capability_discovery_requires_both_declaration_and_provider() -> None:
class Model(SupportsModelLocalCUDAGraph):
supports_model_local_cudagraph = True

def get_model_local_cudagraph_components(self):
return ()

assert supports_model_local_cudagraph(Model())
assert not supports_model_local_cudagraph(object())
103 changes: 103 additions & 0 deletions tests/worker/test_generation_model_local_cudagraph_runner.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project

from contextlib import nullcontext
from types import SimpleNamespace

import pytest
import torch
from vllm.config import CUDAGraphMode

from vllm_omni.worker import gpu_generation_model_runner as generation_runner_module
from vllm_omni.worker.gpu_generation_model_runner import GPUGenerationModelRunner
from vllm_omni.worker.gpu_model_runner import OmniGPUModelRunner

pytestmark = [pytest.mark.core_model, pytest.mark.cpu]


class _GraphModel:
supports_model_local_cudagraph = True

def get_model_local_cudagraph_components(self):
return ()


class _FakeManager:
instances: list["_FakeManager"] = []

def __init__(self, *, vllm_config, device) -> None:
self.vllm_config = vllm_config
self.device = device
self.prepared_model = None
self.cleared = False
self.instances.append(self)

def prepare(self, model) -> None:
self.prepared_model = model

def clear(self) -> None:
self.cleared = True


def _runner(*, enforce_eager: bool, mode: CUDAGraphMode):
runner = object.__new__(GPUGenerationModelRunner)
runner.model_config = SimpleNamespace(enforce_eager=enforce_eager, model_local_cudagraph={"decode": {}})
runner.compilation_config = SimpleNamespace(cudagraph_mode=mode)
runner.vllm_config = SimpleNamespace()
runner.device = torch.device("cpu")
runner.model = _GraphModel()
runner.model_local_cudagraph_manager = None
return runner


def test_load_model_prepares_manager_from_unwrapped_model(monkeypatch) -> None:
_FakeManager.instances.clear()
monkeypatch.setattr(OmniGPUModelRunner, "load_model", lambda self, *args, **kwargs: None)
monkeypatch.setattr(generation_runner_module, "ModelLocalCUDAGraphManager", _FakeManager)
runner = _runner(enforce_eager=False, mode=CUDAGraphMode.FULL)

GPUGenerationModelRunner.load_model(runner)

assert len(_FakeManager.instances) == 1
manager = _FakeManager.instances[0]
assert manager.prepared_model is runner.model
assert runner.model_local_cudagraph_manager is manager


def test_omitted_vocoder_config_keeps_upstream_runner(monkeypatch) -> None:
monkeypatch.setattr(OmniGPUModelRunner, "load_model", lambda self, *args, **kwargs: None)
runner = _runner(enforce_eager=False, mode=CUDAGraphMode.FULL)
runner.model_config.model_local_cudagraph = None
GPUGenerationModelRunner.load_model(runner)
assert runner.model_local_cudagraph_manager is None


def test_enforce_eager_skips_manager_and_shutdown_restores_components(monkeypatch) -> None:
_FakeManager.instances.clear()
monkeypatch.setattr(OmniGPUModelRunner, "load_model", lambda self, *args, **kwargs: None)
monkeypatch.setattr(generation_runner_module, "ModelLocalCUDAGraphManager", _FakeManager)
runner = _runner(enforce_eager=True, mode=CUDAGraphMode.NONE)

GPUGenerationModelRunner.load_model(runner)
assert runner.model_local_cudagraph_manager is None

manager = _FakeManager(vllm_config=runner.vllm_config, device=runner.device)
runner.model_local_cudagraph_manager = manager
shutdown_called = []
monkeypatch.setattr(OmniGPUModelRunner, "shutdown", lambda self: shutdown_called.append(True))
GPUGenerationModelRunner.shutdown(runner)

assert manager.cleared
assert runner.model_local_cudagraph_manager is None
assert shutdown_called == [True]


def test_profile_uses_measured_estimate(monkeypatch):
runner = _runner(enforce_eager=False, mode=CUDAGraphMode.FULL)
manager = SimpleNamespace(profile_memory=lambda: 321)
runner.model_local_cudagraph_manager = manager
monkeypatch.setattr(runner, "_freeze_gc", nullcontext)
monkeypatch.setattr(generation_runner_module, "graph_capture", lambda **kwargs: nullcontext())
monkeypatch.setattr(torch.accelerator, "synchronize", lambda: None)
monkeypatch.setattr(torch.accelerator, "empty_cache", lambda: None)
assert runner.profile_cudagraph_memory() == 321
1 change: 1 addition & 0 deletions tests/worker/test_gpu_generation_model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -131,6 +131,7 @@ def _make_guard_runner():
runner._update_states = lambda scheduler_output: None
runner.synchronize_input_prep = contextlib.nullcontext
runner.attach_omni_connector_output = lambda result: result
runner.model_local_cudagraph_manager = None
return runner


Expand Down
Loading
Loading