From cffe24ffd203ea404965d3e54f55b03f47e8dafb Mon Sep 17 00:00:00 2001 From: jamiepine <32987599+jamiepine@users.noreply.github.com> Date: Sun, 4 Oct 2026 00:09:48 +0000 Subject: [PATCH] fix(mlx): drain the MLX pool after an in-flight op when unload landed mid-generation Review follow-up: unload_model() runs inline on the event loop, so when it lands during a generation it only drops the backend's reference; the worker's local keeps the model alive and mx.clear_cache() finds nothing to free. Once the folded load-and-generate (or transcribe) callable finishes and releases that local, check whether the model was unloaded meanwhile and drain the pool then. Covered by a test that unloads from inside a fake model.generate(). --- backend/backends/mlx_backend.py | 17 +++++++++-- backend/backends/qwen_llm_backend.py | 7 ++++- backend/tests/test_mlx_unload_clears_cache.py | 28 +++++++++++++++++++ 3 files changed, 49 insertions(+), 3 deletions(-) diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index ad287bca..293da323 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -319,7 +319,15 @@ class MLXTTSBackend: """ if self.model is None or self._current_model_size != self.model_size: self._reload_sync(self.model_size) - return _generate_sync() + 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 async with self._op_lock: audio, sample_rate = await _run_on_mlx_thread(_reload_and_generate_sync) @@ -441,7 +449,12 @@ class MLXSTTBackend: """ if self.model is None or self.model_size != resolved_size: self._load_model_sync(resolved_size) - return _transcribe_sync() + 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 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 11b8ca62..debf078c 100644 --- a/backend/backends/qwen_llm_backend.py +++ b/backend/backends/qwen_llm_backend.py @@ -282,7 +282,12 @@ 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) - return self._generate_sync(prompt, system, max_tokens, temperature, examples) + 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 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 05075485..a311a98b 100644 --- a/backend/tests/test_mlx_unload_clears_cache.py +++ b/backend/tests/test_mlx_unload_clears_cache.py @@ -80,3 +80,31 @@ def test_mlx_backends_do_not_clear_cache_when_already_unloaded(monkeypatch): backend.unload_model() mock_clear.assert_not_called() + + +@pytest.mark.asyncio +async def test_mlx_tts_generate_drains_cache_when_unloaded_mid_generation(monkeypatch): + """An unload that lands while generate() runs must still release the pool afterwards.""" + mock_clear = MagicMock() + monkeypatch.setattr(mlx_backend, "empty_mlx_cache", mock_clear) + + backend = mlx_backend.MLXTTSBackend(model_size="0.6B") + + class _FakeResult: + audio = [0.0, 0.0] + sample_rate = 24000 + + class _FakeModel: + def generate(self, text, **kwargs): + # Simulate /models/unload arriving from the event loop mid-generation. + backend.unload_model() + yield _FakeResult() + + backend.model = _FakeModel() + backend._current_model_size = "0.6B" + + audio, sample_rate = await backend.generate("hello", {}, "en") + + assert len(audio) == 2 + assert backend.model is None + mock_clear.assert_called_once()