mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -07:00
fix(backend): release memory when unloading MLX models
Unloading a TTS/Whisper/LLM model on the MLX backend only dropped the
Python reference (`del self.model`). MLX keeps freed array buffers in
its own allocator pool for reuse instead of returning them to the OS,
so the process's memory footprint never actually shrank after unload
on Apple Silicon (the default backend there) until the process exited.
Add empty_mlx_cache() (backend/backends/base.py), wrapping
mx.clear_cache(), and call it from the three MLX unload_model()
implementations: MLXTTSBackend, MLXSTTBackend, MLXQwenLLMBackend.
Separately, the voice-clone prompt cache (backend/utils/cache.py) is a
process-lifetime dict populated by create_voice_prompt() across every
TTS engine, but nothing ever cleared it on model unload — only the
unrelated /tasks/clear-cache endpoint touched it. Add
clear_voice_prompt_memory_cache() (memory only, disk cache untouched
so a later generation still reloads the prompt instead of recomputing
it) and wire it into every TTS unload path (services/tts.py and the
qwen_custom_voice / generic branches of unload_model_by_config).
Whisper and the LLM backends never produce voice prompts, so their
unload paths are left alone.
Testing:
- New unit tests: backend/tests/test_mlx_unload_clears_cache.py,
backend/tests/test_voice_prompt_cache_unload.py (8 tests, all pass).
- Verified end-to-end on Apple Silicon against real cached models
(Qwen TTS 1.7B, Whisper Turbo, Qwen3 0.6B): loaded each via the
running app, unloaded via the real /models/{name}/unload endpoint,
and confirmed via mx.get_cache_memory()/get_active_memory() that the
MLX allocator's cache drops to 0 on every cycle. Ran a real
voice-clone generation end to end and confirmed the in-memory prompt
cache goes from 1 entry to 0 on unload while the on-disk .prompt
file is left intact.
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
5803aaaa91
commit
19f8f51408
@@ -566,6 +566,7 @@ def unload_model_by_config(config: ModelConfig) -> bool:
|
|||||||
"""Unload a model given its config. Returns True if it was loaded, False otherwise."""
|
"""Unload a model given its config. Returns True if it was loaded, False otherwise."""
|
||||||
from . import get_tts_backend_for_engine
|
from . import get_tts_backend_for_engine
|
||||||
from ..services import tts, transcribe, llm as llm_service
|
from ..services import tts, transcribe, llm as llm_service
|
||||||
|
from ..utils.cache import clear_voice_prompt_memory_cache
|
||||||
|
|
||||||
if config.engine == "whisper":
|
if config.engine == "whisper":
|
||||||
whisper_model = transcribe.get_whisper_model()
|
whisper_model = transcribe.get_whisper_model()
|
||||||
@@ -595,6 +596,7 @@ def unload_model_by_config(config: ModelConfig) -> bool:
|
|||||||
loaded_size = getattr(backend, "_current_model_size", None) or getattr(backend, "model_size", None)
|
loaded_size = getattr(backend, "_current_model_size", None) or getattr(backend, "model_size", None)
|
||||||
if backend.is_loaded() and loaded_size == config.model_size:
|
if backend.is_loaded() and loaded_size == config.model_size:
|
||||||
backend.unload_model()
|
backend.unload_model()
|
||||||
|
clear_voice_prompt_memory_cache()
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -602,6 +604,7 @@ def unload_model_by_config(config: ModelConfig) -> bool:
|
|||||||
backend = get_tts_backend_for_engine(config.engine)
|
backend = get_tts_backend_for_engine(config.engine)
|
||||||
if backend.is_loaded():
|
if backend.is_loaded():
|
||||||
backend.unload_model()
|
backend.unload_model()
|
||||||
|
clear_voice_prompt_memory_cache()
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|||||||
@@ -224,6 +224,20 @@ def empty_device_cache(device: str) -> None:
|
|||||||
torch.mps.empty_cache()
|
torch.mps.empty_cache()
|
||||||
|
|
||||||
|
|
||||||
|
def empty_mlx_cache() -> None:
|
||||||
|
"""
|
||||||
|
Free cached memory in the MLX allocator.
|
||||||
|
|
||||||
|
MLX keeps freed array buffers in an internal pool for reuse instead of
|
||||||
|
returning them to the OS. Backends must call this after unloading an
|
||||||
|
MLX model, or the process's memory footprint never shrinks even though
|
||||||
|
the model object itself was dropped.
|
||||||
|
"""
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
mx.clear_cache()
|
||||||
|
|
||||||
|
|
||||||
def manual_seed(seed: int, device: str) -> None:
|
def manual_seed(seed: int, device: str) -> None:
|
||||||
"""
|
"""
|
||||||
Set the random seed on both CPU and the active accelerator.
|
Set the random seed on both CPU and the active accelerator.
|
||||||
|
|||||||
@@ -33,7 +33,12 @@ patch_huggingface_hub_offline()
|
|||||||
ensure_original_qwen_config_cached()
|
ensure_original_qwen_config_cached()
|
||||||
|
|
||||||
from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS
|
from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS
|
||||||
from .base import is_model_cached, combine_voice_prompts as _combine_voice_prompts, model_load_progress
|
from .base import (
|
||||||
|
is_model_cached,
|
||||||
|
combine_voice_prompts as _combine_voice_prompts,
|
||||||
|
model_load_progress,
|
||||||
|
empty_mlx_cache,
|
||||||
|
)
|
||||||
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
|
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
|
||||||
|
|
||||||
|
|
||||||
@@ -149,6 +154,7 @@ class MLXTTSBackend:
|
|||||||
del self.model
|
del self.model
|
||||||
self.model = None
|
self.model = None
|
||||||
self._current_model_size = None
|
self._current_model_size = None
|
||||||
|
empty_mlx_cache()
|
||||||
logger.info("MLX TTS model unloaded")
|
logger.info("MLX TTS model unloaded")
|
||||||
|
|
||||||
async def create_voice_prompt(
|
async def create_voice_prompt(
|
||||||
@@ -381,6 +387,7 @@ class MLXSTTBackend:
|
|||||||
if self.model is not None:
|
if self.model is not None:
|
||||||
del self.model
|
del self.model
|
||||||
self.model = None
|
self.model = None
|
||||||
|
empty_mlx_cache()
|
||||||
logger.info("MLX Whisper model unloaded")
|
logger.info("MLX Whisper model unloaded")
|
||||||
|
|
||||||
async def transcribe(
|
async def transcribe(
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ from .base import (
|
|||||||
is_model_cached,
|
is_model_cached,
|
||||||
get_torch_device,
|
get_torch_device,
|
||||||
empty_device_cache,
|
empty_device_cache,
|
||||||
|
empty_mlx_cache,
|
||||||
manual_seed,
|
manual_seed,
|
||||||
model_load_progress,
|
model_load_progress,
|
||||||
)
|
)
|
||||||
@@ -258,6 +259,7 @@ class MLXQwenLLMBackend:
|
|||||||
self.model = None
|
self.model = None
|
||||||
self.tokenizer = None
|
self.tokenizer = None
|
||||||
self._current_model_size = None
|
self._current_model_size = None
|
||||||
|
empty_mlx_cache()
|
||||||
logger.info("Qwen3 (MLX) unloaded")
|
logger.info("Qwen3 (MLX) unloaded")
|
||||||
|
|
||||||
async def generate(
|
async def generate(
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import io
|
|||||||
import soundfile as sf
|
import soundfile as sf
|
||||||
|
|
||||||
from ..backends import get_tts_backend, TTSBackend
|
from ..backends import get_tts_backend, TTSBackend
|
||||||
|
from ..utils.cache import clear_voice_prompt_memory_cache
|
||||||
|
|
||||||
|
|
||||||
def get_tts_model() -> TTSBackend:
|
def get_tts_model() -> TTSBackend:
|
||||||
@@ -24,6 +25,7 @@ def unload_tts_model():
|
|||||||
"""Unload TTS model to free memory."""
|
"""Unload TTS model to free memory."""
|
||||||
backend = get_tts_backend()
|
backend = get_tts_backend()
|
||||||
backend.unload_model()
|
backend.unload_model()
|
||||||
|
clear_voice_prompt_memory_cache()
|
||||||
|
|
||||||
|
|
||||||
def audio_to_wav_bytes(audio: np.ndarray, sample_rate: int) -> bytes:
|
def audio_to_wav_bytes(audio: np.ndarray, sample_rate: int) -> bytes:
|
||||||
|
|||||||
@@ -0,0 +1,82 @@
|
|||||||
|
"""
|
||||||
|
Regression tests: unloading an MLX-backed model must release MLX's internal
|
||||||
|
buffer cache, not just drop the Python reference.
|
||||||
|
|
||||||
|
MLX keeps freed array buffers in an internal pool for reuse instead of
|
||||||
|
returning them to the OS (see `mlx.core.clear_cache` / `get_cache_memory`).
|
||||||
|
Before this fix, `unload_model()` on the MLX TTS/Whisper/LLM backends only
|
||||||
|
did `del self.model; self.model = None`, so the process's memory footprint
|
||||||
|
never shrank after "unload" (issue: resources stay held after first use on
|
||||||
|
Apple Silicon). These tests assert each backend's `unload_model()` calls the
|
||||||
|
shared `empty_mlx_cache()` helper, without requiring the real `mlx` package
|
||||||
|
to be installed — `empty_mlx_cache` is monkeypatched, so its own lazy
|
||||||
|
`import mlx.core` never executes here.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
pytest.importorskip("torch")
|
||||||
|
|
||||||
|
from backend.backends import mlx_backend, qwen_llm_backend
|
||||||
|
|
||||||
|
|
||||||
|
def test_mlx_tts_backend_unload_clears_mlx_cache(monkeypatch):
|
||||||
|
"""Unloading the MLX TTS backend must call empty_mlx_cache()."""
|
||||||
|
mock_clear = MagicMock()
|
||||||
|
monkeypatch.setattr(mlx_backend, "empty_mlx_cache", mock_clear)
|
||||||
|
|
||||||
|
backend = mlx_backend.MLXTTSBackend()
|
||||||
|
backend.model = MagicMock()
|
||||||
|
backend._current_model_size = "1.7B"
|
||||||
|
|
||||||
|
backend.unload_model()
|
||||||
|
|
||||||
|
assert backend.model is None
|
||||||
|
assert backend._current_model_size is None
|
||||||
|
mock_clear.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_mlx_stt_backend_unload_clears_mlx_cache(monkeypatch):
|
||||||
|
"""Unloading the MLX Whisper backend must call empty_mlx_cache()."""
|
||||||
|
mock_clear = MagicMock()
|
||||||
|
monkeypatch.setattr(mlx_backend, "empty_mlx_cache", mock_clear)
|
||||||
|
|
||||||
|
backend = mlx_backend.MLXSTTBackend()
|
||||||
|
backend.model = MagicMock()
|
||||||
|
|
||||||
|
backend.unload_model()
|
||||||
|
|
||||||
|
assert backend.model is None
|
||||||
|
mock_clear.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_mlx_llm_backend_unload_clears_mlx_cache(monkeypatch):
|
||||||
|
"""Unloading the MLX Qwen3 LLM backend must call empty_mlx_cache()."""
|
||||||
|
mock_clear = MagicMock()
|
||||||
|
monkeypatch.setattr(qwen_llm_backend, "empty_mlx_cache", mock_clear)
|
||||||
|
|
||||||
|
backend = qwen_llm_backend.MLXQwenLLMBackend()
|
||||||
|
backend.model = MagicMock()
|
||||||
|
backend.tokenizer = MagicMock()
|
||||||
|
backend._current_model_size = "4B"
|
||||||
|
|
||||||
|
backend.unload_model()
|
||||||
|
|
||||||
|
assert backend.model is None
|
||||||
|
assert backend.tokenizer is None
|
||||||
|
mock_clear.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_mlx_backends_do_not_clear_cache_when_already_unloaded(monkeypatch):
|
||||||
|
"""Calling unload on an already-unloaded backend is a no-op (no spurious clear)."""
|
||||||
|
mock_clear = MagicMock()
|
||||||
|
monkeypatch.setattr(mlx_backend, "empty_mlx_cache", mock_clear)
|
||||||
|
|
||||||
|
backend = mlx_backend.MLXTTSBackend()
|
||||||
|
assert backend.model is None
|
||||||
|
|
||||||
|
backend.unload_model()
|
||||||
|
|
||||||
|
mock_clear.assert_not_called()
|
||||||
@@ -0,0 +1,91 @@
|
|||||||
|
"""
|
||||||
|
Regression tests: unloading a TTS model must also drop the in-memory
|
||||||
|
voice-prompt cache.
|
||||||
|
|
||||||
|
`backend/utils/cache.py` keeps a process-lifetime `_memory_cache` dict of
|
||||||
|
voice-clone prompts (tensors or device-backed dicts produced by whichever
|
||||||
|
TTS model created them). Before this fix, nothing ever cleared it on
|
||||||
|
unload, so those prompts stayed referenced — and their memory held —
|
||||||
|
indefinitely, even after the model that produced them was gone. Whisper
|
||||||
|
and the LLM backends never produce voice prompts, so their unload paths
|
||||||
|
must leave the cache alone.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from backend import backends as backends_module
|
||||||
|
from backend.backends import ModelConfig
|
||||||
|
from backend.services import tts as tts_service
|
||||||
|
from backend.utils import cache as cache_module
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _reset_memory_cache():
|
||||||
|
"""Isolate each test from the shared process-lifetime _memory_cache dict."""
|
||||||
|
cache_module._memory_cache.clear()
|
||||||
|
yield
|
||||||
|
cache_module._memory_cache.clear()
|
||||||
|
|
||||||
|
|
||||||
|
def _populate_memory_cache():
|
||||||
|
cache_module._memory_cache["some-cache-key"] = {"ref_audio": "x.wav", "ref_text": "hi"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_clear_voice_prompt_memory_cache_leaves_disk_cache_alone(tmp_path, monkeypatch):
|
||||||
|
"""Clearing the memory cache must not touch cached .prompt files on disk."""
|
||||||
|
monkeypatch.setattr(cache_module, "_get_cache_dir", lambda: tmp_path)
|
||||||
|
disk_file = tmp_path / "some-cache-key.prompt"
|
||||||
|
disk_file.write_bytes(b"fake torch.save payload")
|
||||||
|
_populate_memory_cache()
|
||||||
|
|
||||||
|
cache_module.clear_voice_prompt_memory_cache()
|
||||||
|
|
||||||
|
assert cache_module._memory_cache == {}
|
||||||
|
assert disk_file.exists()
|
||||||
|
|
||||||
|
|
||||||
|
def test_unload_tts_model_clears_voice_prompt_memory_cache(monkeypatch):
|
||||||
|
"""/models/unload (the legacy qwen-only endpoint) must clear the prompt cache."""
|
||||||
|
fake_backend = MagicMock()
|
||||||
|
monkeypatch.setattr(tts_service, "get_tts_backend", lambda: fake_backend)
|
||||||
|
_populate_memory_cache()
|
||||||
|
|
||||||
|
tts_service.unload_tts_model()
|
||||||
|
|
||||||
|
fake_backend.unload_model.assert_called_once()
|
||||||
|
assert cache_module._memory_cache == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_unload_model_by_config_clears_cache_for_generic_tts_engine(monkeypatch):
|
||||||
|
"""Unloading any non-qwen TTS engine (e.g. kokoro) must clear the prompt cache."""
|
||||||
|
fake_backend = MagicMock()
|
||||||
|
fake_backend.is_loaded.return_value = True
|
||||||
|
monkeypatch.setattr(backends_module, "get_tts_backend_for_engine", lambda engine: fake_backend)
|
||||||
|
_populate_memory_cache()
|
||||||
|
|
||||||
|
config = ModelConfig(model_name="kokoro", display_name="Kokoro", engine="kokoro", hf_repo_id="x/y")
|
||||||
|
was_loaded = backends_module.unload_model_by_config(config)
|
||||||
|
|
||||||
|
assert was_loaded is True
|
||||||
|
fake_backend.unload_model.assert_called_once()
|
||||||
|
assert cache_module._memory_cache == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_unload_model_by_config_leaves_cache_alone_for_whisper(monkeypatch):
|
||||||
|
"""Whisper never produces voice prompts, so unloading it must not touch the cache."""
|
||||||
|
fake_whisper = MagicMock()
|
||||||
|
fake_whisper.is_loaded.return_value = True
|
||||||
|
fake_whisper.model_size = "base"
|
||||||
|
monkeypatch.setattr("backend.services.transcribe.get_whisper_model", lambda: fake_whisper)
|
||||||
|
_populate_memory_cache()
|
||||||
|
cache_before = dict(cache_module._memory_cache)
|
||||||
|
|
||||||
|
config = ModelConfig(
|
||||||
|
model_name="whisper-base", display_name="Whisper Base", engine="whisper", hf_repo_id="x/y", model_size="base"
|
||||||
|
)
|
||||||
|
was_loaded = backends_module.unload_model_by_config(config)
|
||||||
|
|
||||||
|
assert was_loaded is True
|
||||||
|
assert cache_module._memory_cache == cache_before
|
||||||
@@ -93,6 +93,19 @@ def cache_voice_prompt(
|
|||||||
torch.save(voice_prompt, cache_file)
|
torch.save(voice_prompt, cache_file)
|
||||||
|
|
||||||
|
|
||||||
|
def clear_voice_prompt_memory_cache() -> None:
|
||||||
|
"""
|
||||||
|
Drop the in-memory voice prompt cache without touching the disk cache.
|
||||||
|
|
||||||
|
Backends call this when a TTS model unloads: the cached prompts (tensors
|
||||||
|
or device-backed dicts produced by that model) would otherwise keep
|
||||||
|
referencing memory forever, since nothing else ever clears this
|
||||||
|
process-lifetime dict. The disk cache is left alone, so the next
|
||||||
|
generation just reloads the prompt from disk instead of recomputing it.
|
||||||
|
"""
|
||||||
|
_memory_cache.clear()
|
||||||
|
|
||||||
|
|
||||||
def clear_voice_prompt_cache() -> int:
|
def clear_voice_prompt_cache() -> int:
|
||||||
"""
|
"""
|
||||||
Clear all voice prompt caches (memory and disk).
|
Clear all voice prompt caches (memory and disk).
|
||||||
|
|||||||
Reference in New Issue
Block a user