Repository navigation
fix(models): give back all the memory an idle unload claims to free #2677
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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() | ||
|
Comment on lines
+4595
to
+4599
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
|
||
|
|
||
|
|
||
| def unload_diarization_pipeline() -> bool: | ||
| """Release a resident pyannote pipeline after the runtime changes.""" | ||
| global _diar_pipeline | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Checking only
_fi_graph_cacheskips FlashInfer cleanup after a runtime fallback or RAM offload, because_unapply_flashinfer()removes that attribute without clearing_CTX. The module can still hold the last attention wrapper and GPU tensors, so the new unload path leaves that memory behind. Track outstanding FlashInfer state separately from the current patch, and clear it on the inference thread before flushing.