diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index 293da323..67c5e56f 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -319,15 +319,16 @@ class MLXTTSBackend: """ if self.model is None or self._current_model_size != self.model_size: self._reload_sync(self.model_size) - result = _generate_sync() - # An unload_model() that landed while we were generating only - # dropped the backend's reference; the model's buffers were kept - # alive by _generate_sync's local and have just been returned to - # MLX's pool. Drain it now, or they stay resident until the next - # load/unload cycle. - if self.model is None: - empty_mlx_cache() - return result + try: + return _generate_sync() + finally: + # An unload_model() that landed while we were generating only + # dropped the backend's reference; the model's buffers were + # kept alive by _generate_sync's local and have just been + # returned to MLX's pool (on success or failure). Drain it now, + # or they stay resident until the next load/unload cycle. + if self.model is None: + empty_mlx_cache() async with self._op_lock: audio, sample_rate = await _run_on_mlx_thread(_reload_and_generate_sync) @@ -449,12 +450,13 @@ class MLXSTTBackend: """ if self.model is None or self.model_size != resolved_size: self._load_model_sync(resolved_size) - result = _transcribe_sync() - # See MLXTTSBackend._reload_and_generate_sync: drain the pool if an - # unload landed mid-transcription. - if self.model is None: - empty_mlx_cache() - return result + try: + return _transcribe_sync() + finally: + # See MLXTTSBackend._reload_and_generate_sync: drain the pool + # if an unload landed mid-transcription. + if self.model is None: + empty_mlx_cache() async with self._op_lock: return await _run_on_mlx_thread(_reload_and_transcribe_sync) diff --git a/backend/backends/qwen_llm_backend.py b/backend/backends/qwen_llm_backend.py index debf078c..81ca7a9f 100644 --- a/backend/backends/qwen_llm_backend.py +++ b/backend/backends/qwen_llm_backend.py @@ -282,12 +282,13 @@ class MLXQwenLLMBackend: # unload can swap or null out self.model / self.tokenizer. if self.model is None or self._current_model_size != resolved_size: self._reload_sync(resolved_size) - result = self._generate_sync(prompt, system, max_tokens, temperature, examples) - # Drain the MLX pool if an unload landed mid-generation (see - # MLXTTSBackend._reload_and_generate_sync). - if self.model is None: - empty_mlx_cache() - return result + try: + return self._generate_sync(prompt, system, max_tokens, temperature, examples) + finally: + # Drain the MLX pool if an unload landed mid-generation (see + # MLXTTSBackend._reload_and_generate_sync). + if self.model is None: + empty_mlx_cache() async with self._op_lock: return await _run_on_mlx_thread(_reload_and_generate_sync) diff --git a/backend/tests/test_mlx_unload_clears_cache.py b/backend/tests/test_mlx_unload_clears_cache.py index a311a98b..4c9ba7f7 100644 --- a/backend/tests/test_mlx_unload_clears_cache.py +++ b/backend/tests/test_mlx_unload_clears_cache.py @@ -107,4 +107,30 @@ async def test_mlx_tts_generate_drains_cache_when_unloaded_mid_generation(monkey assert len(audio) == 2 assert backend.model is None - mock_clear.assert_called_once() + # Once from unload_model() itself (reference drop) and once more after the + # in-flight generation released the model it had bound locally. + assert mock_clear.call_count == 2 + + +@pytest.mark.asyncio +async def test_mlx_tts_generate_drains_cache_when_unloaded_and_generation_fails(monkeypatch): + """The post-generation drain also runs when the in-flight op raises.""" + mock_clear = MagicMock() + monkeypatch.setattr(mlx_backend, "empty_mlx_cache", mock_clear) + + backend = mlx_backend.MLXTTSBackend(model_size="0.6B") + + class _FakeModel: + def generate(self, text, **kwargs): + backend.unload_model() + raise RuntimeError("boom") + yield # pragma: no cover - makes this a generator + + backend.model = _FakeModel() + backend._current_model_size = "0.6B" + + with pytest.raises(RuntimeError, match="boom"): + await backend.generate("hello", {}, "en") + + assert backend.model is None + assert mock_clear.call_count == 2