diff --git a/backend/backends/__init__.py b/backend/backends/__init__.py index 142c430c..eb436c2b 100644 --- a/backend/backends/__init__.py +++ b/backend/backends/__init__.py @@ -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.""" from . import get_tts_backend_for_engine from ..services import tts, transcribe, llm as llm_service + from ..utils.cache import clear_voice_prompt_memory_cache if config.engine == "whisper": 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) if backend.is_loaded() and loaded_size == config.model_size: backend.unload_model() + clear_voice_prompt_memory_cache() return True return False @@ -602,6 +604,7 @@ def unload_model_by_config(config: ModelConfig) -> bool: backend = get_tts_backend_for_engine(config.engine) if backend.is_loaded(): backend.unload_model() + clear_voice_prompt_memory_cache() return True return False diff --git a/backend/backends/base.py b/backend/backends/base.py index a8dcb442..30b3a06c 100644 --- a/backend/backends/base.py +++ b/backend/backends/base.py @@ -224,6 +224,20 @@ def empty_device_cache(device: str) -> None: 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: """ Set the random seed on both CPU and the active accelerator. diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index ce3cdc31..84b9b620 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -33,7 +33,12 @@ patch_huggingface_hub_offline() ensure_original_qwen_config_cached() 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 @@ -149,6 +154,7 @@ class MLXTTSBackend: del self.model self.model = None self._current_model_size = None + empty_mlx_cache() logger.info("MLX TTS model unloaded") async def create_voice_prompt( @@ -381,6 +387,7 @@ class MLXSTTBackend: if self.model is not None: del self.model self.model = None + empty_mlx_cache() logger.info("MLX Whisper model unloaded") async def transcribe( diff --git a/backend/backends/qwen_llm_backend.py b/backend/backends/qwen_llm_backend.py index a7b99632..11b8ca62 100644 --- a/backend/backends/qwen_llm_backend.py +++ b/backend/backends/qwen_llm_backend.py @@ -16,6 +16,7 @@ from .base import ( is_model_cached, get_torch_device, empty_device_cache, + empty_mlx_cache, manual_seed, model_load_progress, ) @@ -258,6 +259,7 @@ class MLXQwenLLMBackend: self.model = None self.tokenizer = None self._current_model_size = None + empty_mlx_cache() logger.info("Qwen3 (MLX) unloaded") async def generate( diff --git a/backend/services/tts.py b/backend/services/tts.py index d4f90ff3..9e7505f3 100644 --- a/backend/services/tts.py +++ b/backend/services/tts.py @@ -8,6 +8,7 @@ import io import soundfile as sf from ..backends import get_tts_backend, TTSBackend +from ..utils.cache import clear_voice_prompt_memory_cache def get_tts_model() -> TTSBackend: @@ -24,6 +25,7 @@ def unload_tts_model(): """Unload TTS model to free memory.""" backend = get_tts_backend() backend.unload_model() + clear_voice_prompt_memory_cache() def audio_to_wav_bytes(audio: np.ndarray, sample_rate: int) -> bytes: diff --git a/backend/tests/test_mlx_unload_clears_cache.py b/backend/tests/test_mlx_unload_clears_cache.py new file mode 100644 index 00000000..05075485 --- /dev/null +++ b/backend/tests/test_mlx_unload_clears_cache.py @@ -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() diff --git a/backend/tests/test_voice_prompt_cache_unload.py b/backend/tests/test_voice_prompt_cache_unload.py new file mode 100644 index 00000000..ab09a013 --- /dev/null +++ b/backend/tests/test_voice_prompt_cache_unload.py @@ -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 diff --git a/backend/utils/cache.py b/backend/utils/cache.py index dd4b9f83..4c9a0733 100644 --- a/backend/utils/cache.py +++ b/backend/utils/cache.py @@ -93,6 +93,19 @@ def cache_voice_prompt( 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: """ Clear all voice prompt caches (memory and disk).