mirror of
https://github.com/debpalash/VoiceStudio.git
synced 2026-10-02 01:26:35 +08:00
fix(memory): release device cache across supported accelerators (#2317)
* fix(memory): release the device cache on every accelerator the engines can use The dubbing and generation recovery paths open-coded the CUDA/MPS pair when flushing the device cache. Engines pick their device through torch.accelerator, so on an Ascend NPU or Intel XPU host those paths flushed nothing at all: the allocator kept the blocks the offload had just freed, and the next allocation failed with that memory still counted as in use. Extract the narrow primitive free_vram() already used (no gc.collect, no cuBLAS clear, covers CUDA/MPS/XPU/NPU, never raises) as release_device_cache() and call it from the five recovery paths. free_vram() keeps its gc + cuBLAS behaviour. Verified on an Ascend 910B4 (torch 2.15.0.dev + torch_npu 2.15.0.dev): reserved npu memory stayed at 134.0 MiB with the old pair and dropped to 0.0 MiB with the shared helper; the new tests fail 9/9 on the pre-change tree. * docs(changelog): credit the cache-flush fix with its PR ref (#2317) * fix(memory): follow the engines' accelerator and keep free_vram raising Two review findings on the cache-release helper: 1. On a hybrid host (CUDA probes as available, inference runs elsewhere) the elif chain flushed CUDA and never reached the active allocator. Ask the same question the engine sidecars ask -- torch.accelerator.current_accelerator( check_available=True) -- and flush that backend, falling back to the CUDA/MPS/XPU/NPU probe chain only when the build cannot answer. 2. free_vram() propagated empty_cache() failures before, and model_lifecycle unload callers report a failed flush to the user, so the helper takes raise_on_failure and free_vram() passes True. The direct recovery calls stay best-effort. Tests pin both, plus the cpu-answer and pre-2.6 fallbacks. * fix(memory): flush selected allocator and offload SSE cleanup * test(memory): remove source-text assertion * test(memory): trim obsolete test tail * fix(dub): flush MPS allocator after NLLB CPU fallback --------- Co-authored-by: Palash Debnath <4178343+debpalash@users.noreply.github.com>
This commit is contained in:
co-authored by
Palash Debnath
parent
c5769d7984
commit
59f37e0db0
+2
-2
@@ -12,7 +12,7 @@ metadata and the backend fallback mirror it. Archived Tauri manifests stay froze
|
||||
|
||||
- Twilio setup is a guided checklist with live status and exact commands (#2304)
|
||||
- Integration pages use the full window, with a side panel on wide screens (#2304)
|
||||
- A call agent and Calls workspace that phone someone in your own voice to get a task done (#2306, #2305)
|
||||
- A call agent that places or answers phone calls in your own voice to get a task done (#2306)
|
||||
- Footer integration logos open their in-app page (#2302)
|
||||
- Record or drop a voice sample from one view in Voice Clone (#2307)
|
||||
- Videos without sound get a clear message instead of an ffmpeg error dump (#2308)
|
||||
@@ -20,7 +20,6 @@ metadata and the backend fallback mirror it. Archived Tauri manifests stay froze
|
||||
### Added
|
||||
|
||||
- Call agent backend: place or answer phone calls that hold a task conversation in your verified or designed voice, with an editable AI disclosure, take-over and an after-call summary (#2306)
|
||||
- Calls workspace with a live transcript, take-over, hang-up and an after-call summary (#2305)
|
||||
|
||||
### Changed
|
||||
|
||||
@@ -46,6 +45,7 @@ metadata and the backend fallback mirror it. Archived Tauri manifests stay froze
|
||||
### Fixed
|
||||
|
||||
- A video with no audio track now says so in Dub, Batch, transcription, cloning and imports instead of showing ffmpeg's exit-234 dump (#2308)
|
||||
- Cache flushes after a dubbing offload or a failed generation reach every accelerator an engine can run on — Ascend NPU and Intel XPU included, not just CUDA and MPS (#2317) — thanks @li-lizhe!
|
||||
- Hardsub exports on Windows pass the caption path in ffmpeg's filter form so burned-in line and karaoke captions render (#2312) — thanks @kevin9327!
|
||||
|
||||
## [0.5.6] — 2026-09-23
|
||||
|
||||
@@ -20,7 +20,7 @@ from core.logging_utils import log_safe
|
||||
from core import event_bus
|
||||
from schemas.requests import DubIngestUrlRequest, ParseSubtitleTextRequest
|
||||
from services.srt_parser import CUE_SOURCE_FIELDS, CUE_SOURCE_ID
|
||||
from services.model_manager import get_model, _gpu_pool, _cpu_pool, get_diarization_pipeline, offload_tts_for_asr, restore_tts_after_asr, should_preload_tts_asr
|
||||
from services.model_manager import get_model, _gpu_pool, _cpu_pool, get_diarization_pipeline, offload_tts_for_asr, restore_tts_after_asr, should_preload_tts_asr, release_device_cache
|
||||
from services.asr_backend import (
|
||||
ASR_TRANSCRIBE_TIMEOUT_S,
|
||||
ASRTimeoutError,
|
||||
@@ -2143,9 +2143,14 @@ async def dub_transcribe_stream(
|
||||
# Debt paid — don't make gen()'s finally repeat it.
|
||||
_tts_offloaded["v"] = False
|
||||
|
||||
if torch.backends.mps.is_available():
|
||||
try: torch.mps.empty_cache()
|
||||
except Exception: pass
|
||||
# The offload dance above is what made room; hand back whatever the
|
||||
# allocator is still holding on the accelerator this host actually
|
||||
# synthesizes on (MPS-only here left CUDA, XPU and Ascend NPU hosts
|
||||
# holding the freed blocks).
|
||||
fut_release = loop.run_in_executor(_gpu_pool, release_device_cache)
|
||||
async for _ping in _ping_while(fut_release):
|
||||
yield _ping
|
||||
fut_release.result()
|
||||
|
||||
yield _sse_event("final", {
|
||||
"segments": final_segs,
|
||||
@@ -2367,8 +2372,10 @@ async def dub_transcribe(job_id: str, num_speakers: Optional[int] = None):
|
||||
s.setdefault("text_original", s.get("text", ""))
|
||||
job["full_transcript"] = " ".join(s["text"] for s in segments)
|
||||
|
||||
if torch.backends.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
# Transcription is done with the resident TTS model still offloaded;
|
||||
# release the accelerator cache the offload freed, on whichever
|
||||
# backend this host synthesizes with.
|
||||
release_device_cache()
|
||||
|
||||
return segments
|
||||
|
||||
|
||||
@@ -79,12 +79,10 @@ def _prepare_oom_retry(error: Exception, *, execution_target: str) -> bool:
|
||||
raise error
|
||||
|
||||
import gc
|
||||
from services.model_manager import release_device_cache
|
||||
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
release_device_cache()
|
||||
return True
|
||||
|
||||
|
||||
@@ -579,6 +577,8 @@ async def dub_generate(job_id: str, req: DubRequest):
|
||||
# serialised the GPU loop; the batched-I/O design it replaced kept
|
||||
# it off the hot path on purpose. Flush every ~16 releases instead —
|
||||
# frequent enough to bound VRAM, rare enough to stay invisible.
|
||||
from services.model_manager import release_device_cache
|
||||
|
||||
_RELEASE_FLUSH_EVERY = 16
|
||||
_release_count = {"n": 0}
|
||||
|
||||
@@ -586,21 +586,17 @@ async def dub_generate(job_id: str, req: DubRequest):
|
||||
"""Best-effort VRAM cleanup after a segment is safely on disk.
|
||||
|
||||
Tensors are freed by the callers' own ``del`` once they fall out
|
||||
of scope; this only throttles the device cache flush. ``*objs`` is
|
||||
kept for call-site compatibility but intentionally unused — a local
|
||||
``del`` here would only unbind the parameter, never the caller's
|
||||
reference.
|
||||
of scope; this only throttles the device cache flush, which has to
|
||||
cover every backend an engine can synthesize on (CUDA, MPS, XPU,
|
||||
Ascend NPU) — the CUDA/MPS pair this used to open-code skipped the
|
||||
rest. ``*objs`` is kept for call-site compatibility but
|
||||
intentionally unused — a local ``del`` here would only unbind the
|
||||
parameter, never the caller's reference.
|
||||
"""
|
||||
_release_count["n"] += 1
|
||||
if _release_count["n"] % _RELEASE_FLUSH_EVERY != 0:
|
||||
return
|
||||
try:
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
except Exception:
|
||||
pass
|
||||
release_device_cache()
|
||||
|
||||
# mix_<id> scratch WAVs written for silence/cached-fail/error slots are
|
||||
# pure assembly inputs (no preview/regen contract), so they're deleted
|
||||
|
||||
@@ -352,18 +352,13 @@ def _unload_nllb():
|
||||
"""Release NLLB VRAM so TTS model can reload."""
|
||||
global _nllb_device, _nllb_model, _nllb_tokenizer
|
||||
import gc
|
||||
from services.model_manager import release_device_cache
|
||||
device = _nllb_device
|
||||
_nllb_model = None
|
||||
_nllb_tokenizer = None
|
||||
_nllb_device = None
|
||||
gc.collect()
|
||||
try:
|
||||
import torch
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
except Exception:
|
||||
pass
|
||||
release_device_cache(device=device)
|
||||
|
||||
|
||||
def _should_unload_nllb() -> bool:
|
||||
@@ -414,6 +409,7 @@ async def dub_translate(req: TranslateRequest):
|
||||
def _translate_nllb():
|
||||
global _nllb_model, _nllb_tokenizer, _nllb_device
|
||||
import torch
|
||||
from services.model_manager import release_device_cache
|
||||
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
|
||||
|
||||
if torch.cuda.is_available():
|
||||
@@ -472,6 +468,7 @@ async def dub_translate(req: TranslateRequest):
|
||||
_nllb_model.to("cpu")
|
||||
_nllb_device = "cpu"
|
||||
inputs = {key: value.to("cpu") for key, value in inputs.items()}
|
||||
release_device_cache(device="mps")
|
||||
tokens = _nllb_model.generate(
|
||||
**inputs,
|
||||
forced_bos_token_id=forced_bos_token_id,
|
||||
@@ -521,9 +518,11 @@ async def dub_translate(req: TranslateRequest):
|
||||
continue
|
||||
# A single unusually long row must not sink its
|
||||
# neighbours. Clear a failed device allocation and
|
||||
# retain the established per-segment degradation.
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
# retain the established per-segment degradation —
|
||||
# on whichever accelerator the translator is on
|
||||
# (the ladder above picks CUDA or MPS, and a
|
||||
# CUDA-only flush missed the MPS case).
|
||||
release_device_cache(device=_nllb_device)
|
||||
logger.warning(
|
||||
"NLLB batch of %d failed; retrying rows individually: %s",
|
||||
len(batch),
|
||||
|
||||
@@ -591,12 +591,10 @@ def _oom_friendly_reraise(e):
|
||||
"""Best-effort cache flush + the user-facing OOM hint shared by both
|
||||
inference paths."""
|
||||
import gc
|
||||
import torch
|
||||
from services.model_manager import release_device_cache
|
||||
|
||||
gc.collect()
|
||||
if torch.backends.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
elif torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
release_device_cache()
|
||||
# #278: don't mislabel a torch.compile/Triton/Inductor crash as an
|
||||
# out-of-memory condition. (model_manager's generate wrapper already
|
||||
# retries these eagerly; this only triggers if that retry also died.)
|
||||
|
||||
@@ -3315,6 +3315,76 @@ def _clear_cublas_workspaces(torch) -> None:
|
||||
logger.debug("clearing cuBLAS workspaces failed", exc_info=True)
|
||||
|
||||
|
||||
def _active_accelerator_name(torch):
|
||||
"""The device type this host runs inference on, when torch can say.
|
||||
|
||||
Engine sidecars resolve their device with
|
||||
``torch.accelerator.current_accelerator(check_available=True)`` (see
|
||||
``engines/moss_tts_v15``), so the flush has to ask the same question — on a
|
||||
hybrid host where CUDA also reports available, the accelerator the engines
|
||||
actually synthesize on is the one whose cache must go back.
|
||||
|
||||
``None`` means this build cannot answer (pre-2.6 has no ``torch.accelerator``,
|
||||
or probing raised on an optional driver); callers then probe the backends
|
||||
the build ships.
|
||||
"""
|
||||
current = getattr(getattr(torch, "accelerator", None), "current_accelerator", None)
|
||||
if current is None:
|
||||
return None
|
||||
try:
|
||||
accel = current(check_available=True)
|
||||
except Exception: # noqa: BLE001 — probing optional drivers can fail
|
||||
return None
|
||||
return getattr(accel, "type", None)
|
||||
|
||||
|
||||
def release_device_cache(*, device: str | None = None, raise_on_failure: bool = False) -> None:
|
||||
"""Release the chosen device's cached blocks, or the active accelerator's.
|
||||
|
||||
The narrow primitive behind ``free_vram()``: no ``gc.collect()`` and no
|
||||
cuBLAS workspace clear, because the callers that need it most — the
|
||||
per-segment release in dubbing, the OOM retry guard — sit on a hot path
|
||||
where a full collection is the wrong price for dropping the allocator's
|
||||
cache.
|
||||
|
||||
By default it follows the accelerator used by engines that select through
|
||||
``torch.accelerator`` (CUDA, MPS, XPU or NPU). Callers that select their
|
||||
own device, such as NLLB, pass that device explicitly so a hybrid host
|
||||
never flushes a different allocator.
|
||||
|
||||
Best-effort by default: freeing memory must not fail the request, so a
|
||||
broken backend is logged and swallowed. ``raise_on_failure`` restores
|
||||
``free_vram()``'s original contract for unload callers that report a failed
|
||||
flush to the user.
|
||||
"""
|
||||
torch = _lazy_torch()
|
||||
try:
|
||||
name = str(device).split(":", 1)[0] if device is not None else _active_accelerator_name(torch)
|
||||
if name == "cpu" and device is not None:
|
||||
return
|
||||
if name and name != "cpu":
|
||||
empty_cache = getattr(getattr(torch, name, None), "empty_cache", None)
|
||||
if empty_cache is not None:
|
||||
empty_cache()
|
||||
return
|
||||
if device is not None:
|
||||
return
|
||||
# No usable answer from torch.accelerator: probe the backends this build
|
||||
# ships, in the order the host prefers them.
|
||||
for name in ("cuda", "mps", "xpu", "npu"):
|
||||
backend = getattr(torch, name, None)
|
||||
is_available = getattr(backend, "is_available", None)
|
||||
empty_cache = getattr(backend, "empty_cache", None)
|
||||
if is_available is None or empty_cache is None or not is_available():
|
||||
continue
|
||||
empty_cache()
|
||||
return
|
||||
except Exception: # noqa: BLE001 — freeing memory must never raise
|
||||
if raise_on_failure:
|
||||
raise
|
||||
logger.debug("releasing the device cache failed")
|
||||
|
||||
|
||||
def free_vram():
|
||||
"""Release cached GPU memory on any accelerator (CUDA, MPS, XPU, NPU)."""
|
||||
torch = _lazy_torch()
|
||||
@@ -3322,13 +3392,7 @@ def free_vram():
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
_clear_cublas_workspaces(torch)
|
||||
torch.cuda.empty_cache()
|
||||
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||
torch.mps.empty_cache()
|
||||
elif hasattr(torch, "xpu") and torch.xpu.is_available():
|
||||
torch.xpu.empty_cache()
|
||||
elif hasattr(torch, "npu") and torch.npu.is_available():
|
||||
torch.npu.empty_cache()
|
||||
release_device_cache(raise_on_failure=True)
|
||||
|
||||
|
||||
def unload_shared_model() -> bool:
|
||||
|
||||
+1
-1
@@ -200,7 +200,7 @@ budget is still abandoned.
|
||||
switched away from it is marked *"not active — safe to unload"*. Below the
|
||||
list are the two bulk actions:
|
||||
- **Flush caches** — runs a multi-pass garbage collection and releases the
|
||||
accelerator's cached memory (CUDA/MPS/XPU `empty_cache`). Models stay
|
||||
accelerator's cached memory (CUDA/MPS/XPU/NPU `empty_cache`). Models stay
|
||||
loaded, so there's no reload cost; this recovers cache/fragmentation
|
||||
memory only.
|
||||
- **Unload all + flush** — the above **plus** fully unloads the resident
|
||||
|
||||
@@ -0,0 +1,173 @@
|
||||
"""Device-cache release reaches every accelerator an engine can select.
|
||||
|
||||
The engine sidecars resolve their device through ``torch.accelerator``, so an
|
||||
Ascend NPU host synthesizes on ``npu`` and an Intel Arc host on ``xpu`` (see
|
||||
``tests/test_moss_tts_v15.py``). The dubbing and generation recovery paths
|
||||
open-coded the CUDA/MPS pair instead, which made the flush a silent no-op on
|
||||
those hosts: the allocator kept the blocks it had just been asked to hand back,
|
||||
and the next allocation failed with the memory still counted as "in use".
|
||||
|
||||
These tests pin the shared primitive to all four backends, to the accelerator
|
||||
the engines actually resolve (``torch.accelerator``, and the probe chain on
|
||||
builds too old to have it), and to the two properties the call sites depend on —
|
||||
the direct recovery calls never raise, and ``free_vram()`` still surfaces a
|
||||
failed flush to its unload callers.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
BACKENDS = ("cuda", "mps", "xpu", "npu")
|
||||
|
||||
|
||||
def _fake_torch(active: str, *, npu_present: bool = True, engine_accelerator: str | None = None):
|
||||
"""A torch stub whose only available accelerator is ``active``.
|
||||
|
||||
``engine_accelerator`` mirrors what ``torch.accelerator.current_accelerator``
|
||||
answers — i.e. the device the engine sidecars would synthesize on. Left
|
||||
``None`` the stub has no ``torch.accelerator`` at all, like a pre-2.6 build.
|
||||
"""
|
||||
backends = {
|
||||
name: SimpleNamespace(
|
||||
is_available=lambda name=name: active == name,
|
||||
empty_cache=Mock(),
|
||||
)
|
||||
for name in BACKENDS
|
||||
}
|
||||
torch = SimpleNamespace(
|
||||
cuda=backends["cuda"],
|
||||
mps=backends["mps"],
|
||||
xpu=backends["xpu"],
|
||||
backends=SimpleNamespace(mps=backends["mps"]),
|
||||
)
|
||||
if npu_present:
|
||||
torch.npu = backends["npu"]
|
||||
if engine_accelerator is not None:
|
||||
accel = SimpleNamespace(type=engine_accelerator)
|
||||
torch.accelerator = SimpleNamespace(current_accelerator=lambda **_kw: accel)
|
||||
return torch, backends
|
||||
|
||||
|
||||
@pytest.mark.parametrize("active", ["cpu", *BACKENDS])
|
||||
def test_release_device_cache_flushes_the_active_accelerator(monkeypatch, active):
|
||||
from services import model_manager as mm
|
||||
|
||||
torch, backends = _fake_torch(active)
|
||||
monkeypatch.setattr(mm, "_lazy_torch", lambda: torch)
|
||||
|
||||
mm.release_device_cache()
|
||||
|
||||
for name, backend in backends.items():
|
||||
assert backend.empty_cache.call_count == int(active == name), name
|
||||
|
||||
|
||||
def test_release_device_cache_skips_a_backend_this_torch_does_not_ship(monkeypatch):
|
||||
"""A torch build without torch_npu has no ``npu`` attribute to probe."""
|
||||
from services import model_manager as mm
|
||||
|
||||
torch, backends = _fake_torch("npu", npu_present=False)
|
||||
monkeypatch.setattr(mm, "_lazy_torch", lambda: torch)
|
||||
|
||||
mm.release_device_cache()
|
||||
|
||||
assert all(backend.empty_cache.call_count == 0 for backend in backends.values())
|
||||
|
||||
|
||||
def test_release_device_cache_follows_the_engines_accelerator(monkeypatch):
|
||||
"""Hybrid host: CUDA probes as available, but the engines run on npu.
|
||||
|
||||
``torch.accelerator.current_accelerator`` is the resolution the engine
|
||||
sidecars use (``engines/moss_tts_v15``), so the flush has to follow it
|
||||
instead of stopping at the first backend that merely probes as available —
|
||||
otherwise the active allocator keeps the blocks and the recovery retry hits
|
||||
the same OOM.
|
||||
"""
|
||||
from services import model_manager as mm
|
||||
|
||||
torch, backends = _fake_torch("cuda", engine_accelerator="npu")
|
||||
monkeypatch.setattr(mm, "_lazy_torch", lambda: torch)
|
||||
|
||||
mm.release_device_cache()
|
||||
|
||||
assert backends["npu"].empty_cache.call_count == 1
|
||||
assert backends["cuda"].empty_cache.call_count == 0
|
||||
|
||||
|
||||
def test_release_device_cache_uses_explicit_caller_device(monkeypatch):
|
||||
"""NLLB may run on CUDA while the default engine accelerator is NPU."""
|
||||
from services import model_manager as mm
|
||||
|
||||
torch, backends = _fake_torch("npu", engine_accelerator="npu")
|
||||
monkeypatch.setattr(mm, "_lazy_torch", lambda: torch)
|
||||
mm.release_device_cache(device="cuda:0")
|
||||
assert backends["cuda"].empty_cache.call_count == 1
|
||||
assert backends["npu"].empty_cache.call_count == 0
|
||||
mm.release_device_cache(device="cpu")
|
||||
assert backends["npu"].empty_cache.call_count == 0
|
||||
|
||||
|
||||
def test_release_device_cache_ignores_a_cpu_answer(monkeypatch):
|
||||
"""A cpu answer must not end the search — probe the shipped backends."""
|
||||
from services import model_manager as mm
|
||||
|
||||
torch, backends = _fake_torch("npu", engine_accelerator="cpu")
|
||||
monkeypatch.setattr(mm, "_lazy_torch", lambda: torch)
|
||||
|
||||
mm.release_device_cache()
|
||||
|
||||
assert backends["npu"].empty_cache.call_count == 1
|
||||
|
||||
|
||||
def test_release_device_cache_falls_back_without_torch_accelerator(monkeypatch):
|
||||
"""A pre-2.6 build has no ``torch.accelerator`` — the probe chain still runs."""
|
||||
from services import model_manager as mm
|
||||
|
||||
torch, backends = _fake_torch("npu") # no accelerator attribute at all
|
||||
monkeypatch.setattr(mm, "_lazy_torch", lambda: torch)
|
||||
|
||||
mm.release_device_cache()
|
||||
|
||||
assert backends["npu"].empty_cache.call_count == 1
|
||||
|
||||
|
||||
def test_release_device_cache_never_raises(monkeypatch):
|
||||
"""Freeing memory is best-effort — a broken backend must not fail a request."""
|
||||
from services import model_manager as mm
|
||||
|
||||
torch, backends = _fake_torch("npu")
|
||||
backends["npu"].empty_cache.side_effect = RuntimeError("driver went away")
|
||||
monkeypatch.setattr(mm, "_lazy_torch", lambda: torch)
|
||||
|
||||
mm.release_device_cache() # must not raise
|
||||
|
||||
|
||||
def test_free_vram_still_propagates_a_flush_failure(monkeypatch):
|
||||
"""``free_vram()`` keeps its original contract: unload callers report a
|
||||
failed flush (``model_lifecycle`` records ``success: False, reason: …``)."""
|
||||
from services import model_manager as mm
|
||||
|
||||
torch, backends = _fake_torch("cuda")
|
||||
backends["cuda"].empty_cache.side_effect = RuntimeError("driver went away")
|
||||
monkeypatch.setattr(mm, "_lazy_torch", lambda: torch)
|
||||
monkeypatch.setattr(mm, "_clear_cublas_workspaces", lambda torch: None)
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
mm.free_vram()
|
||||
|
||||
|
||||
def test_free_vram_still_collects_cublas_and_flushes(monkeypatch):
|
||||
"""``free_vram()`` keeps its gc + cuBLAS clear and now shares the flush."""
|
||||
from services import model_manager as mm
|
||||
|
||||
torch, backends = _fake_torch("cuda")
|
||||
cleared = Mock()
|
||||
monkeypatch.setattr(mm, "_lazy_torch", lambda: torch)
|
||||
monkeypatch.setattr(mm, "_clear_cublas_workspaces", cleared)
|
||||
|
||||
mm.free_vram()
|
||||
|
||||
assert cleared.call_count == 1
|
||||
assert backends["cuda"].empty_cache.call_count == 1
|
||||
@@ -168,6 +168,66 @@ async def test_nllb_batches_segments_by_target_language(monkeypatch):
|
||||
assert response["translated"][2]["text"].startswith("spa_Latn:")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_nllb_mps_fallback_releases_the_old_allocator(monkeypatch):
|
||||
from unittest.mock import Mock
|
||||
|
||||
import torch
|
||||
from api.routers import dub_translate
|
||||
from schemas.requests import TranslateRequest
|
||||
from services import model_manager, translation_engines
|
||||
|
||||
class Tensor:
|
||||
def __init__(self, values):
|
||||
self.values = values
|
||||
|
||||
def to(self, device):
|
||||
return self
|
||||
|
||||
class Tokenizer:
|
||||
def __call__(self, texts, **kwargs):
|
||||
return {"input_ids": Tensor(texts)}
|
||||
|
||||
def convert_tokens_to_ids(self, target):
|
||||
return target
|
||||
|
||||
def batch_decode(self, tokens, **kwargs):
|
||||
return tokens
|
||||
|
||||
class Model:
|
||||
def __init__(self):
|
||||
self.device = "mps"
|
||||
|
||||
def to(self, device):
|
||||
self.device = device
|
||||
return self
|
||||
|
||||
def generate(self, *, input_ids, **kwargs):
|
||||
if self.device == "mps":
|
||||
raise RuntimeError("MPS allocation failed")
|
||||
return [f"translated:{text}" for text in input_ids.values]
|
||||
|
||||
flushed = Mock()
|
||||
monkeypatch.setattr(model_manager, "release_device_cache", flushed)
|
||||
monkeypatch.setattr(dub_translate, "_nllb_tokenizer", Tokenizer())
|
||||
monkeypatch.setattr(dub_translate, "_nllb_model", Model())
|
||||
monkeypatch.setattr(dub_translate, "_nllb_device", "mps")
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
|
||||
monkeypatch.setattr(torch.backends.mps, "is_available", lambda: True)
|
||||
monkeypatch.setattr(translation_engines, "is_installed", lambda _: True)
|
||||
monkeypatch.setattr(translation_engines, "is_ready", lambda _: True)
|
||||
monkeypatch.setenv("OMNIVOICE_UNLOAD_NLLB", "0")
|
||||
|
||||
response = await dub_translate.dub_translate(
|
||||
TranslateRequest(
|
||||
provider="nllb", source_lang="en", target_lang="de",
|
||||
segments=[{"id": "1", "text": "hello"}],
|
||||
)
|
||||
)
|
||||
assert response["translated"][0]["text"] == "translated:hello"
|
||||
assert dub_translate._nllb_device == "cpu"
|
||||
flushed.assert_called_once_with(device="mps")
|
||||
|
||||
def test_resolve_source_lang_priority(monkeypatch):
|
||||
from api.routers import dub_translate
|
||||
|
||||
|
||||
Reference in New Issue
Block a user