diff --git a/CHANGELOG.md b/CHANGELOG.md index 69bca6c85..def29cd45 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -30,6 +30,10 @@ metadata and the backend fallback mirror it. ### Fixed +- The voice model really leaves GPU memory after the idle timeout when the OpenAI-compatible speech API, an engine self-test or a remote worker used it last, instead of keeping about 3 GB while showing as unloaded (#2677) +- Steady OpenAI-compatible speech traffic no longer counts as idle, so the voice model is not unloaded under it and then loaded a second time (#2677) +- Unloading a voice model compiled with CUDA graphs, or running FlashInfer, also frees its graph memory and attention workspace (#2677) +- Word-alignment models and the speaker-diarization pipeline are released after the idle timeout, and alignment models also when transcription unloads (#2677) - MCP speech tools wait through model loading and progress-extended CPU renders instead of timing out before the backend (#2609) - On Windows, the GPU report finds your graphics card again, and CPU-only hosts with integrated graphics are no longer told to fix an NVIDIA driver (#2620) — thanks @creatorliao! - Stretch Video exports that keep the original background no longer fail with HTTP 409 when subtitle cues sit close together or have no length (#2616) — thanks @quan0pek! diff --git a/backend/services/asr_backend.py b/backend/services/asr_backend.py index f79e1c0f1..65e20f318 100644 --- a/backend/services/asr_backend.py +++ b/backend/services/asr_backend.py @@ -563,6 +563,8 @@ def ensure_module(self, stacklevel): #: aligner is independent of whatever produced the segments, so MLX (which #: transcribes on the GPU) reuses exactly the aligner WhisperX would have used. _ALIGN_CACHE: dict[tuple[str, str], object] = {} +#: Last aligner use, on the monotonic clock (idle release, like the capture ASR). +_align_last_used: float = 0.0 #: Forced alignment is torch/wav2vec2 (not CTranslate2), so unlike Whisper itself #: it *can* run on MPS — measured on an M2: 20.3 s vs 28.4 s for a 30 s chunk, with @@ -580,6 +582,8 @@ def load_align_model(language_code: str, device: str): language — WhisperX bundles them for ~20 major languages only — or it cannot be loaded on ``device``. The caller then keeps Whisper's own (looser) word timestamps, after trying any fallback device.""" + global _align_last_used + _align_last_used = time.monotonic() key = (language_code, device) if key in _ALIGN_CACHE: return _ALIGN_CACHE[key] @@ -600,6 +604,38 @@ def load_align_model(language_code: str, device: str): return _ALIGN_CACHE[key] +def release_align_models() -> int: + """Drop every loaded aligner; returns how many were released. + + Keeps the ``None`` entries: those record that a language has no aligner + (or none that loads on that device) and hold no memory, and forgetting + them would only make the next transcription probe again. A transcription + mid-alignment keeps its own reference and finishes. + """ + released = 0 + for key, value in list(_ALIGN_CACHE.items()): + if value is not None: + _ALIGN_CACHE.pop(key, None) + released += 1 + return released + + +def release_idle_align_models(idle_s: float, *, now: float | None = None) -> bool: + """Release the aligners once none has been used for ``idle_s`` seconds. + + They load on the first aligned transcription — up to ~1.2 GB per language + — and were otherwise kept for the life of the process. Returns True when + any were released; the caller flushes the device cache. + """ + now = time.monotonic() if now is None else now + if now - _align_last_used < idle_s: + return False + released = release_align_models() + if released: + logger.info("Idle timeout reached. Released %d forced-alignment model(s).", released) + return released > 0 + + def forced_align(segments: list, audio, language_code: str, device: str | None = None) -> list: """Snap word boundaries to the audio with wav2vec2 forced alignment. @@ -666,7 +702,6 @@ class WhisperXBackend(ASRBackend): def __init__(self): self._model_name = os.environ.get("ASR_MODEL_WHISPERX", "large-v3") self._asr = None - self._align_cache = {} # language_code → (align_model, metadata) self._device, self._compute_type = self._pick_device() @staticmethod @@ -1100,7 +1135,10 @@ def transcribe(self, audio_path: str, *, word_timestamps: bool = True, def unload(self) -> None: self._asr = None - self._align_cache.clear() + # The aligners live in the module cache, shared with MLX Whisper; a + # per-instance dict here was cleared on every unload but never filled, + # so they stayed resident. + release_align_models() import gc gc.collect() try: @@ -1417,6 +1455,8 @@ def warmup(self) -> None: def unload(self) -> None: + # The wav2vec2 aligners forced_align() loaded for this backend. + release_align_models() # mlx-whisper owns the weights in a library-level singleton, not on # this wrapper. Dropping the wrapper alone retains unified memory. import sys diff --git a/backend/services/model_manager.py b/backend/services/model_manager.py index beaa0bc3b..5b55afa72 100644 --- a/backend/services/model_manager.py +++ b/backend/services/model_manager.py @@ -1973,6 +1973,11 @@ def _resolve_compile_mode() -> str: _compiled_inference_executor: "ThreadPoolExecutor | None" = None _compiled_inference_thread_ident: "int | None" = None +# How long an unload waits for the compiled-inference thread to run the reset +# before flushing without it; the reset still runs, and flushes, once the +# render holding the thread finishes. Bounded because idle_worker unloads on +# the event loop. +_COMPILE_RESET_WAIT_S = 5.0 def _get_compiled_inference_executor() -> ThreadPoolExecutor: @@ -3223,6 +3228,24 @@ async def get_model(*, allow_load: bool = True): return model +def acquire_resident_model(): + """The resident shared model for a synchronous caller, or None. Never loads. + + ``OmniVoiceBackend``'s warm path: it runs on a pool thread, cannot await + :func:`get_model`, and must not keep the model between calls. It still + has to touch the idle clock — the cached adapter used to skip it, so + steady ``/v1/audio/speech`` traffic looked idle and ``idle_worker`` + unloaded the model under it — and heal placement, as the warm half of + ``get_model()`` does. + """ + global _last_used + _last_used = time.time() + resident = model + if resident is not None: + ensure_tts_on_device() + return resident + + def _make_room_before_tts_load() -> None: """Evict-then-load: free what we already own before a tight TTS load. @@ -3559,6 +3582,19 @@ async def idle_worker(): release_idle_models(idle_timeout) except Exception: # noqa: BLE001 — the reaper must never kill idle_worker logger.warning("idle watermark-model release failed", exc_info=True) + # The wav2vec2 aligners and the pyannote pipeline: loaded by the first + # transcription that needed them, then held for the life of the process. + try: + from services.asr_backend import release_idle_align_models + + if release_idle_align_models(idle_timeout): + free_vram() + except Exception: # noqa: BLE001 — the reaper must never kill idle_worker + logger.warning("idle aligner release failed", exc_info=True) + try: + release_idle_diarization_pipeline(idle_timeout) + except Exception: # noqa: BLE001 — the reaper must never kill idle_worker + logger.warning("idle diarization release failed", exc_info=True) def release_tts_side_caches(): """Drop caches keyed to the TTS model, for when the model itself is released. @@ -3724,17 +3760,72 @@ def unload_shared_model() -> bool: ``_model_lock`` simply keep holding it across the call. Assignment is GIL-atomic, so the worst a race costs is a redundant reload. Idempotent — returns False when nothing was resident. + + A compiled or FlashInfer model leaves state outside the model that holds + memory too; that goes before the flush (``_release_inference_state``). """ global model, _ram_offload if model is None: return False + compiled = getattr(getattr(model, "llm", None), "_orig_mod", None) is not None + flashinfer = getattr(model, "_fi_graph_cache", None) is not None model = None _ram_offload = None release_tts_side_caches() - free_vram() + if compiled or flashinfer: + _release_inference_state(compiled=compiled, flashinfer=flashinfer) + else: + free_vram() return True +def _release_inference_state(*, compiled: bool, flashinfer: bool) -> None: + """Drop what a compiled or FlashInfer model leaves behind, then flush. + + Dropping the model does not free ``torch.compile``'s state: Dynamo's code + caches outlive it, and ``mode="reduce-overhead"`` keeps its CUDA-graph + memory pools in Inductor's cudagraph trees until ``torch._dynamo.reset()``. + Those trees are thread-local — a reset from any other thread asserts out + in ``reset_cudagraph_trees()`` and the pools stay allocated — so it runs on + the single thread that captured them (#315). FlashInfer's module context + keeps the last attention workspace, and its attention calls read that + context mid-render, so it is cleared on the same thread. + + Queued behind a render still on that thread rather than racing it. If one + holds the thread past ``_COMPILE_RESET_WAIT_S`` the unload flushes what is + free now, and the queued reset flushes again when it gets its turn. + """ + torch = _lazy_torch() + + def _release(): + if compiled: + try: + torch._dynamo.reset() + except Exception: # noqa: BLE001 — freeing memory must never raise + logger.warning("could not reset torch.compile state on unload", exc_info=True) + if flashinfer: + fi = sys.modules.get("omnivoice.models.omnivoice_flashinfer") + ctx = getattr(fi, "_CTX", None) + if isinstance(ctx, dict): + for key in list(ctx): + ctx[key] = None + free_vram() + + executor = _compiled_inference_executor + if executor is None or threading.get_ident() == _compiled_inference_thread_ident: + _release() + return + released = executor.submit(_release) + try: + released.result(timeout=_COMPILE_RESET_WAIT_S) + except TimeoutError: + logger.info( + "A render still holds the compiled-model thread; its CUDA graphs are " + "released when it finishes." + ) + free_vram() + + def _has_dedicated_vram(): """Check if the current device has limited dedicated VRAM that needs offloading.""" torch = _lazy_torch() @@ -4334,6 +4425,9 @@ def _note_gpu_pool_idle() -> None: _diar_pipeline = None +# Monotonic, like the capture-ASR and watermark clocks: an NTP step or a laptop +# resume must not read as an idle timeout. +_diar_last_used: float = 0.0 # Sentinel error classes used by callers (dub_core) to decide whether to # emit a structured SSE warning with a docs deeplink. Kept as module-level @@ -4426,7 +4520,7 @@ def get_diarization_pipeline(return_error: bool = False): streaming `_diarize` path uses to emit a structured SSE warning with a docs deeplink — issue #78. """ - global _diar_pipeline + global _diar_pipeline, _diar_last_used from services.diarization_runtime import SORTFORMER, selected_backend if selected_backend() == SORTFORMER: try: @@ -4436,6 +4530,7 @@ def get_diarization_pipeline(return_error: bool = False): except Exception as exc: logger.exception("Could not prepare native Sortformer") return (None, _classify_diarization_error(exc)) if return_error else None + _diar_last_used = time.monotonic() if _diar_pipeline is not None: return (_diar_pipeline, None) if return_error else _diar_pipeline @@ -4489,6 +4584,21 @@ def get_diarization_pipeline(return_error: bool = False): return (None, err_class) if return_error else None +def release_idle_diarization_pipeline(idle_s: float, *, now: float | None = None) -> bool: + """Release the pyannote pipeline once it has gone unused for ``idle_s``. + + It loaded on the first diarized transcription and then stayed for the life + of the process. A dub still diarizing past the timeout keeps its own + reference and finishes; only the cache lets go. Returns True when a + pipeline was released. + """ + now = time.monotonic() if now is None else now + if _diar_pipeline is None or now - _diar_last_used < idle_s: + return False + logger.info("Idle timeout reached. Releasing the speaker diarization pipeline.") + return unload_diarization_pipeline() + + def unload_diarization_pipeline() -> bool: """Release a resident pyannote pipeline after the runtime changes.""" global _diar_pipeline diff --git a/backend/services/tts_backend.py b/backend/services/tts_backend.py index b2fbae6ec..78de228c7 100644 --- a/backend/services/tts_backend.py +++ b/backend/services/tts_backend.py @@ -22,6 +22,7 @@ import logging import os import re +import sys import threading import time from abc import ABC, abstractmethod @@ -1503,28 +1504,47 @@ class OmniVoiceBackend(TTSBackend): ref_strategy = "best_window" def __init__(self, model=None): - # The live OmniVoice instance. Reuses the singleton owned by - # model_manager so memory isn't doubled. + # Set only for a per-call view: the native routes pass the model they + # just got from get_model(). A cached adapter (_active_instance, + # _ENGINE_INSTANCES, a TTS stream) leaves it None and never fills it: + # model_manager owns the shared model, and a second owner is how an + # idle-unloaded model stayed resident — and kept generating — while + # /model/status said idle. self._model = model + def _resident_model(self): + """The model this adapter would run on now, without loading one. + + Looked up through ``sys.modules``: a diagnostic read must not import + model_manager, and if it was never imported nothing is resident. + """ + if self._model is not None: + return self._model + return getattr(sys.modules.get("services.model_manager"), "model", None) + + def execution_evidence_loaded(self) -> bool: + return self._resident_model() is not None + @property def execution_device(self) -> str | None: """Actual device of the shared model, for live engine diagnostics.""" - if self._model is None: + model = self._resident_model() + if model is None: return None try: - return str(next(self._model.parameters()).device) + return str(next(model.parameters()).device) except Exception: # noqa: BLE001 - third-party model wrappers vary - device = getattr(self._model, "device", None) + device = getattr(model, "device", None) return str(device) if device is not None else None @property def dtype(self) -> str | None: """Actual parameter precision of the shared model when resident.""" - if self._model is None: + model = self._resident_model() + if model is None: return None try: - return str(next(self._model.parameters()).dtype) + return str(next(model.parameters()).dtype) except Exception: # noqa: BLE001 - diagnostics must remain best effort return None @@ -1540,9 +1560,10 @@ def is_available(cls) -> tuple[bool, str]: @property def sample_rate(self) -> int: - if self._model is None: + model = self._resident_model() + if model is None: return self._DEFAULT_SAMPLE_RATE # canonical OmniVoice rate - return getattr(self._model, "sampling_rate", 24000) + return getattr(model, "sampling_rate", 24000) @property def supported_languages(self) -> list[str]: @@ -1550,17 +1571,26 @@ def supported_languages(self) -> list[str]: return ["multi"] def _ensure_loaded(self): + """Return the model to run this call on — never stored on ``self``. + + The next unload has to be able to free it, and the call after that has + to run on whatever model_manager holds then, not on the copy this + adapter saw first. + """ if self._model is not None: - # The cached instance skips get_model(), and with it the placement - # heal: put the shared model back on its device if the opt-in - # post-generation offload (#2618) or an unbalanced ASR offload - # (#1191) left it in RAM. One parameter probe when it is in place. + # A per-call view skips get_model() from here on, and with it the + # placement heal: put the shared model back on its device if the + # opt-in post-generation offload (#2618) or an unbalanced ASR + # offload (#1191) left it in RAM. One parameter probe when in place. from services.model_manager import ensure_tts_on_device ensure_tts_on_device() - return - # Reuse model_manager's cached instance so we don't double-load. - from services.model_manager import get_model + return self._model + from services.model_manager import acquire_resident_model, get_model + + model = acquire_resident_model() + if model is not None: + return model import asyncio # Caller is sync; spin up a fresh loop if needed. get_running_loop() # raises only when *no* loop is running — that's the safe path where @@ -1568,15 +1598,14 @@ def _ensure_loaded(self): try: asyncio.get_running_loop() except RuntimeError: - self._model = asyncio.run(get_model()) - return + return asyncio.run(get_model()) raise RuntimeError( "OmniVoiceBackend.generate() called inside an async context without a pre-loaded model. " "Pass `model=await get_model()` to the constructor." ) def generate(self, text, **kw) -> torch.Tensor: - self._ensure_loaded() + model = self._ensure_loaded() language = kw.get("language") ref_audio = kw.get("ref_audio") ref_text = kw.get("ref_text") @@ -1603,7 +1632,7 @@ def generate(self, text, **kw) -> torch.Tensor: # one place here and a subtly different one there is exactly how the cache # came to be wired into the adapter and nowhere else. audios = generate_with_cached_ref( - self._model, ref_audio=ref_audio, ref_text=ref_text, **gen_kw + model, ref_audio=ref_audio, ref_text=ref_text, **gen_kw ) return audios[0] @@ -1615,7 +1644,7 @@ def generate_batch(self, texts: list[str], **kw) -> list[torch.Tensor]: model together; an incomplete prompt batch falls back to the proven single-item path instead of changing synthesis semantics. """ - self._ensure_loaded() + model = self._ensure_loaded() if not texts: return [] from services.model_manager import tts_inference @@ -1623,9 +1652,9 @@ def generate_batch(self, texts: list[str], **kw) -> list[torch.Tensor]: # Holds the shared model in place for prompt encoding + the batch # generate (an offload can't move it mid-batch; #2618). with tts_inference(): - return self._generate_batch_on_model(texts, **kw) + return self._generate_batch_on_model(model, texts, **kw) - def _generate_batch_on_model(self, texts: list[str], **kw) -> list[torch.Tensor]: + def _generate_batch_on_model(self, model, texts: list[str], **kw) -> list[torch.Tensor]: def _items(value): if isinstance(value, list): return value @@ -1648,7 +1677,7 @@ def _item_kwargs(index): prompts = [] break prompt = _get_clone_prompt( - self._model, + model, ref_audio, ref_text, preprocess_prompt, @@ -1679,7 +1708,7 @@ def _item_kwargs(index): else: gen_kw["ref_audio"] = None gen_kw["ref_text"] = None - return self._model.generate(text=texts, **gen_kw) + return model.generate(text=texts, **gen_kw) def unload(self) -> None: """Release the OmniVoice model (MM2-02). OmniVoice shares the singleton diff --git a/docs/engines/mlx-whisper.md b/docs/engines/mlx-whisper.md index eb8abcdeb..17d818c93 100644 --- a/docs/engines/mlx-whisper.md +++ b/docs/engines/mlx-whisper.md @@ -40,7 +40,8 @@ platforms use the CUDA/CPU engines instead. - `OMNIVOICE_ALIGN_DEVICE` — force the wav2vec2 aligner's device. The aligner runs on MPS when it can and falls back to CPU; languages without a bundled aligner (~20 major languages have one) keep Whisper's native word - timestamps. + timestamps. Loaded aligners are released when the engine unloads and after + the idle timeout (`OMNIVOICE_IDLE_TIMEOUT_S`). ## Quirks diff --git a/docs/engines/whisperx.md b/docs/engines/whisperx.md index b8235d895..e493d0edd 100644 --- a/docs/engines/whisperx.md +++ b/docs/engines/whisperx.md @@ -35,7 +35,8 @@ prefers it wherever CTranslate2 can use the GPU. download on first load — see [downloading-models](../downloading-models.md). - `OMNIVOICE_ALIGN_DEVICE` — force the wav2vec2 aligner's device. Aligners exist for ~20 major languages; other languages keep Whisper's native word - timestamps instead of failing. + timestamps instead of failing. Loaded aligners are released when the engine + unloads and after the idle timeout (`OMNIVOICE_IDLE_TIMEOUT_S`). ## VRAM preflight and degradation diff --git a/docs/performance.md b/docs/performance.md index 63c137c4e..0088a7320 100644 --- a/docs/performance.md +++ b/docs/performance.md @@ -89,7 +89,7 @@ None of them are required — the defaults are chosen for the common case. | `CUDA_VISIBLE_DEVICES` | all NVIDIA GPUs | On a multi-GPU NVIDIA host, expose only the selected physical adapter to VoiceStudio and its engine subprocesses. **Settings → Performance & Device → CUDA** lists adapters by index and name, persists their stable GPU UUID, and applies the choice after restart. An externally supplied environment variable wins over the saved UI choice. | | `OMNIVOICE_FLASHINFER` | `0` | CUDA-only accelerated decoding for the default engine via [FlashInfer](https://github.com/flashinfer-ai/flashinfer) kernels (packed CFG attention, fused RMSNorm/RoPE/GEMM) — ~2x on upstream's benchmarks. `1` enables it; `graph` also captures CUDA graphs (best when you render one thing at a time). Requires installing the optional `flashinfer-python` package into the backend environment first (`uv pip install flashinfer-python flashinfer-jit-cache --extra-index-url https://flashinfer.ai/whl/cu128/`, matching your CUDA build). Replaces `torch.compile` for that session, pins inference to a single GPU thread (the FlashInfer attention plan is per-generation state), and keeps fused copies of the attention/MLP weights resident (~roughly half the LLM's weight size extra VRAM) — leave it off on tight-VRAM cards. If the package is missing or a FlashInfer/CUDA-graph kernel fails at runtime, the app logs the reason and falls back to the standard path; failures outside those kernels (e.g. a genuine out-of-memory) surface normally. | | `OMNIVOICE_PROMPT_DISK_CACHE` | `1` | Persist encoded voice-clone references (`prompt_cache/` in the app data dir, ~10 KB per voice, 32 newest kept) so the first generation with a known voice after a restart skips the reference re-encode and any auto-transcription. Set `0` to keep the cache in memory only. | -| `OMNIVOICE_IDLE_TIMEOUT_S` | `900` | Seconds of idle before the TTS model unloads to free memory. Raise it (e.g. `3600`) if you generate in bursts and dislike the ~8 s reload; lower it on tight-memory machines. | +| `OMNIVOICE_IDLE_TIMEOUT_S` | `900` | Seconds of idle before the TTS model unloads to free memory. The dictation ASR, watermark models, word-alignment models and speaker-diarization pipeline are each released after the same idle time. Raise it (e.g. `3600`) if you generate in bursts and dislike the ~8 s reload; lower it on tight-memory machines. | | `OMNIVOICE_OFFLOAD_AFTER_GENERATION` | off | `1` moves the built-in TTS model to system RAM once generation finishes, and back on the next generation. Same toggle as **Settings → Performance & Device → Memory management** (the env var wins over the UI). See [Offload to RAM after generation](#offload-to-ram-after-generation). | | `OMNIVOICE_OFFLOAD_AFTER_GENERATION_GRACE_S` | `3` | How long the GPU must stay idle after a generation before that offload runs. | | `OMNIVOICE_SIDECAR_IDLE_TIMEOUT_S` | `300` | Same idea for sidecar engines (IndexTTS 2.5 etc.). | diff --git a/docs/remote-workers.md b/docs/remote-workers.md index c3578ddc8..be783e8a0 100644 --- a/docs/remote-workers.md +++ b/docs/remote-workers.md @@ -356,7 +356,7 @@ same model and they are configured separately: | Timer | Default | Set with | |---|---|---| | Engine registry — drops the cached engine instance and, for VoiceStudio, the shared model with it | 600 s | `OMNIVOICE_ENGINE_IDLE_UNLOAD_SECONDS` | -| In-process model reaper — the backstop, also releases the dictation ASR and the watermark models | 900 s | `OMNIVOICE_IDLE_TIMEOUT` (or Settings) | +| In-process model reaper — the backstop, also releases the dictation ASR, the watermark models, the word-alignment models and the speaker-diarization pipeline | 900 s | `OMNIVOICE_IDLE_TIMEOUT` (or Settings) | In practice the first one gets there first and the second finds nothing to do. Shortening only `OMNIVOICE_ENGINE_IDLE_UNLOAD_SECONDS` is the right move when diff --git a/tests/test_idle_unload_releases_memory.py b/tests/test_idle_unload_releases_memory.py new file mode 100644 index 000000000..1dc59d47e --- /dev/null +++ b/tests/test_idle_unload_releases_memory.py @@ -0,0 +1,461 @@ +"""An idle unload has to give the memory back, not just report that it did. + +Seen on an RTX 5090 Docker host: ``GET /model/status`` answered +``{"status": "idle", "loaded": false}`` while the backend still held ~3 GB of +VRAM. ``unload_shared_model()`` cleared ``model_manager.model`` — but that was +not the only reference, and the model was not the only thing left behind: + +* The cached ``OmniVoiceBackend`` adapters (``tts_backend._active_instance`` + for ``/v1/audio/speech``, ``_ENGINE_INSTANCES`` for the engine self-test, + explicit-engine ``/generate`` and worker assignments) stored the model on + their first generate and kept it. The adapter then skipped ``get_model()`` + for good, so it never touched the idle clock either: an API-only user had + the model "unloaded" mid-traffic, the adapter carried on with the orphan, + and the next native generate loaded a second copy beside it. +* ``torch.compile(mode="reduce-overhead")`` keeps CUDA-graph pools and + compiled-code caches that outlive the model. Nothing reset them — and the + reset only works on the thread that captured the graphs. +* FlashInfer's module-global context kept the last attention workspace. +* The wav2vec2 aligners (``asr_backend._ALIGN_CACHE``) were never released: + ``WhisperXBackend.unload()`` cleared a per-instance dict nothing ever filled. +* The pyannote diarization pipeline had a manual unload but no idle release. + +Fail-before: the adapters kept the unloaded model (and kept using it after a +reload), the warm adapter path left ``_last_used`` alone, no unload path ran +``torch._dynamo.reset()`` or cleared FlashInfer's context, WhisperX's unload +left the aligners resident, and ``idle_worker`` released neither the aligners +nor the diarization pipeline. +""" +from __future__ import annotations + +import asyncio +import gc +import importlib +import sys +import threading +import time +import weakref +from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace + +import pytest +import torch + + +class _Weights: + """Stands in for the model. Identity and collectability are all we need.""" + + +class _Compiled: + """What ``torch.compile`` leaves on ``model.llm``: an OptimizedModule.""" + + def __init__(self): + self._orig_mod = object() + + +@pytest.fixture +def mods(monkeypatch): + # Resolved per test: other suites pop and reimport services.*, and the + # adapter imports model_manager at call time, so patch what is live now. + tb = importlib.import_module("services.tts_backend") + mm = importlib.import_module("services.model_manager") + monkeypatch.setattr(mm, "model", None) + tb.reset_active_backend() + yield tb, mm + tb.reset_active_backend() + with tb._ENGINE_CACHE_LOCK: + tb._ENGINE_INSTANCES.pop(tb.OmniVoiceBackend, None) + tb._ENGINE_LAST_USED.pop(tb.OmniVoiceBackend, None) + + +def _quiet_manager(monkeypatch, mm, flushes=None): + """No real allocator, placement or make-room work — just the references.""" + monkeypatch.setattr(mm, "free_vram", lambda: flushes.append(mm.model) if flushes is not None else None) + monkeypatch.setattr(mm, "release_tts_side_caches", lambda: None) + monkeypatch.setattr(mm, "make_room_before_generate", lambda: None) + monkeypatch.setattr(mm, "_stranded_tts_target", lambda: None) + + +def _record_generations(monkeypatch, tb): + used: list = [] + + def _fake(model, **_kw): + used.append(model) + return [torch.zeros(1)] + + monkeypatch.setattr(tb, "generate_with_cached_ref", _fake) + return used + + +def _active_adapter(monkeypatch, tb): + monkeypatch.setattr(tb, "active_backend_id", lambda: "omnivoice") + # The in-process adapter, as on CUDA/ROCm/CPU hosts; an MPS host runs + # OmniVoice in a sidecar process that owns its own memory instead. + monkeypatch.setattr(tb, "_effective_backend_class", lambda _bid, cls, host_family=None: cls) + return tb.get_active_tts_backend() + + +def _engine_adapter(_monkeypatch, tb): + return tb.get_engine_instance(tb.OmniVoiceBackend) + + +_ADAPTERS = pytest.mark.parametrize( + "adapter", [_active_adapter, _engine_adapter], ids=["active-instance", "engine-instances"], +) + + +# ── The cached adapters must not own the shared model ────────────────────── + +@_ADAPTERS +def test_cached_adapter_lets_the_unload_free_the_model(monkeypatch, mods, adapter): + tb, mm = mods + _quiet_manager(monkeypatch, mm) + used = _record_generations(monkeypatch, tb) + weights = _Weights() + monkeypatch.setattr(mm, "model", weights) + backend = adapter(monkeypatch, tb) + + backend.generate("hello") + assert used == [weights] + + gone = weakref.ref(weights) + del weights + used.clear() + assert mm.unload_shared_model() is True + gc.collect() + + assert gone() is None, "a cached adapter still holds the model the unload released" + + +@_ADAPTERS +def test_cached_adapter_uses_the_reloaded_model_not_the_released_one(monkeypatch, mods, adapter): + """Two copies resident at once: the adapter's orphan plus a fresh load.""" + tb, mm = mods + _quiet_manager(monkeypatch, mm) + used = _record_generations(monkeypatch, tb) + monkeypatch.setattr(mm, "model", _Weights()) + backend = adapter(monkeypatch, tb) + backend.generate("before the idle unload") + + mm.unload_shared_model() + reloaded = _Weights() + monkeypatch.setattr(mm, "model", reloaded) # what the next /generate loads + backend.generate("after it") + + assert used[-1] is reloaded + + +def test_cached_adapter_cold_loads_through_the_manager_after_an_unload(monkeypatch, mods): + tb, mm = mods + _quiet_manager(monkeypatch, mm) + used = _record_generations(monkeypatch, tb) + loaded = _Weights() + + def _load(): + mm.model = loaded + return loaded + + monkeypatch.setattr(mm, "_load_model_with_timeout", lambda: asyncio.sleep(0, _load())) + monkeypatch.setattr(importlib.import_module("core.run_sentinel"), "touch_activity", lambda *_a: None) + backend = _active_adapter(monkeypatch, tb) + + backend.generate("cold") + + assert used == [loaded] + assert mm.model is loaded + + +def test_warm_cached_adapter_keeps_the_idle_clock_running(monkeypatch, mods): + """Otherwise steady /v1/audio/speech traffic looks idle to idle_worker.""" + tb, mm = mods + _quiet_manager(monkeypatch, mm) + _record_generations(monkeypatch, tb) + monkeypatch.setattr(mm, "model", _Weights()) + backend = _active_adapter(monkeypatch, tb) + backend.generate("first") + + monkeypatch.setattr(mm, "_last_used", 0.0) + backend.generate("second") + + assert mm._last_used > time.time() - 60 + + +def test_cached_adapter_reports_the_resident_model(monkeypatch, mods): + tb, mm = mods + backend = _active_adapter(monkeypatch, tb) + monkeypatch.setattr(mm, "model", SimpleNamespace(sampling_rate=48000)) + assert backend.sample_rate == 48000 + monkeypatch.setattr(mm, "model", None) + assert backend.sample_rate == tb.OmniVoiceBackend._DEFAULT_SAMPLE_RATE + + +def test_explicit_model_view_still_uses_its_model(monkeypatch, mods): + tb, mm = mods + _quiet_manager(monkeypatch, mm) + used = _record_generations(monkeypatch, tb) + monkeypatch.setattr(mm, "model", _Weights()) + pinned = _Weights() + + tb.OmniVoiceBackend(model=pinned).generate("hi") + + assert used == [pinned] + + +# ── Compiled-model state goes with the model ─────────────────────────────── + +@pytest.fixture +def infer_thread(monkeypatch, mods): + """A stand-in for the single compiled-inference thread (#315).""" + _, mm = mods + executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="compiled-infer-test") + ident = executor.submit(threading.get_ident).result() + monkeypatch.setattr(mm, "_compiled_inference_executor", executor) + monkeypatch.setattr(mm, "_compiled_inference_thread_ident", ident) + yield executor, ident + executor.shutdown(wait=True) + + +def _watch_resets(monkeypatch, mm): + """Record dynamo resets and flushes: (event, thread, what mm.model held).""" + events: list = [] + dynamo = SimpleNamespace(reset=lambda: events.append(("reset", threading.get_ident(), mm.model))) + monkeypatch.setattr(mm, "_lazy_torch", lambda: SimpleNamespace(_dynamo=dynamo)) + monkeypatch.setattr(mm, "free_vram", lambda: events.append(("free", threading.get_ident(), mm.model))) + monkeypatch.setattr(mm, "release_tts_side_caches", lambda: None) + return events + + +def _compiled_weights(): + weights = _Weights() + weights.llm = _Compiled() + return weights + + +def test_unload_resets_compiled_state_on_the_thread_that_captured_it(monkeypatch, mods, infer_thread): + """CUDA-graph trees are thread-local: reset_cudagraph_trees() from any other + thread asserts out and the graph pools stay allocated.""" + _, mm = mods + _, ident = infer_thread + events = _watch_resets(monkeypatch, mm) + monkeypatch.setattr(mm, "model", _compiled_weights()) + + assert mm.unload_shared_model() is True + + assert [e[0] for e in events] == ["reset", "free"] + assert {e[1] for e in events} == {ident} + assert all(e[2] is None for e in events), "reset/flush ran while the model was referenced" + + +def test_unload_resets_inline_without_a_compiled_inference_thread(monkeypatch, mods): + """Non-cudagraph compile modes (pre-Ampere "default") never start the thread.""" + _, mm = mods + monkeypatch.setattr(mm, "_compiled_inference_executor", None) + events = _watch_resets(monkeypatch, mm) + monkeypatch.setattr(mm, "model", _compiled_weights()) + + assert mm.unload_shared_model() is True + + assert [(e[0], e[1]) for e in events] == [ + ("reset", threading.get_ident()), ("free", threading.get_ident()), + ] + + +def test_unload_of_an_eager_model_leaves_dynamo_alone(monkeypatch, mods, infer_thread): + """A reset is process-wide; an uncompiled unload must not cost other + engines their compile caches.""" + _, mm = mods + events = _watch_resets(monkeypatch, mm) + monkeypatch.setattr(mm, "model", _Weights()) + + assert mm.unload_shared_model() is True + + assert [e[0] for e in events] == ["free"] + + +def test_unload_from_the_inference_thread_does_not_wait_on_itself(monkeypatch, mods, infer_thread): + _, mm = mods + executor, ident = infer_thread + events = _watch_resets(monkeypatch, mm) + monkeypatch.setattr(mm, "model", _compiled_weights()) + + assert executor.submit(mm.unload_shared_model).result(timeout=5) is True + + assert [(e[0], e[1]) for e in events] == [("reset", ident), ("free", ident)] + + +def test_unload_behind_a_running_render_flushes_now_and_resets_after(monkeypatch, mods, infer_thread): + """The reset queues behind the render that owns the thread; the unload must + not block the event loop on it, and the late reset still hands memory back.""" + _, mm = mods + executor, ident = infer_thread + events = _watch_resets(monkeypatch, mm) + monkeypatch.setattr(mm, "_COMPILE_RESET_WAIT_S", 0.05) + monkeypatch.setattr(mm, "model", _compiled_weights()) + render = threading.Event() + executor.submit(render.wait, 10) + + started = time.monotonic() + assert mm.unload_shared_model() is True + assert time.monotonic() - started < 5 + assert [(e[0], e[1]) for e in events] == [("free", threading.get_ident())] + + render.set() + executor.submit(lambda: None).result(timeout=5) + assert [(e[0], e[1]) for e in events[1:]] == [("reset", ident), ("free", ident)] + + +def test_unload_clears_flashinfer_context_on_the_inference_thread(monkeypatch, mods, infer_thread): + """Its attention calls read the module-global context on every layer, so it + must be cleared on the thread they run on, never under a render.""" + _, mm = mods + _, ident = infer_thread + _watch_resets(monkeypatch, mm) + writers: set = set() + + class _Ctx(dict): + def __setitem__(self, key, value): + writers.add(threading.get_ident()) + super().__setitem__(key, value) + + ctx = _Ctx(wrapper=object(), pos_ids=object(), doc_slots=object()) + monkeypatch.setitem(sys.modules, "omnivoice.models.omnivoice_flashinfer", SimpleNamespace(_CTX=ctx)) + weights = _Weights() + weights._fi_graph_cache = {} + monkeypatch.setattr(mm, "model", weights) + + assert mm.unload_shared_model() is True + + assert dict(ctx) == {"wrapper": None, "pos_ids": None, "doc_slots": None} + assert writers == {ident} + + +# ── wav2vec2 aligners ────────────────────────────────────────────────────── + +@pytest.fixture +def ab(): + return importlib.import_module("services.asr_backend") + + +def test_whisperx_unload_releases_the_shared_aligners(monkeypatch, ab): + monkeypatch.setattr(ab.WhisperXBackend, "_pick_device", staticmethod(lambda: ("cpu", "int8"))) + aligner = _Weights() + gone = weakref.ref(aligner) + monkeypatch.setattr(ab, "_ALIGN_CACHE", {("en", "cpu"): (aligner, {}), ("xx", "cpu"): None}) + del aligner + + ab.WhisperXBackend().unload() + gc.collect() + + assert gone() is None, "WhisperX unloaded but its aligner stayed resident" + # "No aligner for this language" is a memo, not memory: keep it so the + # next transcription doesn't probe for one again. + assert ab._ALIGN_CACHE == {("xx", "cpu"): None} + + +def test_idle_aligners_are_released(monkeypatch, ab): + monkeypatch.setattr(ab, "_ALIGN_CACHE", {("en", "cpu"): (object(), {})}) + monkeypatch.setattr(ab, "_align_last_used", 100.0) + + assert ab.release_idle_align_models(60.0, now=200.0) is True + assert ab._ALIGN_CACHE == {} + + +def test_recently_used_aligners_are_kept(monkeypatch, ab): + held = (object(), {}) + monkeypatch.setattr(ab, "_ALIGN_CACHE", {("en", "cpu"): held}) + monkeypatch.setattr(ab, "_align_last_used", 100.0) + + assert ab.release_idle_align_models(60.0, now=130.0) is False + assert ab._ALIGN_CACHE == {("en", "cpu"): held} + + +def test_idle_aligner_release_is_a_noop_with_only_memos(monkeypatch, ab): + monkeypatch.setattr(ab, "_ALIGN_CACHE", {("xx", "cpu"): None}) + monkeypatch.setattr(ab, "_align_last_used", 0.0) + + assert ab.release_idle_align_models(0.0, now=1e9) is False + assert ab._ALIGN_CACHE == {("xx", "cpu"): None} + + +def test_aligning_restarts_the_idle_clock(monkeypatch, ab): + monkeypatch.setattr(ab, "_ALIGN_CACHE", {("en", "cpu"): (object(), {})}) + monkeypatch.setattr(ab, "_align_last_used", 0.0) + monkeypatch.setattr(ab.time, "monotonic", lambda: 500.0) + + ab.load_align_model("en", "cpu") + + assert ab._align_last_used == 500.0 + + +# ── pyannote diarization ─────────────────────────────────────────────────── + +def test_idle_diarization_pipeline_is_released(monkeypatch, mods): + _, mm = mods + monkeypatch.setattr(mm, "free_vram", lambda: None) + monkeypatch.setattr(mm, "_diar_pipeline", object()) + monkeypatch.setattr(mm, "_diar_last_used", 100.0) + + assert mm.release_idle_diarization_pipeline(60.0, now=200.0) is True + assert mm._diar_pipeline is None + + +def test_recently_used_diarization_pipeline_is_kept(monkeypatch, mods): + _, mm = mods + pipeline = object() + monkeypatch.setattr(mm, "_diar_pipeline", pipeline) + monkeypatch.setattr(mm, "_diar_last_used", 100.0) + + assert mm.release_idle_diarization_pipeline(60.0, now=130.0) is False + assert mm._diar_pipeline is pipeline + + +def test_idle_diarization_release_is_a_noop_when_nothing_loaded(monkeypatch, mods): + _, mm = mods + monkeypatch.setattr(mm, "_diar_pipeline", None) + assert mm.release_idle_diarization_pipeline(0.0, now=1e9) is False + + +def test_getting_the_pipeline_restarts_the_idle_clock(monkeypatch, mods): + _, mm = mods + runtime = importlib.import_module("services.diarization_runtime") + monkeypatch.setattr(runtime, "selected_backend", lambda: runtime.PYANNOTE) + pipeline = object() + monkeypatch.setattr(mm, "_diar_pipeline", pipeline) + monkeypatch.setattr(mm, "_diar_last_used", 0.0) + monkeypatch.setattr(mm.time, "monotonic", lambda: 500.0) + + assert mm.get_diarization_pipeline() is pipeline + assert mm._diar_last_used == 500.0 + + +# ── idle_worker wires all of it in ───────────────────────────────────────── + +class _Stop(Exception): + pass + + +def test_idle_worker_releases_idle_aligners_and_diarization(monkeypatch, mods, ab): + _, mm = mods + calls: list = [] + monkeypatch.setattr(mm, "_resolve_idle_timeout", lambda: 42.0) + monkeypatch.setattr(mm, "free_vram", lambda: calls.append("flush")) + monkeypatch.setattr(ab, "release_idle_capture_backend", lambda _s: False) + monkeypatch.setattr(ab, "release_idle_align_models", lambda s: calls.append(("aligners", s)) or True) + monkeypatch.setattr(mm, "release_idle_diarization_pipeline", lambda s: calls.append(("diarization", s)) or True) + monkeypatch.setattr(importlib.import_module("services.watermark"), "release_idle_models", lambda _s: False) + + ticks: list = [] + + async def _one_tick(_seconds): + if ticks: + raise _Stop + ticks.append(1) + + monkeypatch.setattr(mm, "asyncio", SimpleNamespace(sleep=_one_tick)) + + with pytest.raises(_Stop): + asyncio.run(mm.idle_worker()) + + assert ("aligners", 42.0) in calls + assert calls[calls.index(("aligners", 42.0)) + 1] == "flush" + assert ("diarization", 42.0) in calls