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:
jamiepine
2026-10-04 00:25:53 +00:00
committed by capy-ai-staging[bot]
parent cffe24ffd2
commit fb67ad436c
3 changed files with 51 additions and 22 deletions
+17 -15
View File
@@ -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)
+7 -6
View File
@@ -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)