diff --git a/backend/backends/__init__.py b/backend/backends/__init__.py index eb436c2b..e3cfe049 100644 --- a/backend/backends/__init__.py +++ b/backend/backends/__init__.py @@ -595,16 +595,16 @@ def unload_model_by_config(config: ModelConfig) -> bool: backend = get_tts_backend_for_engine(config.engine) 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() + backend.unload_model() return True return False # All other TTS engines backend = get_tts_backend_for_engine(config.engine) if backend.is_loaded(): - backend.unload_model() clear_voice_prompt_memory_cache() + backend.unload_model() return True return False diff --git a/backend/backends/base.py b/backend/backends/base.py index 30b3a06c..a00b1604 100644 --- a/backend/backends/base.py +++ b/backend/backends/base.py @@ -232,6 +232,12 @@ def empty_mlx_cache() -> None: 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. + + Safe from any thread: ``mx.clear_cache`` only drains the global + allocator pool and never touches the per-thread stream registry, so + unlike load/generate it does not have to run on the MLX worker thread + (verified from the FastAPI event loop with a generation in flight on + the worker). """ import mlx.core as mx diff --git a/backend/services/tts.py b/backend/services/tts.py index 9e7505f3..76b23971 100644 --- a/backend/services/tts.py +++ b/backend/services/tts.py @@ -24,8 +24,8 @@ def get_tts_model() -> TTSBackend: def unload_tts_model(): """Unload TTS model to free memory.""" backend = get_tts_backend() - backend.unload_model() clear_voice_prompt_memory_cache() + backend.unload_model() def audio_to_wav_bytes(audio: np.ndarray, sample_rate: int) -> bytes: