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().
This commit is contained in:
jamiepine
2026-10-04 00:25:53 +00:00
committed by capy-ai-staging[bot]
parent b1323ed0a8
commit cffe24ffd2
3 changed files with 49 additions and 3 deletions
+15 -2
View File
@@ -319,7 +319,15 @@ class MLXTTSBackend:
""" """
if self.model is None or self._current_model_size != self.model_size: if self.model is None or self._current_model_size != self.model_size:
self._reload_sync(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: async with self._op_lock:
audio, sample_rate = await _run_on_mlx_thread(_reload_and_generate_sync) 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: if self.model is None or self.model_size != resolved_size:
self._load_model_sync(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: async with self._op_lock:
return await _run_on_mlx_thread(_reload_and_transcribe_sync) return await _run_on_mlx_thread(_reload_and_transcribe_sync)
+6 -1
View File
@@ -282,7 +282,12 @@ class MLXQwenLLMBackend:
# unload can swap or null out self.model / self.tokenizer. # unload can swap or null out self.model / self.tokenizer.
if self.model is None or self._current_model_size != resolved_size: if self.model is None or self._current_model_size != resolved_size:
self._reload_sync(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: async with self._op_lock:
return await _run_on_mlx_thread(_reload_and_generate_sync) return await _run_on_mlx_thread(_reload_and_generate_sync)
@@ -80,3 +80,31 @@ def test_mlx_backends_do_not_clear_cache_when_already_unloaded(monkeypatch):
backend.unload_model() backend.unload_model()
mock_clear.assert_not_called() 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()