From fb67ad436c2da000af1f6e19ac4ce74fcc93ced7 Mon Sep 17 00:00:00 2001 From: jamiepine <32987599+jamiepine@users.noreply.github.com> Date: Sun, 4 Oct 2026 00:13:36 +0000 Subject: [PATCH] fix(mlx): drain the pool in a finally so a failing in-flight op releases it too; fix the mid-generation test's call count Review follow-ups: the post-op drain now runs in a finally block for the TTS, STT and LLM folded callables, so an exception escaping the generation (e.g. the voice-clone fallback failing) still releases the buffers the worker held. The mid-generation test asserted one drain call but the path legitimately produces two (unload_model's own and the post-op one); a second test covers the error path. --- backend/backends/mlx_backend.py | 32 ++++++++++--------- backend/backends/qwen_llm_backend.py | 13 ++++---- backend/tests/test_mlx_unload_clears_cache.py | 28 +++++++++++++++- 3 files changed, 51 insertions(+), 22 deletions(-) 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