mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -07:00
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.
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
cffe24ffd2
commit
fb67ad436c
@@ -319,15 +319,16 @@ 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)
|
||||||
result = _generate_sync()
|
try:
|
||||||
# An unload_model() that landed while we were generating only
|
return _generate_sync()
|
||||||
# dropped the backend's reference; the model's buffers were kept
|
finally:
|
||||||
# alive by _generate_sync's local and have just been returned to
|
# An unload_model() that landed while we were generating only
|
||||||
# MLX's pool. Drain it now, or they stay resident until the next
|
# dropped the backend's reference; the model's buffers were
|
||||||
# load/unload cycle.
|
# kept alive by _generate_sync's local and have just been
|
||||||
if self.model is None:
|
# returned to MLX's pool (on success or failure). Drain it now,
|
||||||
empty_mlx_cache()
|
# or they stay resident until the next load/unload cycle.
|
||||||
return result
|
if self.model is None:
|
||||||
|
empty_mlx_cache()
|
||||||
|
|
||||||
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)
|
||||||
@@ -449,12 +450,13 @@ 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)
|
||||||
result = _transcribe_sync()
|
try:
|
||||||
# See MLXTTSBackend._reload_and_generate_sync: drain the pool if an
|
return _transcribe_sync()
|
||||||
# unload landed mid-transcription.
|
finally:
|
||||||
if self.model is None:
|
# See MLXTTSBackend._reload_and_generate_sync: drain the pool
|
||||||
empty_mlx_cache()
|
# if an unload landed mid-transcription.
|
||||||
return result
|
if self.model is None:
|
||||||
|
empty_mlx_cache()
|
||||||
|
|
||||||
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)
|
||||||
|
|||||||
@@ -282,12 +282,13 @@ 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)
|
||||||
result = self._generate_sync(prompt, system, max_tokens, temperature, examples)
|
try:
|
||||||
# Drain the MLX pool if an unload landed mid-generation (see
|
return self._generate_sync(prompt, system, max_tokens, temperature, examples)
|
||||||
# MLXTTSBackend._reload_and_generate_sync).
|
finally:
|
||||||
if self.model is None:
|
# Drain the MLX pool if an unload landed mid-generation (see
|
||||||
empty_mlx_cache()
|
# MLXTTSBackend._reload_and_generate_sync).
|
||||||
return result
|
if self.model is None:
|
||||||
|
empty_mlx_cache()
|
||||||
|
|
||||||
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)
|
||||||
|
|||||||
@@ -107,4 +107,30 @@ async def test_mlx_tts_generate_drains_cache_when_unloaded_mid_generation(monkey
|
|||||||
|
|
||||||
assert len(audio) == 2
|
assert len(audio) == 2
|
||||||
assert backend.model is None
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user