diff --git a/tests/config/test_omni_config.py b/tests/config/test_omni_config.py index da69ec5f1a2..c87e13cb6cc 100644 --- a/tests/config/test_omni_config.py +++ b/tests/config/test_omni_config.py @@ -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", @@ -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", diff --git a/tests/engine/test_stage_engine_args.py b/tests/engine/test_stage_engine_args.py index 3e8f0ef87d5..86e666954ed 100644 --- a/tests/engine/test_stage_engine_args.py +++ b/tests/engine/test_stage_engine_args.py @@ -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", } ) @@ -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 diff --git a/tests/model_executor/models/interfaces/test_model_local_cudagraph.py b/tests/model_executor/models/interfaces/test_model_local_cudagraph.py new file mode 100644 index 00000000000..380d06f0e7b --- /dev/null +++ b/tests/model_executor/models/interfaces/test_model_local_cudagraph.py @@ -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()) diff --git a/tests/worker/test_generation_model_local_cudagraph_runner.py b/tests/worker/test_generation_model_local_cudagraph_runner.py new file mode 100644 index 00000000000..9ec182d1340 --- /dev/null +++ b/tests/worker/test_generation_model_local_cudagraph_runner.py @@ -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 diff --git a/tests/worker/test_gpu_generation_model_runner.py b/tests/worker/test_gpu_generation_model_runner.py index 0cd0a1ab33b..06d272dd03d 100644 --- a/tests/worker/test_gpu_generation_model_runner.py +++ b/tests/worker/test_gpu_generation_model_runner.py @@ -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 diff --git a/tests/worker/test_model_local_cudagraph_manager.py b/tests/worker/test_model_local_cudagraph_manager.py new file mode 100644 index 00000000000..5629ec54984 --- /dev/null +++ b/tests/worker/test_model_local_cudagraph_manager.py @@ -0,0 +1,974 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project + +from __future__ import annotations + +import logging +from collections import Counter, OrderedDict +from collections.abc import Callable, Set +from contextlib import contextmanager, nullcontext +from dataclasses import dataclass +from types import SimpleNamespace +from typing import Any, NamedTuple, cast +from unittest.mock import patch + +import pytest +import torch +from vllm.platforms import current_platform + +from vllm_omni.model_executor.models.interfaces.model_local_cudagraph import ( + BaseModelLocalCUDAGraphRoutine, + ModelLocalCaptureMode, + ModelLocalCUDAGraphComponent, + ModelLocalCUDAGraphDescriptor, + ModelLocalGraphHandle, + ModelLocalRuntimeKey, + ModelLocalRuntimeResolution, + SupportsModelLocalCUDAGraph, +) +from vllm_omni.worker.model_local_cudagraph_manager import ( + ManagedComponent, + ModelLocalCUDAGraphEntry, + ModelLocalCUDAGraphManager, + ModelLocalGraphStatsSink, + clone_tensor_tree, +) + +pytestmark = [pytest.mark.core_model, pytest.mark.cpu] + + +class _NestedOutput(NamedTuple): + tensor: torch.Tensor + metadata: dict[str, object] + + +def test_clone_tensor_tree_clones_supported_tensor_containers() -> None: + tensor = torch.tensor([1.0, 2.0]) + output = { + "tensor": tensor, + "tuple": (tensor, [tensor]), + "namedtuple": _NestedOutput(tensor, {"label": "audio"}), + } + + cloned = clone_tensor_tree(output) + assert isinstance(cloned, dict) + + assert isinstance(cloned["namedtuple"], _NestedOutput) + assert cloned["namedtuple"].metadata == {"label": "audio"} + cloned_tensors = ( + cloned["tensor"], + cloned["tuple"][0], + cloned["tuple"][1][0], + cloned["namedtuple"].tensor, + ) + for cloned_tensor in cloned_tensors: + torch.testing.assert_close(cloned_tensor, tensor) + assert cloned_tensor.data_ptr() != tensor.data_ptr() + + +@dataclass +class _Buffers: + input: torch.Tensor + output: torch.Tensor + + +class _Graph: + def __init__(self, buffers: _Buffers, *, fail: bool = False) -> None: + self.buffers = buffers + self.fail = fail + + def reset(self) -> None: + self.reset_called = True + + def replay(self) -> None: + if self.fail: + raise RuntimeError("replay failed") + self.buffers.output.copy_(self.buffers.input * 2) + + +class _Routine(BaseModelLocalCUDAGraphRoutine): + def __init__(self) -> None: + self.eager_calls = 0 + self.validate_calls = 0 + self._runnable: Callable[[torch.Tensor], torch.Tensor] = lambda value: value * 2 + + @property + def runnable(self) -> Callable[..., Any]: + return self._runnable + + def eager_call(self, value: torch.Tensor) -> torch.Tensor: + self.eager_calls += 1 + return self._runnable(value) + + def validate_runtime_inputs(self, args: tuple[Any, ...], kwargs: dict[str, Any]) -> None: + self.validate_calls += 1 + if kwargs or len(args) != 1 or not isinstance(args[0], torch.Tensor): + raise ValueError("invalid invocation") + if args[0].numel() == 0: + raise ValueError("empty input") + + def resolve_runtime( + self, + args: tuple[Any, ...], + kwargs: dict[str, Any], + available: Set[ModelLocalCUDAGraphDescriptor], + ) -> ModelLocalRuntimeResolution: + del kwargs + size = int(args[0].numel()) + descriptor = min( + (item for item in available if isinstance(item.variant, int) and item.variant >= size), + key=lambda item: item.variant if isinstance(item.variant, int) else 0, + default=None, + ) + return ModelLocalRuntimeResolution(ModelLocalRuntimeKey(size), descriptor) + + def allocate_buffers(self, descriptor: ModelLocalCUDAGraphDescriptor, device: torch.device) -> _Buffers: + assert isinstance(descriptor.variant, int) + size = descriptor.variant + return _Buffers(torch.zeros(size, device=device), torch.zeros(size, device=device)) + + def forward_for_capture(self, buffers: object) -> torch.Tensor: + assert isinstance(buffers, _Buffers) + buffers.output.copy_(buffers.input * 2) + return buffers.output + + def copy_runtime_inputs( + self, + args: tuple[Any, ...], + kwargs: dict[str, Any], + buffers: object, + ) -> None: + del kwargs + assert isinstance(buffers, _Buffers) + buffers.input.zero_() + buffers.input[: args[0].numel()].copy_(args[0]) + + def output_after_replay( + self, + args: tuple[Any, ...], + kwargs: dict[str, Any], + buffers: object, + captured_output: object, + ) -> torch.Tensor: + del kwargs, buffers + assert isinstance(captured_output, torch.Tensor) + return captured_output[: args[0].numel()] + + +class _LifecycleRoutine(_Routine): + def __init__(self, *, fail_forward: bool = False, context_only: bool = False) -> None: + super().__init__() + self.events: list[str] = [] + self.fail_forward = fail_forward + self.context_only = context_only + + def prepare_for_capture(self, buffers: object) -> None: + del buffers + if self.context_only: + raise AssertionError("manager called prepare_for_capture directly") + self.events.append("prepare") + + def forward_for_capture(self, buffers: object) -> torch.Tensor: + if self.context_only: + return _Routine.forward_for_capture(self, buffers) + self.events.append("forward") + if self.fail_forward: + raise RuntimeError("capture forward failed") + return super().forward_for_capture(buffers) + + def after_capture(self, buffers: object) -> None: + del buffers + if self.context_only: + raise AssertionError("manager called after_capture directly") + self.events.append("after") + + @contextmanager + def capture_context(self, descriptor: ModelLocalCUDAGraphDescriptor, buffers: object): + if self.context_only: + del descriptor, buffers + self.events.append("context-enter") + try: + yield + finally: + self.events.append("context-exit") + return + with super().capture_context(descriptor, buffers): + yield + + +class _ScopedLifecycleRoutine(_LifecycleRoutine): + @contextmanager + def capture_context(self, descriptor: ModelLocalCUDAGraphDescriptor, buffers: object): + self.events.append("scope-enter") + try: + with super().capture_context(descriptor, buffers): + yield + finally: + self.events.append("scope-exit") + + +class _TestManager(ModelLocalCUDAGraphManager): + def __init__(self, *, config: dict[str, Any] | None = None, log_stats: bool = False) -> None: + vllm_config = SimpleNamespace( + model_config=SimpleNamespace(model_local_cudagraph=config), + compilation_config=SimpleNamespace(cudagraph_num_of_warmups=0), + observability_config=SimpleNamespace(cudagraph_metrics=log_stats), + ) + super().__init__(vllm_config=vllm_config, device=torch.device("cpu")) + self._use_default_component_config = config is None + self.components_during_capture: list[tuple[bool, ...]] = [] + self.fail_replay_for: set[tuple[str, object]] = set() + self.fail_capture_for: set[tuple[str, object]] = set() + self.capture_attempts: Counter[tuple[str, object]] = Counter() + self.capture_pools: list[object | None] = [] + + def prepare(self, model: SupportsModelLocalCUDAGraph) -> None: + if self._use_default_component_config: + self.config = {component.component_id: {} for component in model.get_model_local_cudagraph_components()} + super().prepare(model) + + def capture_entry( + self, + component: ModelLocalCUDAGraphComponent, + descriptor: ModelLocalCUDAGraphDescriptor, + *, + graph_pool: object | None = None, + ) -> ModelLocalCUDAGraphEntry | None: + self.capture_pools.append(graph_pool) + self.components_during_capture.append(tuple(item._bound_handle is not None for item in self.components)) + key = (component.component_id, descriptor.variant) + self.capture_attempts[key] += 1 + if not component.validate_descriptor(descriptor): + return None + if key in self.fail_capture_for: + return None + buffers = component.routine.allocate_buffers(descriptor, self.device) + assert isinstance(buffers, _Buffers) + output = component.routine.forward_for_capture(buffers) + graph = _Graph( + buffers, + fail=(component.component_id, descriptor.variant) in self.fail_replay_for, + ) + return ModelLocalCUDAGraphEntry( + descriptor=descriptor, + graph=cast(torch.cuda.CUDAGraph, graph), + buffers=buffers, + captured_output=output, + ) + + +class _Model: + supports_model_local_cudagraph = True + model_local_cudagraph_shared_config_keys = frozenset({"shared_shape_policy"}) + + def __init__(self, components: tuple[ModelLocalCUDAGraphComponent, ...]) -> None: + self.components = components + + def get_model_local_cudagraph_components(self) -> tuple[ModelLocalCUDAGraphComponent, ...]: + return self.components + + +def _component( + component_id: str, + *sizes: int, + capture_order_key: Callable[[ModelLocalCUDAGraphDescriptor], Any] | None = None, + capture_mode: ModelLocalCaptureMode = ModelLocalCaptureMode.PRECAPTURE, +) -> tuple[ModelLocalCUDAGraphComponent, _Routine]: + routine = _Routine() + component = ModelLocalCUDAGraphComponent( + component_id, + routine, + [ModelLocalCUDAGraphDescriptor(size) for size in sizes], + supported_config_keys=frozenset({"bucket_policy"}), + capture_order_key=capture_order_key, + capture_mode=capture_mode, + ) + return component, routine + + +def _manager_for_capture(*, warmups: int) -> ModelLocalCUDAGraphManager: + vllm_config = SimpleNamespace( + model_config=SimpleNamespace(model_local_cudagraph=None), + compilation_config=SimpleNamespace(cudagraph_num_of_warmups=warmups), + ) + return ModelLocalCUDAGraphManager(vllm_config=vllm_config, device=torch.device("cpu")) + + +def _mock_cuda_capture(monkeypatch) -> None: + monkeypatch.setattr(torch.cuda, "current_stream", lambda _device: SimpleNamespace(synchronize=lambda: None)) + monkeypatch.setattr(torch.cuda, "CUDAGraph", object) + monkeypatch.setattr(torch.cuda, "graph", lambda *_args, **_kwargs: nullcontext()) + + +def test_handle_exposes_runtime_call_and_read_only_graph_coverage() -> None: + handle = ModelLocalGraphHandle(lambda value, *, offset=0: value + offset) + + assert handle(2, offset=3) == 5 + assert handle.available_descriptors == frozenset() + assert not hasattr(handle, "capture") + assert not hasattr(handle, "replay") + assert not hasattr(handle, "entries") + + +def test_component_descriptor_validation_defaults_to_capture() -> None: + component, _ = _component("decode", 2) + + assert component.validate_descriptor(ModelLocalCUDAGraphDescriptor(2)) + + +def test_capture_context_wraps_each_warmup_and_capture_forward(monkeypatch) -> None: + routine = _LifecycleRoutine() + component = ModelLocalCUDAGraphComponent("decode", routine, [ModelLocalCUDAGraphDescriptor(2)]) + manager = _manager_for_capture(warmups=2) + _mock_cuda_capture(monkeypatch) + + manager.capture_entry(component, ModelLocalCUDAGraphDescriptor(2), graph_pool=object()) + + assert routine.events == ["prepare", "forward", "after"] * 3 + + +def test_capture_context_runs_cleanup_when_forward_raises() -> None: + routine = _LifecycleRoutine(fail_forward=True) + component = ModelLocalCUDAGraphComponent("decode", routine, [ModelLocalCUDAGraphDescriptor(2)]) + manager = _manager_for_capture(warmups=1) + + with pytest.raises(RuntimeError, match="capture forward failed"): + manager.capture_entry(component, ModelLocalCUDAGraphDescriptor(2)) + + assert routine.events == ["prepare", "forward", "after"] + + +def test_graph_capture_failure_resets_graph_and_uses_warmup_stream(monkeypatch) -> None: + routine = _LifecycleRoutine() + component = ModelLocalCUDAGraphComponent("decode", routine, [ModelLocalCUDAGraphDescriptor(2)]) + manager = _manager_for_capture(warmups=1) + capture_stream = SimpleNamespace(synchronize=lambda: None) + graph = SimpleNamespace(reset_called=False) + + def reset_graph(): + graph.reset_called = True + + graph.reset = reset_graph + graph_streams = [] + + def graph_context(*_args, **kwargs): + graph_streams.append(kwargs["stream"]) + return nullcontext() + + original_forward = routine.forward_for_capture + forward_calls = 0 + + def forward(buffers): + nonlocal forward_calls + forward_calls += 1 + if forward_calls == 2: + routine.events.append("forward") + raise RuntimeError("graph capture failed") + return original_forward(buffers) + + monkeypatch.setattr(routine, "forward_for_capture", forward) + monkeypatch.setattr(torch.cuda, "current_stream", lambda _device: capture_stream) + monkeypatch.setattr(torch.cuda, "CUDAGraph", lambda: graph) + monkeypatch.setattr(torch.cuda, "graph", graph_context) + + with pytest.raises(RuntimeError, match="graph capture failed"): + manager.capture_entry(component, ModelLocalCUDAGraphDescriptor(2), graph_pool=object()) + + assert routine.events == ["prepare", "forward", "after"] * 2 + assert graph_streams == [capture_stream] + assert graph.reset_called + + +def test_manager_uses_capture_context_instead_of_prepare_or_after(monkeypatch) -> None: + routine = _LifecycleRoutine(context_only=True) + component = ModelLocalCUDAGraphComponent("decode", routine, [ModelLocalCUDAGraphDescriptor(2)]) + manager = _manager_for_capture(warmups=1) + _mock_cuda_capture(monkeypatch) + + manager.capture_entry(component, ModelLocalCUDAGraphDescriptor(2), graph_pool=object()) + + assert routine.events == ["context-enter", "context-exit"] * 2 + + +def test_capture_context_override_composes_default_lifecycle(monkeypatch) -> None: + routine = _ScopedLifecycleRoutine() + component = ModelLocalCUDAGraphComponent("decode", routine, [ModelLocalCUDAGraphDescriptor(2)]) + manager = _manager_for_capture(warmups=1) + _mock_cuda_capture(monkeypatch) + + manager.capture_entry(component, ModelLocalCUDAGraphDescriptor(2), graph_pool=object()) + + assert routine.events == ["scope-enter", "prepare", "forward", "after", "scope-exit"] * 2 + + +def test_stats_sink_bounds_detail_items_but_preserves_aggregate_counters() -> None: + sink = ModelLocalGraphStatsSink(enabled=True, max_log_items=2) + for variant in (1, 2, 3): + resolution = ModelLocalRuntimeResolution( + runtime_key=ModelLocalRuntimeKey(variant), + descriptor=ModelLocalCUDAGraphDescriptor(variant), + ) + sink.record("hit", "decode", resolution) + + snapshot = sink.snapshot() + assert snapshot["calls"] == {"decode": 3} + assert snapshot["outcomes"] == {("decode", "hit"): 3} + assert snapshot["descriptors"] == {("decode", 2): 1, ("decode", 3): 1} + assert snapshot["runtime_keys"] == {("decode", 2): 1, ("decode", 3): 1} + + +def test_stats_sink_logs_every_100_component_calls() -> None: + sink = ModelLocalGraphStatsSink(enabled=True) + resolution = ModelLocalRuntimeResolution(ModelLocalRuntimeKey(2), ModelLocalCUDAGraphDescriptor(2)) + + with patch.object(logging.Logger, "info") as log_info: + for _ in range(99): + sink.record("hit", "decode", resolution) + log_info.assert_not_called() + sink.record("hit", "decode", resolution) + + log_info.assert_called_once() + assert log_info.call_args.args[1] == 100 + assert log_info.call_args.args[2]["outcomes"] == {("decode", "hit"): 100} + + +def test_runtime_lazy_capture_logs_only_new_entries(monkeypatch) -> None: + component, _ = _component("decode", 2) + manager = _TestManager() + descriptor = ModelLocalCUDAGraphDescriptor(3) + fake_buffers = _Buffers(torch.zeros(1), torch.zeros(1)) + fake_entry = ModelLocalCUDAGraphEntry( + descriptor=descriptor, + graph=cast(torch.cuda.CUDAGraph, _Graph(fake_buffers)), + buffers=fake_buffers, + captured_output=fake_buffers.output, + ) + managed = ManagedComponent( + component=component, + entries=OrderedDict(), + capture_mode=ModelLocalCaptureMode.PRECAPTURE_LAZY, + max_graphs=1, + ) + calls: list[ModelLocalCUDAGraphDescriptor] = [] + + def capture(managed_component, requested_descriptor): + assert managed_component is managed + calls.append(requested_descriptor) + return fake_entry + + monkeypatch.setattr(manager, "_capture_and_register", capture) + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + monkeypatch.setattr(manager, "_runtime_capture_scope", nullcontext) + with patch.object(logging.Logger, "info") as log_info: + result = manager._runtime_capture_and_register(managed, descriptor) + + assert result is fake_entry + assert calls == [descriptor] + log_info.assert_called_once_with( + "Lazy-captured model-local CUDA Graph Component %s Descriptor %r", + "decode", + descriptor, + ) + + managed.entries[descriptor] = fake_entry + calls.clear() + with patch.object(logging.Logger, "info") as log_info: + result = manager._runtime_capture_and_register(managed, descriptor) + + assert result is fake_entry + assert calls == [] + log_info.assert_not_called() + + +def test_capture_binds_only_after_all_components_are_captured_and_clear_restores_eager() -> None: + first, first_routine = _component("first", 4) + second, _ = _component("second", 8) + manager = _TestManager() + manager.prepare(_Model((first, second))) + + manager.capture_and_bind() + + assert manager.components_during_capture + assert all(not any(bound) for bound in manager.components_during_capture) + assert set(manager.managed_components) == {"first", "second"} + assert isinstance(first._bound_handle, ModelLocalGraphHandle) + assert isinstance(second._bound_handle, ModelLocalGraphHandle) + assert first.available_descriptors == frozenset({ModelLocalCUDAGraphDescriptor(4)}) + assert second.available_descriptors == frozenset({ModelLocalCUDAGraphDescriptor(8)}) + value = torch.tensor([1.0, 2.0]) + first_output = first(value) + assert torch.equal(first_output, value * 2) + assert first_routine.eager_calls == 0 + + # clone_output=True prevents the next replay from overwriting a retained result. + retained = first_output.clone() + first(torch.tensor([4.0, 5.0])) + assert torch.equal(first_output, retained) + + manager.clear() + assert first._bound_handle is None + assert first.available_descriptors == frozenset() + assert torch.equal(first(value), value * 2) + assert first_routine.eager_calls == 1 + + +def test_unified_fallback_uses_segmented_graph_after_all_components_bind() -> None: + segmented, segmented_routine = _component("segmented", 4) + + class _UnifiedRoutine(_Routine): + def eager_call(self, value: torch.Tensor) -> torch.Tensor: + self.eager_calls += 1 + return segmented(value) + 1 + + def forward_for_capture(self, buffers: object) -> torch.Tensor: + assert isinstance(buffers, _Buffers) + buffers.output.copy_(segmented(buffers.input) + 1) + return buffers.output + + unified_routine = _UnifiedRoutine() + unified = ModelLocalCUDAGraphComponent("unified", unified_routine, [ModelLocalCUDAGraphDescriptor(2)]) + manager = _TestManager() + manager.prepare(_Model((segmented, unified))) + manager.capture_and_bind() + + # Unified capture called segmented after its capture, but before binding. + assert segmented_routine.eager_calls == 1 + assert all(not any(bound) for bound in manager.components_during_capture) + + segmented_entry = manager.managed_components["segmented"].entries[ModelLocalCUDAGraphDescriptor(4)] + replay_calls = 0 + original_replay = segmented_entry.graph.replay + + def replay() -> None: + nonlocal replay_calls + replay_calls += 1 + original_replay() + + segmented_entry.graph.replay = replay + value = torch.tensor([1.0, 2.0, 3.0]) # Misses the unified graph, fits the segmented graph. + torch.testing.assert_close(unified(value), value * 2 + 1) + assert unified_routine.eager_calls == 1 + assert segmented_routine.eager_calls == 1 + assert replay_calls == 1 + + +def test_coverage_miss_falls_back_but_validation_and_replay_errors_propagate() -> None: + component, routine = _component("decode", 2) + manager = _TestManager(log_stats=True) + manager.fail_replay_for.add(("decode", 2)) + manager.prepare(_Model((component,))) + manager.capture_and_bind() + + fallback_input = torch.tensor([1.0, 2.0, 3.0]) + assert torch.equal(component(fallback_input), fallback_input * 2) + assert routine.eager_calls == 1 + + with pytest.raises(ValueError, match="empty input"): + component(torch.tensor([])) + assert routine.eager_calls == 1 + + with pytest.raises(RuntimeError, match="replay failed"): + component(torch.tensor([1.0])) + assert routine.eager_calls == 1 + outcomes = manager.stats_sink.snapshot()["outcomes"] + assert isinstance(outcomes, dict) + assert outcomes[("decode", "fallback")] == 1 + assert outcomes[("decode", "replay_error")] == 1 + + +def test_copy_and_postprocess_errors_propagate_without_eager_retry() -> None: + component, routine = _component("decode", 2) + manager = _TestManager(log_stats=True) + manager.prepare(_Model((component,))) + manager.capture_and_bind() + + def fail_copy(*_args, **_kwargs): + raise RuntimeError("copy failed") + + routine.copy_runtime_inputs = fail_copy + with pytest.raises(RuntimeError, match="copy failed"): + component(torch.tensor([1.0])) + assert routine.eager_calls == 0 + + component, routine = _component("decode-postprocess", 2) + manager = _TestManager(log_stats=True) + manager.prepare(_Model((component,))) + manager.capture_and_bind() + + def fail_output(*_args, **_kwargs): + raise RuntimeError("postprocess failed") + + routine.output_after_replay = fail_output + with pytest.raises(RuntimeError, match="postprocess failed"): + component(torch.tensor([1.0])) + assert routine.eager_calls == 0 + + +def test_config_validation_catches_unknown_component_and_extension_keys() -> None: + component, _ = _component("decode", 2) + + manager = _TestManager(config={"unknown": 1}) + with pytest.raises(ValueError, match="Unknown model-local CUDA Graph component"): + manager.prepare(_Model((component,))) + + manager = _TestManager(config={"missing": {}}) + with pytest.raises(ValueError, match="Unknown model-local CUDA Graph component"): + manager.prepare(_Model((component,))) + + manager = _TestManager(config={"decode": {"unknown_bucket_policy": [2]}}) + with pytest.raises(ValueError, match="Unknown config key"): + manager.prepare(_Model((component,))) + + +@pytest.mark.parametrize( + ("key", "value", "message"), + [ + ("max_extra_graphs", True, "must be a non-negative integer"), + ("max_extra_graphs", -1, "must be a non-negative integer"), + ("max_extra_graphs", 1.0, "must be a non-negative integer"), + ], +) +def test_component_policy_validation_rejects_invalid_types(key, value, message) -> None: + component, _ = _component("decode", 2, capture_mode=ModelLocalCaptureMode.PRECAPTURE_LAZY) + manager = _TestManager(config={"decode": {key: value}}) + + with pytest.raises(TypeError, match=message): + manager.prepare(_Model((component,))) + + +def test_component_registry_rejects_duplicate_ids_and_descriptors() -> None: + first, _ = _component("decode", 2) + duplicate_id, _ = _component("decode", 3) + manager = _TestManager() + with pytest.raises(ValueError, match="Duplicate model-local CUDA Graph component_id"): + manager.prepare(_Model((first, duplicate_id))) + + duplicate_descriptor = ModelLocalCUDAGraphComponent( + "duplicate", + _Routine(), + [ModelLocalCUDAGraphDescriptor(2), ModelLocalCUDAGraphDescriptor(2)], + ) + manager = _TestManager() + with pytest.raises(ValueError, match="Duplicate Descriptor"): + manager.prepare(_Model((duplicate_descriptor,))) + + +def test_omitted_component_remains_on_original_eager_callable() -> None: + component, routine = _component("decode", 2) + manager = _TestManager(config={}) + manager.prepare(_Model((component,))) + manager.capture_and_bind() + + assert not manager.managed_components + assert component._bound_handle is None + component(torch.tensor([1.0])) + assert routine.eager_calls == 1 + + +def test_pure_lazy_profiles_descriptors_but_does_not_capture_at_startup(monkeypatch) -> None: + component, routine = _component("decode", 2, capture_mode=ModelLocalCaptureMode.PURE_LAZY) + manager = _TestManager(config={"decode": {"max_extra_graphs": 1}}) + manager.prepare(_Model((component,))) + manager.capture_and_bind() + + assert manager.capture_attempts[("decode", 2)] == 0 + assert component.available_descriptors == frozenset() + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + monkeypatch.setattr(manager, "_runtime_capture_scope", nullcontext) + value = torch.ones(2) + assert torch.equal(component(value), value * 2) + assert component.available_descriptors == frozenset({ModelLocalCUDAGraphDescriptor(2)}) + assert manager.capture_attempts[("decode", 2)] == 1 + larger = torch.ones(3) + assert torch.equal(component(larger), larger * 2) + assert manager.capture_attempts[("decode", 3)] == 0 + assert routine.eager_calls == 1 + + +def test_pure_lazy_requires_profiling_descriptors() -> None: + component, _ = _component("decode", capture_mode=ModelLocalCaptureMode.PURE_LAZY) + manager = _TestManager(config={"decode": {}}) + with pytest.raises(ValueError, match="needs descriptors for memory profiling"): + manager.prepare(_Model((component,))) + + +def test_descriptor_rejection_falls_back_to_eager_without_capture(monkeypatch) -> None: + component, routine = _component("decode", 2, capture_mode=ModelLocalCaptureMode.PRECAPTURE_LAZY) + manager = _TestManager(config={"decode": {"max_extra_graphs": 1}}) + validation_calls = [] + + def reject(descriptor): + validation_calls.append(descriptor) + return False + + component.validate_descriptor = reject + manager.prepare(_Model((component,))) + manager.capture_and_bind() + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + monkeypatch.setattr(manager, "_runtime_capture_scope", nullcontext) + + value = torch.tensor([1.0, 2.0]) + assert torch.equal(component(value), value * 2) + assert torch.equal(component(value), value * 2) + assert routine.eager_calls == 2 + assert manager.capture_attempts[("decode", 2)] == 3 + assert validation_calls == [ModelLocalCUDAGraphDescriptor(2)] * 3 + + +def test_successful_lazy_miss_registers_descriptor_and_replays_current_call(monkeypatch) -> None: + component, routine = _component("decode", 2, capture_mode=ModelLocalCaptureMode.PRECAPTURE_LAZY) + manager = _TestManager(config={"decode": {}}) + manager.prepare(_Model((component,))) + manager.capture_and_bind() + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + monkeypatch.setattr(manager, "_runtime_capture_scope", lambda: nullcontext()) + + value = torch.tensor([1.0, 2.0, 3.0]) + assert torch.equal(component(value), value * 2) + assert len(manager.managed_components["decode"].entries) == 2 + assert manager.capture_attempts[("decode", 3)] == 1 + assert routine.eager_calls == 0 + + +@pytest.mark.parametrize("mode", [ModelLocalCaptureMode.PRECAPTURE_LAZY, ModelLocalCaptureMode.PURE_LAZY]) +def test_zero_lazy_graph_limit_keeps_all_captured_descriptors(monkeypatch, mode) -> None: + component, routine = _component("decode", 2, capture_mode=mode) + manager = _TestManager(config={"decode": {}}) + manager.prepare(_Model((component,))) + manager.capture_and_bind() + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + monkeypatch.setattr(manager, "_runtime_capture_scope", nullcontext) + + for size in (3, 4, 5): + value = torch.ones(size) + assert torch.equal(component(value), value * 2) + + assert manager.managed_components["decode"].max_graphs is None + expected = {2, 3, 4, 5} if mode is ModelLocalCaptureMode.PRECAPTURE_LAZY else {3, 4, 5} + assert {descriptor.variant for descriptor in component.available_descriptors} == expected + assert routine.eager_calls == 0 + + +def test_lazy_capture_rejects_nested_outer_capture(monkeypatch) -> None: + component, routine = _component("decode", 2, capture_mode=ModelLocalCaptureMode.PRECAPTURE_LAZY) + manager = _TestManager(config={"decode": {"max_extra_graphs": 1}}) + manager.prepare(_Model((component,))) + manager.capture_and_bind() + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: True) + + value = torch.tensor([1.0, 2.0, 3.0]) + assert torch.equal(component(value), value * 2) + assert routine.eager_calls == 1 + assert manager.capture_attempts[("decode", 3)] == 0 + + +def test_lazy_capacity_falls_back_without_evicting_startup_entries(monkeypatch) -> None: + component, routine = _component("decode", 2, 3, capture_mode=ModelLocalCaptureMode.PRECAPTURE_LAZY) + manager = _TestManager(config={"decode": {"max_extra_graphs": 1}}) + manager.prepare(_Model((component,))) + manager.capture_and_bind() + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False) + monkeypatch.setattr(manager, "_runtime_capture_scope", lambda: nullcontext()) + + component(torch.tensor([1.0, 2.0, 3.0, 4.0])) + component(torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])) + entries = manager.managed_components["decode"].entries + assert [descriptor.variant for descriptor in entries] == [3, 2, 4] + assert component.available_descriptors == frozenset( + { + ModelLocalCUDAGraphDescriptor(3), + ModelLocalCUDAGraphDescriptor(2), + ModelLocalCUDAGraphDescriptor(4), + } + ) + assert manager.capture_attempts[("decode", 5)] == 0 + assert routine.eager_calls == 1 + + +def test_binding_failure_propagates() -> None: + first, _ = _component("first", 2) + second, _ = _component("second", 2) + manager = _TestManager() + manager.prepare(_Model((first, second))) + + captured_graphs = [] + original_capture = manager.capture_entry + + def capture(component, descriptor): + entry = original_capture(component, descriptor) + captured_graphs.append(entry.graph) + return entry + + manager.capture_entry = capture + original_bind = second._bind_handle + + def fail_bind(handle) -> None: + original_bind(handle) + raise RuntimeError("bind failed") + + second._bind_handle = fail_bind + with pytest.raises(RuntimeError, match="bind failed"): + manager.capture_and_bind() + + assert first._bound_handle is None + assert second._bound_handle is None + assert manager.managed_components == {} + assert all(graph.reset_called for graph in captured_graphs) + + +def test_startup_capture_failure_releases_previous_entries() -> None: + component, _ = _component("decode", 2, 3) + manager = _TestManager() + manager.prepare(_Model((component,))) + captured_graphs = [] + original_capture = manager.capture_entry + + def capture(component, descriptor): + if descriptor.variant == 2: + raise RuntimeError("second startup capture failed") + entry = original_capture(component, descriptor) + captured_graphs.append(entry.graph) + return entry + + manager.capture_entry = capture + with pytest.raises(RuntimeError, match="second startup capture failed"): + manager.capture_and_bind() + + assert component._bound_handle is None + assert manager.managed_components == {} + assert len(captured_graphs) == 1 + assert captured_graphs[0].reset_called + + +def test_prepare_and_capture_are_single_use_lifecycle_operations() -> None: + component, _ = _component("decode", 2) + manager = _TestManager() + model = _Model((component,)) + manager.prepare(model) + with pytest.raises(RuntimeError, match=r"prepare\(\).*more than once"): + manager.prepare(model) + + manager.capture_and_bind() + with pytest.raises(RuntimeError, match="capture has already been attempted"): + manager.capture_and_bind() + + +def test_descriptor_rejection_isolated_to_sibling_component() -> None: + first, _ = _component("first", 2, 3) + second, _ = _component("second", 4) + manager = _TestManager() + first.validate_descriptor = lambda descriptor: descriptor.variant != 2 + manager.prepare(_Model((first, second))) + manager.capture_and_bind() + + assert list(manager.managed_components["first"].entries) == [ModelLocalCUDAGraphDescriptor(3)] + assert list(manager.managed_components["second"].entries) == [ModelLocalCUDAGraphDescriptor(4)] + + +def test_capture_failure_propagates() -> None: + class _FailingRoutine(_Routine): + def allocate_buffers(self, descriptor, device): + del descriptor, device + raise RuntimeError("capture allocation failed") + + component = ModelLocalCUDAGraphComponent( + "decode", + _FailingRoutine(), + [ModelLocalCUDAGraphDescriptor(2)], + ) + vllm_config = SimpleNamespace( + model_config=SimpleNamespace(model_local_cudagraph={"decode": {}}), + compilation_config=SimpleNamespace(cudagraph_num_of_warmups=0), + ) + manager = ModelLocalCUDAGraphManager(vllm_config=vllm_config, device=torch.device("cpu")) + manager.prepare(_Model((component,))) + + with pytest.raises(RuntimeError, match="capture allocation failed"): + manager.capture_and_bind() + + +def test_capture_and_bind_reports_total_memory_delta(monkeypatch): + component, _ = _component("decode", 2) + manager = _TestManager() + manager.device = torch.device("cuda") + free_memory = iter((100, 70)) + monkeypatch.setattr(torch.accelerator, "synchronize", lambda: None) + monkeypatch.setattr(manager, "_synchronized_free_memory", lambda: next(free_memory)) + + def capture(component, descriptor): + buffers = component.routine.allocate_buffers(descriptor, torch.device("cpu")) + output = component.routine.forward_for_capture(buffers) + return ModelLocalCUDAGraphEntry( + descriptor=descriptor, + graph=cast(torch.cuda.CUDAGraph, _Graph(buffers)), + buffers=buffers, + captured_output=output, + ) + + monkeypatch.setattr(manager, "capture_entry", capture) + manager.prepare(_Model((component,))) + + assert manager.capture_and_bind() == 30 + + +def test_capture_and_bind_orders_descriptors_largest_first(): + component, _ = _component("decode", 2, 8, 4, capture_order_key=lambda descriptor: descriptor.variant) + manager = _TestManager() + manager.prepare(_Model((component,))) + + manager.capture_and_bind() + + assert [variant for component_id, variant in manager.capture_attempts if component_id == "decode"] == [8, 4, 2] + + +def test_profile_memory_uses_first_capture_and_per_graph_increment(monkeypatch): + first, _ = _component( + "first", + 2, + 3, + 4, + capture_order_key=lambda descriptor: descriptor.variant, + capture_mode=ModelLocalCaptureMode.PRECAPTURE_LAZY, + ) + second, _ = _component("second", 5) + manager = _TestManager(config={"first": {"max_extra_graphs": 2}, "second": {}}) + manager.prepare(_Model((first, second))) + free_memory = iter((100, 90, 90, 87, 87, 80)) + captured = [] + capture_pools = [] + captured_variants = [] + + def capture(component, descriptor, *, graph_pool=None): + del component + capture_pools.append(graph_pool) + captured_variants.append(descriptor.variant) + buffers = _Buffers(torch.zeros(1), torch.zeros(1)) + entry = ModelLocalCUDAGraphEntry( + descriptor=descriptor, + graph=cast(torch.cuda.CUDAGraph, _Graph(buffers)), + buffers=buffers, + captured_output=buffers.output, + ) + captured.append(entry) + return entry + + monkeypatch.setattr(manager, "capture_entry", capture) + profiling_pool = object() + monkeypatch.setattr(current_platform, "graph_pool_handle", lambda: profiling_pool) + monkeypatch.setattr(torch.accelerator, "synchronize", lambda: None) + monkeypatch.setattr(manager, "_synchronized_free_memory", lambda: next(free_memory)) + monkeypatch.setattr(torch.accelerator, "empty_cache", lambda: None) + + # first: 10 + max(3, 1 MiB) * (3 - 1 + 2); second: 7. + assert manager.profile_memory() == 10 + (1 << 20) * 4 + 7 + assert len(captured) == 3 + assert capture_pools == [profiling_pool] * 3 + assert captured_variants == [4, 3, 5] + assert all(entry.graph.reset_called for entry in captured) + + +def test_precapture_rejects_lazy_graph_capacity() -> None: + component, _ = _component("decode", 2) + manager = _TestManager(config={"decode": {"max_extra_graphs": 4}}) + with pytest.raises(ValueError, match="requires a lazy capture mode"): + manager.prepare(_Model((component,))) diff --git a/tests/worker_v2/test_omni_generation_model_runner.py b/tests/worker_v2/test_omni_generation_model_runner.py index 63c48029151..d2c9decba7d 100644 --- a/tests/worker_v2/test_omni_generation_model_runner.py +++ b/tests/worker_v2/test_omni_generation_model_runner.py @@ -4,20 +4,156 @@ """OmniGenerationModelRunner contracts: empty-step lifecycle, output partition, CPU-sync vs CUDA-async dispatch, and async-chunk slot recycling.""" +from contextlib import nullcontext from types import SimpleNamespace from unittest.mock import MagicMock import numpy as np import pytest import torch +from vllm.config import CUDAGraphMode from vllm_omni.model_executor.models.output_templates import OmniOutput from vllm_omni.outputs import OmniModelRunnerOutput +from vllm_omni.worker.base import OmniGPUWorkerBase +from vllm_omni.worker.gpu_generation_worker import GPUGenerationWorker from vllm_omni.worker_v2.omni_generation_model_runner import OmniGenerationModelRunner +from vllm_omni.worker_v2.omni_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 () + + +def _graph_runner( + *, + enabled: bool = True, + enforce_eager: bool = False, + mode: CUDAGraphMode = CUDAGraphMode.FULL, +) -> OmniGenerationModelRunner: + runner = object.__new__(OmniGenerationModelRunner) + runner.device = torch.device("cpu") + runner.model = _GraphModel() + runner.model_config = SimpleNamespace( + enforce_eager=enforce_eager, + model_local_cudagraph={"decode": {}} if enabled else None, + ) + runner.compilation_config = SimpleNamespace(cudagraph_mode=mode) + runner.vllm_config = SimpleNamespace() + runner.model_local_cudagraph_manager = None + return runner + + +def test_model_local_graph_load_prepares_manager_only_when_configured(monkeypatch): + from vllm_omni.worker_v2 import omni_generation_model_runner as generation_runner + + class FakeManager: + def __init__(self, *, vllm_config, device): + assert vllm_config is runner.vllm_config + assert device == runner.device + self.prepared_model = None + + def prepare(self, model): + self.prepared_model = model + + monkeypatch.setattr(OmniGPUModelRunner, "load_model", lambda self, *args, **kwargs: None) + monkeypatch.setattr(generation_runner, "ModelLocalCUDAGraphManager", FakeManager) + runner = _graph_runner() + runner.load_model() + assert runner.model_local_cudagraph_manager.prepared_model is runner.model + + +@pytest.mark.parametrize( + "kwargs", + [ + {"enabled": False}, + {"enforce_eager": True}, + {"mode": CUDAGraphMode.NONE}, + ], +) +def test_model_local_graph_disabled_keeps_upstream_lifecycle(monkeypatch, kwargs): + monkeypatch.setattr(OmniGPUModelRunner, "load_model", lambda self, *args, **kwargs: None) + monkeypatch.setattr(OmniGPUModelRunner, "needs_cudagraph_capture", lambda self: False) + monkeypatch.setattr(OmniGPUModelRunner, "profile_cudagraph_memory", lambda self: 7) + monkeypatch.setattr(OmniGPUModelRunner, "capture_model", lambda self: 11) + runner = _graph_runner(**kwargs) + runner.load_model() + assert runner.model_local_cudagraph_manager is None + assert not runner.needs_cudagraph_capture() + assert runner.profile_cudagraph_memory() == 7 + assert runner.capture_model() == 11 + + +def test_model_local_graph_profile_capture_and_shutdown_use_local_manager(monkeypatch): + from vllm_omni.worker_v2 import omni_generation_model_runner as generation_runner + + runner = _graph_runner() + events = [] + manager = SimpleNamespace( + profile_memory=lambda: events.append("profile") or 123, + capture_and_bind=lambda: events.append("capture"), + clear=lambda: events.append("clear"), + ) + runner.model_local_cudagraph_manager = manager + monkeypatch.setattr(generation_runner, "freeze_gc_for_cudagraph_capture", nullcontext) + monkeypatch.setattr(generation_runner, "graph_capture", lambda **kwargs: nullcontext()) + monkeypatch.setattr(generation_runner, "lock_workspace", lambda: events.append("lock_workspace")) + monkeypatch.setattr(generation_runner, "set_cudagraph_capturing_enabled", lambda value: events.append(value)) + monkeypatch.setattr(torch.accelerator, "synchronize", lambda: None) + monkeypatch.setattr(torch.accelerator, "empty_cache", lambda: None) + memory_info = iter(((1000, 0), (900, 0))) + monkeypatch.setattr(torch.accelerator, "get_memory_info", lambda: next(memory_info)) + monkeypatch.setattr(OmniGPUModelRunner, "shutdown", lambda self: events.append("shutdown")) + + assert runner.needs_cudagraph_capture() + assert runner.profile_cudagraph_memory() == 123 + assert runner.capture_model() == 100 + runner.shutdown() + assert runner.model_local_cudagraph_manager is None + assert events == [True, "profile", False, True, "capture", False, "lock_workspace", "clear", "shutdown"] + + +@pytest.mark.parametrize("has_manager", [True, False]) +def test_v2_generation_worker_captures_only_with_model_local_manager(has_manager): + worker = object.__new__(GPUGenerationWorker) + worker.use_v2_model_runner = True + worker.model_config = SimpleNamespace(enforce_eager=False) + events = [] + worker.model_runner = SimpleNamespace( + model_local_cudagraph_manager=object() if has_manager else None, + profile_run=lambda: events.append("profile"), + capture_model=lambda: events.append("capture"), + ) + worker._get_cudagraph_capture_context = nullcontext + + worker.compile_or_warm_up_model() + + assert events == (["profile", "capture"] if has_manager else ["profile"]) + + +@pytest.mark.parametrize("has_manager", [True, False]) +def test_generation_worker_reserves_model_local_graph_profile_memory(monkeypatch, has_manager): + monkeypatch.setattr(OmniGPUWorkerBase, "determine_available_memory", lambda self: 500) + worker = object.__new__(GPUGenerationWorker) + worker.cache_config = SimpleNamespace(kv_cache_memory_bytes=None) + profile_calls = [] + worker.model_runner = SimpleNamespace( + model_local_cudagraph_manager=object() if has_manager else None, + profile_cudagraph_memory=lambda: profile_calls.append(True) or 123, + ) + + assert worker.determine_available_memory() == (377 if has_manager else 500) + assert profile_calls == ([True] if has_manager else []) + if has_manager: + assert worker.cudagraph_memory_estimate == 123 + assert worker.available_kv_cache_memory_bytes == 377 + + class _FakeStagedField: def __init__(self, data): self.np = data diff --git a/vllm_omni/config/model.py b/vllm_omni/config/model.py index a56e6148cbc..167b91c8fa1 100644 --- a/vllm_omni/config/model.py +++ b/vllm_omni/config/model.py @@ -126,6 +126,8 @@ class OmniModelConfig(ModelConfig): stage_id: int = 0 async_chunk: bool = False session_mode: str = "turn" + # Resolved per-stage runner-owned model-local CUDA Graph configuration. + model_local_cudagraph: dict[str, Any] | None = None retains_state_across_chunks: bool = False use_v2_model_runner: bool = False supports_native_mrv2_data_plane: bool = False diff --git a/vllm_omni/config/omni_config.py b/vllm_omni/config/omni_config.py index 3a650dd03d2..d7e72f6ea4b 100644 --- a/vllm_omni/config/omni_config.py +++ b/vllm_omni/config/omni_config.py @@ -188,6 +188,7 @@ class _ModelEngineOverrides(TypedDict, total=False): enable_broadcast_weight_load: bool num_weight_load_threads: int disable_autocast: bool + model_local_cudagraph: dict[str, Any] # Upstream ModelConfig inputs that users pass as global CLI flags. served_model_name: str | list[str] allowed_local_media_path: str @@ -523,6 +524,9 @@ class OmniStageModelConfig(_TrackExplicitConfigFields): model_subdir: str | None = None tokenizer_subdir: str | None = None requires_full_payload_input: bool = False + # User-facing per-stage runner-owned model-local graph configuration. It is + # projected to OmniModelConfig.model_local_cudagraph by OmniEngineArgs. + model_local_cudagraph: dict[str, Any] | None = None # Upstream ModelConfig inputs that users pass as global CLI flags. served_model_name: str | list[str] | None = None allowed_local_media_path: str | None = None @@ -1939,7 +1943,7 @@ def _build_stage_config( builder = _STAGE_CONFIG_BUILDERS[topology.execution_type] except KeyError as exc: raise ValueError(f"Unsupported stage execution type: {topology.execution_type!r}") from exc - return cast( + stage_config = cast( StageConfigType, builder( pipeline, @@ -1950,6 +1954,15 @@ def _build_stage_config( model=model, ), ) + if ( + stage_config.model_config.model_local_cudagraph is not None + and topology.execution_type != StageExecutionType.LLM_GENERATION + ): + raise ValueError( + "model_local_cudagraph is supported only for LLM_GENERATION stages; " + f"stage {topology.stage_id} uses {topology.execution_type.value}" + ) + return stage_config def _build_quantization_config( diff --git a/vllm_omni/engine/arg_utils.py b/vllm_omni/engine/arg_utils.py index b9ed35de04c..1737fd1a122 100644 --- a/vllm_omni/engine/arg_utils.py +++ b/vllm_omni/engine/arg_utils.py @@ -197,6 +197,8 @@ class OmniEngineArgs(EngineArgs): silence_ban_frames: int = 0 async_chunk: bool = False session_mode: str = "turn" + # Public deploy YAML and the resolved model config use the same key. + model_local_cudagraph: dict[str, Any] | None = None retains_state_across_chunks: bool = False use_v2_model_runner: bool = False supports_native_mrv2_data_plane: bool = False @@ -314,6 +316,12 @@ def create_model_config(self) -> OmniModelConfig: Returns: OmniModelConfig instance with all configuration fields set """ + if self.model_local_cudagraph is not None and self.worker_type != "generation": + raise ValueError( + "model_local_cudagraph is supported only for LLM_GENERATION stages; " + f"got worker_type={self.worker_type!r}" + ) + # register omni models to avoid model not found error self._ensure_omni_models_registered() @@ -429,6 +437,7 @@ def create_model_config(self) -> OmniModelConfig: stage_id=self.stage_id, async_chunk=self.async_chunk, session_mode=self.session_mode, + model_local_cudagraph=self.model_local_cudagraph, retains_state_across_chunks=self.retains_state_across_chunks, use_v2_model_runner=self.use_v2_model_runner, supports_native_mrv2_data_plane=self.supports_native_mrv2_data_plane, diff --git a/vllm_omni/model_executor/models/interfaces/__init__.py b/vllm_omni/model_executor/models/interfaces/__init__.py new file mode 100644 index 00000000000..4ff70f83410 --- /dev/null +++ b/vllm_omni/model_executor/models/interfaces/__init__.py @@ -0,0 +1,28 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project + +from .model_local_cudagraph import ( + BaseModelLocalCUDAGraphRoutine, + ModelLocalCaptureMode, + ModelLocalCUDAGraphComponent, + ModelLocalCUDAGraphDescriptor, + ModelLocalCUDAGraphRoutine, + ModelLocalGraphHandle, + ModelLocalRuntimeKey, + ModelLocalRuntimeResolution, + SupportsModelLocalCUDAGraph, + supports_model_local_cudagraph, +) + +__all__ = [ + "BaseModelLocalCUDAGraphRoutine", + "ModelLocalCaptureMode", + "SupportsModelLocalCUDAGraph", + "ModelLocalCUDAGraphDescriptor", + "ModelLocalCUDAGraphRoutine", + "ModelLocalCUDAGraphComponent", + "ModelLocalGraphHandle", + "ModelLocalRuntimeKey", + "ModelLocalRuntimeResolution", + "supports_model_local_cudagraph", +] diff --git a/vllm_omni/model_executor/models/interfaces/model_local_cudagraph.py b/vllm_omni/model_executor/models/interfaces/model_local_cudagraph.py new file mode 100644 index 00000000000..eaf79350a64 --- /dev/null +++ b/vllm_omni/model_executor/models/interfaces/model_local_cudagraph.py @@ -0,0 +1,318 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project + +"""Model-facing contracts for runner-owned model-local CUDA Graphs. + +This module intentionally contains no worker lifecycle implementation. Model +packages depend on these declarations; the worker-side manager consumes them. +""" + +from __future__ import annotations + +from collections.abc import Callable, Generator, Hashable, Sequence, Set +from contextlib import AbstractContextManager, contextmanager +from dataclasses import dataclass +from enum import Enum +from typing import Any, ClassVar, Generic, Literal, Protocol, TypeVar, runtime_checkable + +import torch + +VariantT = TypeVar("VariantT", bound=Hashable) + + +class ModelLocalCaptureMode(str, Enum): + """Model-owned capture timing for one Component; runtime graphs are never evicted.""" + + PRECAPTURE = "precapture" + PRECAPTURE_LAZY = "precapture_lazy" + PURE_LAZY = "pure_lazy" + + @property + def allows_lazy_capture(self) -> bool: + return self is not ModelLocalCaptureMode.PRECAPTURE + + +@dataclass(frozen=True) +class ModelLocalCUDAGraphDescriptor(Generic[VariantT]): + """One immutable, Component-scoped capture specialization.""" + + variant: VariantT + + +@dataclass(frozen=True) +class ModelLocalRuntimeKey(Generic[VariantT]): + """Actual runtime input identity before capture-bucket selection.""" + + variant: VariantT + + +@dataclass(frozen=True) +class ModelLocalRuntimeResolution: + """A valid runtime invocation mapped to an available Descriptor, if any.""" + + runtime_key: ModelLocalRuntimeKey + descriptor: ModelLocalCUDAGraphDescriptor | None + + +class ModelLocalGraphHandle: + """Opaque runtime endpoint for one bound model-local CUDA Graph Component. + + The Handle type is shared by the model-facing Component and worker-side + Manager, but Handle instances are manager-created and manager-owned. It + deliberately hides graph entries, replay buffers, and dispatch policy + from model code. It exposes only the descriptors currently backed by + active graphs so model routing can align grouping with graph coverage. + """ + + __slots__ = ("_call", "_get_available_descriptors") + + def __init__( + self, + call: Callable[..., Any], + get_available_descriptors: Callable[[], frozenset[ModelLocalCUDAGraphDescriptor]] = frozenset, + ) -> None: + self._call = call + self._get_available_descriptors = get_available_descriptors + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + return self._call(*args, **kwargs) + + @property + def available_descriptors(self) -> frozenset[ModelLocalCUDAGraphDescriptor]: + """A read-only snapshot of descriptors currently backed by graphs.""" + + return self._get_available_descriptors() + + +class ModelLocalCUDAGraphRoutine(Protocol): + """Model-specific adapter for graph-shape resolution and static-buffer handling. + + A Routine owns graph-shape resolution and static-buffer adaptation. Request/model + semantic state transitions are generally kept in the model execution path to + preserve ownership boundaries. + + Runtime ``args`` and ``kwargs`` provide context for descriptor selection, + replay-input preparation, and logical output materialization. + """ + + def eager_call(self, *args: Any, **kwargs: Any) -> Any: ... + + def validate_runtime_inputs( + self, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> None: ... + + def resolve_runtime( + self, + args: tuple[Any, ...], + kwargs: dict[str, Any], + available: Set[ModelLocalCUDAGraphDescriptor], + ) -> ModelLocalRuntimeResolution: ... + + def make_lazy_descriptor( + self, + runtime_key: ModelLocalRuntimeKey, + ) -> ModelLocalCUDAGraphDescriptor | None: ... + + def allocate_buffers( + self, + descriptor: ModelLocalCUDAGraphDescriptor, + device: torch.device, + ) -> object: + """Allocate descriptor-owned static buffers used by capture and replay. + + The returned object is owned by the graph entry. Request-local semantic + state should generally remain model-owned rather than being transferred + into the graph-buffer lifecycle. + """ + ... + + def prepare_for_capture(self, buffers: object) -> None: + """Prepare graph-owned buffers for warmup/capture. + + This hook is intended for initializing or normalizing static graph state + required by the captured callable. Runtime invocation-specific semantics + are preferably kept outside the capture lifecycle. + """ + ... + + def capture_context( + self, descriptor: ModelLocalCUDAGraphDescriptor, buffers: object + ) -> AbstractContextManager[None]: ... + + def forward_for_capture(self, buffers: object) -> object: ... + + def after_capture(self, buffers: object) -> None: + """Clean up model-specific state after the capture region exits, even on failure.""" + ... + + def copy_runtime_inputs( + self, + args: tuple[Any, ...], + kwargs: dict[str, Any], + buffers: object, + ) -> None: + """Adapt one runtime invocation into descriptor-owned replay buffers. + + ``args`` and ``kwargs`` provide the runtime context needed to populate or + normalize static graph buffers. Request/model semantic state transitions + are preferably kept in the model execution path rather than coupled to + replay-buffer preparation. + """ + ... + + def output_after_replay( + self, + args: tuple[Any, ...], + kwargs: dict[str, Any], + buffers: object, + captured_output: object, + ) -> object: + """Materialize a runtime-visible result from captured graph output. + + ``args`` and ``kwargs`` provide runtime context such as the logical batch + size or output extent. Implementations may use them to select, slice, or + clone graph-owned outputs. Request/model semantic state commit is generally + better kept in the model execution path to preserve ownership boundaries. + """ + ... + + +class BaseModelLocalCUDAGraphRoutine: + """Model-specific adapter between runtime calls and static CUDA Graph buffers. + + A Routine owns graph-shape resolution and static-buffer adaptation. + Request/model semantic state transitions are generally kept in the model + execution path to preserve ownership boundaries. Runtime ``args`` and + ``kwargs`` provide context for descriptor selection, replay-input + preparation, and logical output materialization. + """ + + def validate_runtime_inputs( + self, + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> None: + del args, kwargs + + def make_lazy_descriptor( + self, + runtime_key: ModelLocalRuntimeKey, + ) -> ModelLocalCUDAGraphDescriptor | None: + return ModelLocalCUDAGraphDescriptor(runtime_key.variant) + + def prepare_for_capture(self, buffers: object) -> None: + del buffers + + def after_capture(self, buffers: object) -> None: + del buffers + + @contextmanager + def capture_context( + self, descriptor: ModelLocalCUDAGraphDescriptor, buffers: object + ) -> Generator[None, None, None]: + del descriptor + self.prepare_for_capture(buffers) + try: + yield + finally: + self.after_capture(buffers) + + +class ModelLocalCUDAGraphComponent: + """Resolved planning declaration and stable model-owned call site. + + Before Manager binding, the delegate is ``Routine.eager_call``. After + binding, it is the Manager-created ``ModelLocalGraphHandle``. Restoring or + clearing the Component returns the delegate to ``Routine.eager_call``. + """ + + def __init__( + self, + component_id: str, + routine: ModelLocalCUDAGraphRoutine, + descriptors: Sequence[ModelLocalCUDAGraphDescriptor], + clone_output: bool = True, + *, + supported_config_keys: Set[str] = frozenset(), + capture_order_key: Callable[[ModelLocalCUDAGraphDescriptor], Any] | None = None, + capture_mode: ModelLocalCaptureMode = ModelLocalCaptureMode.PRECAPTURE, + ) -> None: + self.component_id = component_id + self.routine = routine + # PURE_LAZY uses these only to profile the largest supported capture + # shapes; the other modes also use them for startup capture. + self.descriptors = tuple(descriptors) + self.clone_output = bool(clone_output) + self.supported_config_keys = frozenset(supported_config_keys) + self._capture_order_key = capture_order_key + self.capture_mode = ModelLocalCaptureMode(capture_mode) + self._delegate: Callable[..., Any] = routine.eager_call + # The Component owns the stable call site, not the runtime Handle + # lifecycle. The Manager constructs and binds the Handle after capture. + self._bound_handle: ModelLocalGraphHandle | None = None + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + return self._delegate(*args, **kwargs) + + def validate_descriptor(self, descriptor: ModelLocalCUDAGraphDescriptor) -> bool: + """Validate a descriptor against the Component's supported variants.""" + del descriptor + return True + + @property + def capture_descriptors(self) -> tuple[ModelLocalCUDAGraphDescriptor, ...]: + """Startup descriptors ordered to establish the largest graph first. + + By default the Descriptor variant is the ordering key. Components whose + variants do not directly express graph size may provide + ``capture_order_key`` instead. + """ + + key = self._capture_order_key or (lambda descriptor: descriptor.variant) + try: + return tuple(sorted(self.descriptors, key=key, reverse=True)) + except TypeError as exc: + raise TypeError( + f"Component {self.component_id} needs capture_order_key for non-orderable Descriptor variants" + ) from exc + + @property + def is_graph_bound(self) -> bool: + """Whether Manager has replaced the eager delegate with a Handle.""" + + return self._bound_handle is not None + + @property + def available_descriptors(self) -> frozenset[ModelLocalCUDAGraphDescriptor]: + """Descriptors currently backed by active non-eager graph entries.""" + + if self._bound_handle is None: + return frozenset() + return self._bound_handle.available_descriptors + + def _bind_handle(self, handle: ModelLocalGraphHandle) -> None: + if not isinstance(handle, ModelLocalGraphHandle): + raise TypeError("ModelLocalCUDAGraphComponent requires a ModelLocalGraphHandle") + if self._bound_handle is not None: + raise RuntimeError(f"Component already bound: {self.component_id}") + self._bound_handle = handle + self._delegate = handle + + def _restore_eager(self) -> None: + self._delegate = self.routine.eager_call + self._bound_handle = None + + +@runtime_checkable +class SupportsModelLocalCUDAGraph(Protocol): + supports_model_local_cudagraph: ClassVar[Literal[True]] + + def get_model_local_cudagraph_components(self) -> Sequence[ModelLocalCUDAGraphComponent]: ... + + +def supports_model_local_cudagraph(model: object) -> bool: + return bool(getattr(model, "supports_model_local_cudagraph", False)) and callable( + getattr(model, "get_model_local_cudagraph_components", None) + ) diff --git a/vllm_omni/worker/gpu_generation_model_runner.py b/vllm_omni/worker/gpu_generation_model_runner.py index c82cf71e2cc..c17539a3931 100644 --- a/vllm_omni/worker/gpu_generation_model_runner.py +++ b/vllm_omni/worker/gpu_generation_model_runner.py @@ -13,13 +13,15 @@ import logging from collections.abc import Mapping from dataclasses import replace +from typing import cast import numpy as np import torch +from vllm.compilation.monitor import set_cudagraph_capturing_enabled from vllm.config import CUDAGraphMode from vllm.distributed.ec_transfer import get_ec_transfer, has_ec_transfer from vllm.distributed.kv_transfer import get_kv_transfer_group, has_kv_transfer_group -from vllm.distributed.parallel_state import get_pp_group +from vllm.distributed.parallel_state import get_pp_group, graph_capture from vllm.forward_context import set_forward_context from vllm.sequence import IntermediateTensors from vllm.utils.math_utils import cdiv @@ -37,12 +39,18 @@ ) from vllm.v1.worker.ubatch_utils import maybe_create_ubatch_slices from vllm.v1.worker.utils import sanity_check_mm_encoder_outputs +from vllm.v1.worker.workspace import lock_workspace +from vllm_omni.model_executor.models.interfaces.model_local_cudagraph import ( + SupportsModelLocalCUDAGraph, + supports_model_local_cudagraph, +) from vllm_omni.outputs import OmniModelRunnerOutput from vllm_omni.utils.mm_outputs import partition_payload_list from vllm_omni.worker.gpu_ar_model_runner import ExecuteModelState, _ensure_tensor_values from vllm_omni.worker.gpu_model_runner import OmniGPUModelRunner from vllm_omni.worker.mixins import maybe_unpad_input_ids +from vllm_omni.worker.model_local_cudagraph_manager import ModelLocalCUDAGraphManager from vllm_omni.worker.omni_connector_model_runner_mixin import ( OmniConnectorModelRunnerMixin, needs_omni_connector, @@ -61,12 +69,83 @@ class GPUGenerationModelRunner(OmniGPUModelRunner, OmniConnectorModelRunnerMixin def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) + self.model_local_cudagraph_manager: ModelLocalCUDAGraphManager | None = None self._async_chunk = getattr(self.model_config, "async_chunk", False) if needs_omni_connector(self.model_config): self.init_omni_connectors( model_config=self.model_config, ) + def _model_local_cudagraph_enabled(self) -> bool: + return ( + not self.model_config.enforce_eager + and self.compilation_config.cudagraph_mode != CUDAGraphMode.NONE + and bool(getattr(self.model_config, "model_local_cudagraph", None)) + and supports_model_local_cudagraph(self.get_model()) + ) + + def load_model(self, *args, **kwargs) -> None: + super().load_model(*args, **kwargs) + if not self._model_local_cudagraph_enabled(): + return + + raw_model = self.get_model() + manager = ModelLocalCUDAGraphManager( + vllm_config=self.vllm_config, + device=self.device, + ) + manager.prepare(cast(SupportsModelLocalCUDAGraph, raw_model)) + self.model_local_cudagraph_manager = manager + logger.info("Initialized runner-owned model-local CUDA Graph manager") + + @torch.inference_mode() + def profile_cudagraph_memory(self) -> int: + manager = self.model_local_cudagraph_manager + if manager is not None: + set_cudagraph_capturing_enabled(True) + try: + with self._freeze_gc(), graph_capture(device=self.device): + torch.accelerator.synchronize() + torch.accelerator.empty_cache() + return manager.profile_memory() + finally: + set_cudagraph_capturing_enabled(False) + return super().profile_cudagraph_memory() + + @torch.inference_mode() + def capture_model(self) -> int: + manager = self.model_local_cudagraph_manager + if manager is None: + return super().capture_model() + + set_cudagraph_capturing_enabled(True) + try: + with self._freeze_gc(), graph_capture(device=self.device): + torch.accelerator.synchronize() + torch.accelerator.empty_cache() + free_before = torch.accelerator.get_memory_info()[0] + manager.capture_and_bind() + torch.accelerator.synchronize() + free_after = torch.accelerator.get_memory_info()[0] + finally: + set_cudagraph_capturing_enabled(False) + + torch.accelerator.synchronize() + torch.accelerator.empty_cache() + lock_workspace() + captured_bytes = max(0, free_before - free_after) + logger.info( + "Runner-owned model-local CUDA Graph capture replaced upstream model capture (%.2f MiB)", + captured_bytes / (1 << 20), + ) + return captured_bytes + + def shutdown(self) -> None: + if self.model_local_cudagraph_manager is not None: + self.model_local_cudagraph_manager.clear() + self.model_local_cudagraph_manager = None + super().shutdown() + def _update_request_states(self, scheduler_output: SchedulerOutput): # remove requests for req_id in scheduler_output.finished_req_ids: @@ -223,6 +302,9 @@ def execute_model( max_num_scheduled_tokens=max_num_scheduled_tokens, use_cascade_attn=cascade_attn_prefix_lens is not None, num_encoder_reqs=len(scheduler_output.scheduled_encoder_inputs), + # The raw model's declared Components own graph replay for this + # stage. Keep the upstream root wrapper on its eager runnable. + force_eager=self.model_local_cudagraph_manager is not None, ) logger.debug( @@ -603,6 +685,11 @@ def _dummy_run( of max_query_len. Used to profile attention workspace that scales with context length. """ + if self.model_local_cudagraph_manager is not None: + # Warmup/profile calls must not accidentally trigger the upstream + # root CUDAGraphWrapper once model-local Components own graph capture. + cudagraph_runtime_mode = CUDAGraphMode.NONE + mm_config = self.vllm_config.model_config.multimodal_config if mm_config and mm_config.mm_encoder_only: # The current dummy run only covers LM execution, so we can skip it. diff --git a/vllm_omni/worker/gpu_generation_worker.py b/vllm_omni/worker/gpu_generation_worker.py index 0e939957e7d..4df186a3c7d 100644 --- a/vllm_omni/worker/gpu_generation_worker.py +++ b/vllm_omni/worker/gpu_generation_worker.py @@ -35,6 +35,19 @@ class GPUGenerationWorker(OmniWorkerMixin, OmniGPUWorkerBase): model_runner_cls = GPUGenerationModelRunner + def determine_available_memory(self) -> int: + available = super().determine_available_memory() + if ( + self.cache_config.kv_cache_memory_bytes + or getattr(self.model_runner, "model_local_cudagraph_manager", None) is None + ): + return available + + estimate = self.model_runner.profile_cudagraph_memory() + self.cudagraph_memory_estimate = estimate + self.available_kv_cache_memory_bytes = max(0, available - estimate) + return self.available_kv_cache_memory_bytes + @instrument(span_name="Init device") def init_device(self): if self.device_config.device_type in ("cuda", "musa"): @@ -126,4 +139,10 @@ def compile_or_warm_up_model(self) -> CompilationTimes: start = time.perf_counter() self.model_runner.profile_run() + if ( + not self.model_config.enforce_eager + and getattr(self.model_runner, "model_local_cudagraph_manager", None) is not None + ): + with self._get_cudagraph_capture_context(): + self.model_runner.capture_model() return CompilationTimes(language_model=time.perf_counter() - start, encoder=0.0) diff --git a/vllm_omni/worker/model_local_cudagraph_manager.py b/vllm_omni/worker/model_local_cudagraph_manager.py new file mode 100644 index 00000000000..3a6679a0674 --- /dev/null +++ b/vllm_omni/worker/model_local_cudagraph_manager.py @@ -0,0 +1,595 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Omni project + +"""Runner-owned lifecycle for model-declared model-local CUDA Graph Components.""" + +from __future__ import annotations + +import logging +import time +from collections import Counter, OrderedDict +from collections.abc import Callable, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from types import MappingProxyType +from typing import Any, cast + +import torch +from tqdm import tqdm +from vllm.compilation.monitor import set_cudagraph_capturing_enabled +from vllm.config import VllmConfig +from vllm.platforms import current_platform + +from vllm_omni.model_executor.models.interfaces.model_local_cudagraph import ( + ModelLocalCaptureMode, + ModelLocalCUDAGraphComponent, + ModelLocalCUDAGraphDescriptor, + ModelLocalGraphHandle, + ModelLocalRuntimeResolution, + SupportsModelLocalCUDAGraph, +) + +logger = logging.getLogger(__name__) + +MAX_LOG_ITEMS = 100 +LOG_EVERY_CALLS = 100 + + +@dataclass +class ModelLocalCUDAGraphEntry: + """Worker-owned resources for one captured Component Descriptor.""" + + descriptor: ModelLocalCUDAGraphDescriptor + graph: torch.cuda.CUDAGraph + buffers: object + captured_output: object + + +_COMPONENT_POLICY_KEYS = frozenset({"max_extra_graphs"}) + + +def clone_tensor_tree(value: object) -> object: + """Clone tensor leaves after model-specific graph output processing.""" + + if isinstance(value, torch.Tensor): + return value.clone() + if isinstance(value, tuple): + cloned = [clone_tensor_tree(item) for item in value] + if hasattr(value, "_fields"): + constructor = type(value) + return constructor(*cast(tuple[Any, ...], tuple(cloned))) + return tuple(cloned) + if isinstance(value, list): + return [clone_tensor_tree(item) for item in value] + if isinstance(value, dict): + return {key: clone_tensor_tree(item) for key, item in value.items()} + return value + + +@dataclass(frozen=True) +class ModelLocalComponentRuntimeConfig: + max_extra_graphs: int = 0 + + +@dataclass +class ManagedComponent: + component: ModelLocalCUDAGraphComponent + entries: OrderedDict[ModelLocalCUDAGraphDescriptor, ModelLocalCUDAGraphEntry] + capture_mode: ModelLocalCaptureMode + max_graphs: int | None = None + + +class _NoOpRecorder: + __slots__ = () + + def record_graph_hit(self, resolution: ModelLocalRuntimeResolution) -> None: + del resolution + + def record_fallback(self, resolution: ModelLocalRuntimeResolution) -> None: + del resolution + + def record_replay_error(self, resolution: ModelLocalRuntimeResolution) -> None: + del resolution + + +class _ComponentRecorder: + __slots__ = ("_sink", "_component_id") + + def __init__(self, sink: ModelLocalGraphStatsSink, component_id: str) -> None: + self._sink = sink + self._component_id = component_id + + def record_graph_hit(self, resolution: ModelLocalRuntimeResolution) -> None: + self._sink.record("hit", self._component_id, resolution) + + def record_fallback(self, resolution: ModelLocalRuntimeResolution) -> None: + self._sink.record("fallback", self._component_id, resolution) + + def record_replay_error(self, resolution: ModelLocalRuntimeResolution) -> None: + self._sink.record("replay_error", self._component_id, resolution) + + +class ModelLocalGraphStatsSink: + """Manager-read, Recorder-write runtime counters.""" + + def __init__(self, *, enabled: bool, max_log_items: int = MAX_LOG_ITEMS) -> None: + if max_log_items <= 0: + raise ValueError("max_log_items must be positive") + self.enabled = enabled + self.max_log_items = max_log_items + self._calls: Counter[str] = Counter() + self._outcomes: Counter[tuple[str, str]] = Counter() + # Counts graph specializations actually selected for replay. + self._descriptors: OrderedDict[tuple[str, object], int] = OrderedDict() + # Counts runtime keys observed before Descriptor bucket selection. + self._runtime_keys: OrderedDict[tuple[str, object], int] = OrderedDict() + + def recorder_for(self, component_id: str) -> _ComponentRecorder | _NoOpRecorder: + if not self.enabled: + return _NoOpRecorder() + return _ComponentRecorder(self, component_id) + + def record( + self, + outcome: str, + component_id: str, + resolution: ModelLocalRuntimeResolution, + ) -> None: + self._calls[component_id] += 1 + self._outcomes[(component_id, outcome)] += 1 + if resolution.descriptor is not None: + self._record_detail(self._descriptors, (component_id, resolution.descriptor.variant)) + key = (component_id, resolution.runtime_key.variant) + self._record_detail(self._runtime_keys, key) + total_calls = sum(self._calls.values()) + if total_calls % LOG_EVERY_CALLS == 0: + logger.info("Model-local CUDA Graph runtime stats after %d calls: %s", total_calls, self.snapshot()) + + def _record_detail(self, items: OrderedDict[tuple[str, object], int], key: tuple[str, object]) -> None: + if key not in items and len(items) == self.max_log_items: + items.popitem(last=False) + items[key] = items.get(key, 0) + 1 + items.move_to_end(key) + + def snapshot(self) -> dict[str, object]: + return { + "calls": dict(self._calls), + "outcomes": dict(self._outcomes), + "descriptors": dict(self._descriptors), + "runtime_keys": dict(self._runtime_keys), + } + + +class ModelLocalCUDAGraphManager: + """Consumes resolved Components and owns capture/bind/restore lifecycle.""" + + def __init__(self, *, vllm_config: VllmConfig, device: torch.device) -> None: + self.vllm_config = vllm_config + self.device = device + self.components: tuple[ModelLocalCUDAGraphComponent, ...] = () + self.managed_components: dict[str, ManagedComponent] = {} + self._component_configs: dict[str, ModelLocalComponentRuntimeConfig] = {} + self._runtime_capture_stream: torch.cuda.Stream | None = None + self._prepared = False + self._capture_attempted = False + + raw_config = getattr(vllm_config.model_config, "model_local_cudagraph", None) + if raw_config is None: + raw_config = {} + if not isinstance(raw_config, Mapping): + raise TypeError("model_local_cudagraph must be a mapping") + self.config = dict(raw_config) + log_stats = bool(getattr(getattr(vllm_config, "observability_config", None), "cudagraph_metrics", False)) + self.stats_sink = ModelLocalGraphStatsSink(enabled=log_stats) + + def _synchronized_free_memory(self) -> int: + if self.device.type == "cpu": + return 0 + torch.accelerator.synchronize() + return int(torch.accelerator.get_memory_info()[0]) + + def _validate_nonnegative(self, value: object) -> int: + if isinstance(value, bool) or not isinstance(value, int) or value < 0: + raise TypeError(f"{value!r} must be a non-negative integer") + return value + + def prepare(self, model: SupportsModelLocalCUDAGraph) -> None: + if self._prepared: + raise RuntimeError("ModelLocalCUDAGraphManager.prepare() called more than once") + + components = tuple(model.get_model_local_cudagraph_components()) + component_by_id: dict[str, ModelLocalCUDAGraphComponent] = {} + for component in components: + if not isinstance(component, ModelLocalCUDAGraphComponent): + raise TypeError( + "get_model_local_cudagraph_components() must return ModelLocalCUDAGraphComponent objects" + ) + if not component.component_id: + raise ValueError("Model-local CUDA Graph component_id must not be empty") + if component.component_id in component_by_id: + raise ValueError(f"Duplicate model-local CUDA Graph component_id: {component.component_id}") + try: + if len(set(component.descriptors)) != len(component.descriptors): + raise ValueError(f"Duplicate Descriptor in Component {component.component_id}") + except TypeError as exc: + raise TypeError(f"Descriptors for Component {component.component_id} must be hashable") from exc + component_by_id[component.component_id] = component + if not component.descriptors: + logger.info( + "Model-local CUDA Graph Component %s is known but has no startup " + "Descriptors for the resolved model configuration", + component.component_id, + ) + + unknown_component_ids = set(self.config) - set(component_by_id) + if unknown_component_ids: + names = ", ".join(sorted(str(name) for name in unknown_component_ids)) + raise ValueError(f"Unknown model-local CUDA Graph component override(s): {names}") + + for component_id, raw_component in self.config.items(): + component = component_by_id[component_id] + if raw_component is None: + raw_component = {} + if not isinstance(raw_component, Mapping): + raise TypeError(f"model_local_cudagraph.{component_id} must be a mapping") + unknown_keys = set(raw_component) - _COMPONENT_POLICY_KEYS - set(component.supported_config_keys) + if unknown_keys: + names = ", ".join(sorted(unknown_keys)) + raise ValueError(f"Unknown config key(s) for model-local Component {component_id}: {names}") + mode = component.capture_mode + if mode is ModelLocalCaptureMode.PRECAPTURE and "max_extra_graphs" in raw_component: + raise ValueError(f"max_extra_graphs requires a lazy capture mode for Component {component_id}") + if mode is ModelLocalCaptureMode.PURE_LAZY and not component.descriptors: + raise ValueError(f"Pure lazy Component {component_id} needs descriptors for memory profiling") + max_extra = self._validate_nonnegative(raw_component.get("max_extra_graphs", 0)) + if mode.allows_lazy_capture and max_extra == 0: + logger.warning( + "Model-local CUDA Graph Component %s has unbounded lazy capture; memory profiling cannot reserve " + "for every future runtime Descriptor", + component_id, + ) + self._component_configs[component_id] = ModelLocalComponentRuntimeConfig( + max_extra_graphs=max_extra, + ) + + self.components = components + self._prepared = True + logger.info( + "Prepared runner-owned model-local CUDA Graph Components: %s", + [component.component_id for component in components], + ) + + def capture_entry( + self, + component: ModelLocalCUDAGraphComponent, + descriptor: ModelLocalCUDAGraphDescriptor, + *, + graph_pool: object | None = None, + ) -> ModelLocalCUDAGraphEntry | None: + if not component.validate_descriptor(descriptor): + return None + routine = component.routine + buffers = routine.allocate_buffers(descriptor, self.device) + num_warmups = max( + 1, + int(getattr(self.vllm_config.compilation_config, "cudagraph_num_of_warmups", 0)), + ) + for _ in range(num_warmups): + with routine.capture_context(descriptor, buffers): + routine.forward_for_capture(buffers) + capture_stream = torch.cuda.current_stream(self.device) + capture_stream.synchronize() + graph = torch.cuda.CUDAGraph() + try: + with routine.capture_context(descriptor, buffers): + with ( + torch.inference_mode(), + torch.cuda.graph( + graph, + pool=(current_platform.get_global_graph_pool() if graph_pool is None else graph_pool), + stream=capture_stream, + ), + ): + captured_output = routine.forward_for_capture(buffers) + except BaseException: + graph.reset() + raise + return ModelLocalCUDAGraphEntry( + descriptor=descriptor, + graph=graph, + buffers=buffers, + captured_output=captured_output, + ) + + def profile_memory(self) -> int: + """Estimate startup graph memory with throwaway captures.""" + if not self._prepared: + raise RuntimeError("ModelLocalCUDAGraphManager must be prepared before profiling") + + profiling_pool = current_platform.graph_pool_handle() + captured: list[ModelLocalCUDAGraphEntry] = [] + estimate = 0 + try: + for component in self.components: + component_config = self._component_configs.get(component.component_id) + if component_config is None or not component.descriptors: + continue + samples: list[int] = [] + for descriptor in component.capture_descriptors[:2]: + free_before = self._synchronized_free_memory() + entry = self.capture_entry(component, descriptor, graph_pool=profiling_pool) + if entry is None: + continue + free_after = self._synchronized_free_memory() + captured.append(entry) + samples.append(max(0, free_before - free_after)) + if not samples: + continue + first_capture = samples[0] + per_graph = max(samples[1] if len(samples) > 1 else 0, 1 << 20) + if component.capture_mode is ModelLocalCaptureMode.PURE_LAZY: + extra_graphs = max(0, component_config.max_extra_graphs - 1) + else: + lazy_graphs = component_config.max_extra_graphs if component.capture_mode.allows_lazy_capture else 0 + extra_graphs = len(component.capture_descriptors) - 1 + lazy_graphs + estimate += first_capture + per_graph * extra_graphs + logger.debug( + "Estimated model-local Component %s CUDA graph memory: " + "%.2f MiB first-capture + %d x %.2f MiB per-graph", + component.component_id, + first_capture / (1 << 20), + extra_graphs, + per_graph / (1 << 20), + ) + finally: + for entry in captured: + self._destroy_entry(entry) + torch.accelerator.synchronize() + torch.accelerator.empty_cache() + logger.info("Estimated runner-owned vocoder CUDA graph memory: %.2f MiB", estimate / (1 << 20)) + return estimate + + def _capture_and_register( + self, + managed: ManagedComponent, + descriptor: ModelLocalCUDAGraphDescriptor, + ) -> ModelLocalCUDAGraphEntry | None: + existing = managed.entries.get(descriptor) + if existing is not None: + return existing + + entry = self.capture_entry(managed.component, descriptor) + if entry is not None: + managed.entries[descriptor] = entry + return entry + + @staticmethod + def _destroy_entry(entry: ModelLocalCUDAGraphEntry) -> None: + entry.graph.reset() + entry.buffers = None + entry.captured_output = None + + def _available_descriptors( + self, + managed: ManagedComponent, + ) -> frozenset[ModelLocalCUDAGraphDescriptor]: + return frozenset(managed.entries.keys()) + + @contextmanager + def _runtime_capture_scope(self): + if self._runtime_capture_stream is None: + self._runtime_capture_stream = torch.cuda.Stream(device=self.device) + caller = torch.cuda.current_stream(self.device) + ready = torch.cuda.Event() + ready.record(caller) + self._runtime_capture_stream.wait_event(ready) + set_cudagraph_capturing_enabled(True) + try: + with torch.cuda.stream(self._runtime_capture_stream): + yield + complete = torch.cuda.Event() + complete.record(self._runtime_capture_stream) + caller.wait_event(complete) + finally: + set_cudagraph_capturing_enabled(False) + + def _runtime_capture_and_register( + self, + managed: ManagedComponent, + descriptor: ModelLocalCUDAGraphDescriptor, + ) -> ModelLocalCUDAGraphEntry | None: + if not managed.capture_mode.allows_lazy_capture: + return None + if torch.cuda.is_current_stream_capturing(): + return None + existing = managed.entries.get(descriptor) + if existing is not None: + return existing + if managed.max_graphs is not None and len(managed.entries) >= managed.max_graphs: + return None + with self._runtime_capture_scope(): + entry = self._capture_and_register(managed, descriptor) + if entry is not None: + logger.info( + "Lazy-captured model-local CUDA Graph Component %s Descriptor %r", + managed.component.component_id, + descriptor, + ) + return entry + + def _make_runtime_miss_handler( + self, + managed: ManagedComponent, + ) -> Callable[[ModelLocalRuntimeResolution], ModelLocalCUDAGraphEntry | None]: + if not managed.capture_mode.allows_lazy_capture: + return lambda resolution: None + + def on_runtime_miss(resolution: ModelLocalRuntimeResolution) -> ModelLocalCUDAGraphEntry | None: + descriptor = resolution.descriptor + if descriptor is None: + descriptor = managed.component.routine.make_lazy_descriptor(resolution.runtime_key) + if descriptor is None: + return None + return self._runtime_capture_and_register(managed, descriptor) + + return on_runtime_miss + + def _build_runtime_callable( + self, + managed: ManagedComponent, + recorder: _ComponentRecorder | _NoOpRecorder, + ) -> Callable[..., Any]: + component = managed.component + entries = MappingProxyType(managed.entries) + routine = component.routine + on_runtime_miss = self._make_runtime_miss_handler(managed) + clone_output = component.clone_output + + def runtime_callable(*args: Any, **kwargs: Any) -> Any: + routine.validate_runtime_inputs(args, kwargs) + resolution = routine.resolve_runtime(args, kwargs, self._available_descriptors(managed)) + entry = entries.get(resolution.descriptor) if resolution.descriptor is not None else None + if entry is None: + entry = on_runtime_miss(resolution) + if entry is None: + recorder.record_fallback(resolution) + return routine.eager_call(*args, **kwargs) + + descriptor = entry.descriptor + graph_resolution = ( + resolution + if resolution.descriptor == descriptor + else ModelLocalRuntimeResolution( + runtime_key=resolution.runtime_key, + descriptor=descriptor, + ) + ) + try: + routine.copy_runtime_inputs(args, kwargs, entry.buffers) + entry.graph.replay() + output = routine.output_after_replay(args, kwargs, entry.buffers, entry.captured_output) + if clone_output: + output = clone_tensor_tree(output) + except Exception: + recorder.record_replay_error(graph_resolution) + raise + recorder.record_graph_hit(graph_resolution) + return output + + return runtime_callable + + def capture_and_bind(self) -> int: + # Capture every Component before binding any Handle. During capture, + # calls through other Components must stay eager; after binding, a + # unified graph's eager_call fallback can invoke segmented graph + # Components and interleave them with uncaptured eager operations. + if not self._prepared: + raise RuntimeError("ModelLocalCUDAGraphManager must be prepared before capture") + if self._capture_attempted: + raise RuntimeError("Model-local CUDA Graph capture has already been attempted") + self._capture_attempted = True + + capture_start = time.perf_counter() + free_before = self._synchronized_free_memory() + prepared: dict[str, ManagedComponent] = {} + captured: dict[str, ManagedComponent] = {} + active: dict[str, ManagedComponent] = {} + selected = [component for component in self.components if component.component_id in self._component_configs] + for component in selected: + if component._bound_handle is not None: + raise RuntimeError(f"Component already bound before capture: {component.component_id}") + try: + # Prepare every Component before capturing any of them. + for component in selected: + prepared[component.component_id] = ManagedComponent( + component=component, + entries=OrderedDict(), + capture_mode=component.capture_mode, + ) + + for component_id, managed in prepared.items(): + component = managed.component + component_config = self._component_configs[component_id] + capture_descriptors = ( + component.capture_descriptors + if component.capture_mode is not ModelLocalCaptureMode.PURE_LAZY + else () + ) + progress = ( + tqdm( + total=len(capture_descriptors), + desc=f"Capture {component.component_id}", + unit="graph", + leave=True, + ) + if capture_descriptors + else None + ) + try: + for descriptor in capture_descriptors: + try: + self._capture_and_register(managed, descriptor) + finally: + assert progress is not None + progress.update(1) + finally: + if progress is not None: + progress.close() + logger.info( + "Model-local CUDA Graph Component %s captured %d/%d startup Descriptors", + component.component_id, + len(managed.entries), + len(capture_descriptors), + ) + # Zero means no runtime graph-count limit, not zero lazy slots. + managed.max_graphs = ( + None + if managed.capture_mode.allows_lazy_capture and component_config.max_extra_graphs == 0 + else len(managed.entries) + component_config.max_extra_graphs + ) + if not managed.entries and not managed.capture_mode.allows_lazy_capture: + continue + captured[component_id] = managed + + for component_id, managed in captured.items(): + component = managed.component + runtime_callable = self._build_runtime_callable( + managed, + self.stats_sink.recorder_for(component_id), + ) + component._bind_handle( + ModelLocalGraphHandle( + runtime_callable, + lambda managed=managed: self._available_descriptors(managed), + ) + ) + active[component_id] = managed + except BaseException: + for managed in prepared.values(): + managed.component._restore_eager() + for entry in managed.entries.values(): + self._destroy_entry(entry) + managed.entries.clear() + raise + + self.managed_components = active + free_after = self._synchronized_free_memory() + captured_memory = max(0, free_before - free_after) + logger.info( + "Model-local CUDA Graph capture finished in %.2fs, bound=%s, memory=%.2f MiB", + time.perf_counter() - capture_start, + list(active), + captured_memory / (1 << 20), + ) + return captured_memory + + def clear(self) -> None: + for managed in self.managed_components.values(): + managed.component._restore_eager() + for entry in managed.entries.values(): + self._destroy_entry(entry) + managed.entries.clear() + self.managed_components.clear() + self._runtime_capture_stream = None + if self.stats_sink.enabled: + logger.info("Model-local CUDA Graph runtime stats: %s", self.stats_sink.snapshot()) diff --git a/vllm_omni/worker_v2/omni_generation_model_runner.py b/vllm_omni/worker_v2/omni_generation_model_runner.py index 68e16f42e38..5cd00f98630 100644 --- a/vllm_omni/worker_v2/omni_generation_model_runner.py +++ b/vllm_omni/worker_v2/omni_generation_model_runner.py @@ -12,15 +12,18 @@ from __future__ import annotations from dataclasses import replace -from typing import Any +from typing import Any, cast import torch +from vllm.compilation.monitor import set_cudagraph_capturing_enabled from vllm.config.compilation import CUDAGraphMode +from vllm.distributed.parallel_state import graph_capture from vllm.forward_context import set_forward_context from vllm.logger import init_logger from vllm.model_executor.layers.fused_moe.all2all_utils import ( get_ep_all2all_manager, ) +from vllm.utils.gc_utils import freeze_gc_for_cudagraph_capture from vllm.utils.torch_utils import PIN_MEMORY from vllm.v1.core.sched.output import GrammarOutput, SchedulerOutput from vllm.v1.outputs import AsyncModelRunnerOutput, ModelRunnerOutput @@ -29,11 +32,17 @@ ExecuteModelState, IntermediateTensors, ) +from vllm.v1.worker.workspace import lock_workspace from vllm_omni.core.sched.output import OmniCachedRequestData, OmniNewRequestData from vllm_omni.data_entry_keys import flatten_payload +from vllm_omni.model_executor.models.interfaces.model_local_cudagraph import ( + SupportsModelLocalCUDAGraph, + supports_model_local_cudagraph, +) from vllm_omni.model_executor.models.output_templates import OmniOutput from vllm_omni.outputs import OmniModelRunnerOutput +from vllm_omni.worker.model_local_cudagraph_manager import ModelLocalCUDAGraphManager from vllm_omni.worker_v2.omni_ar_model_runner import ( _async_copy_mm, _ensure_tensor_values, @@ -156,12 +165,80 @@ class OmniGenerationModelRunner(OmniGPUModelRunner): def __init__(self, *args: Any, **kwargs: Any) -> None: super().__init__(*args, **kwargs) + self.model_local_cudagraph_manager: ModelLocalCUDAGraphManager | None = None self._gen_model_output: Any = None self._gen_input_batch: Any = None # Placeholder for ExecuteModelState.hidden_states — allocated # once and reused every step to avoid per-forward allocation. self._dummy_hidden = torch.zeros(1, dtype=self.dtype, device=self.device) + def load_model(self, *args: Any, **kwargs: Any) -> None: + super().load_model(*args, **kwargs) + if ( + self.model_config.enforce_eager + or self.compilation_config.cudagraph_mode == CUDAGraphMode.NONE + or not getattr(self.model_config, "model_local_cudagraph", None) + or not supports_model_local_cudagraph(self.get_model()) + ): + return + + manager = ModelLocalCUDAGraphManager(vllm_config=self.vllm_config, device=self.device) + manager.prepare(cast(SupportsModelLocalCUDAGraph, self.get_model())) + self.model_local_cudagraph_manager = manager + logger.info("Initialized runner-owned model-local CUDA Graph manager for MRV2 generation") + + @torch.inference_mode() + def profile_cudagraph_memory(self) -> int: + manager = self.model_local_cudagraph_manager + if manager is None: + return super().profile_cudagraph_memory() + + set_cudagraph_capturing_enabled(True) + try: + with freeze_gc_for_cudagraph_capture(), graph_capture(device=self.device): + torch.accelerator.synchronize() + torch.accelerator.empty_cache() + return manager.profile_memory() + finally: + set_cudagraph_capturing_enabled(False) + + def needs_cudagraph_capture(self) -> bool: + if self.model_local_cudagraph_manager is not None: + return True + return super().needs_cudagraph_capture() + + @torch.inference_mode() + def capture_model(self) -> int: + manager = self.model_local_cudagraph_manager + if manager is None: + return super().capture_model() + + set_cudagraph_capturing_enabled(True) + try: + with freeze_gc_for_cudagraph_capture(), graph_capture(device=self.device): + torch.accelerator.synchronize() + torch.accelerator.empty_cache() + free_before = torch.accelerator.get_memory_info()[0] + manager.capture_and_bind() + torch.accelerator.synchronize() + free_after = torch.accelerator.get_memory_info()[0] + finally: + set_cudagraph_capturing_enabled(False) + + torch.accelerator.synchronize() + torch.accelerator.empty_cache() + lock_workspace() + captured_bytes = max(0, free_before - free_after) + logger.info("MRV2 model-local CUDA Graph capture finished (%.2f MiB)", captured_bytes / (1 << 20)) + return captured_bytes + + def shutdown(self) -> None: + manager = self.model_local_cudagraph_manager + if manager is not None: + manager.clear() + self.model_local_cudagraph_manager = None + super().shutdown() + # ------------------------------------------------------------------ # Async-chunk support: replace prompt_token_ids for cached requests # ------------------------------------------------------------------