mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
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:
committed by
capy-ai-staging[bot]
parent
b1323ed0a8
commit
cffe24ffd2
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user