mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -07:00
fix(backend): drop the prompt cache before the allocator-empty; note clear_cache is thread-agnostic
Review follow-ups: clear_voice_prompt_memory_cache() now runs before backend.unload_model() at all three call sites so the device-resident prompt tensors are already released when empty_device_cache() / empty_mlx_cache() runs, instead of going back into the caching allocator afterwards. empty_mlx_cache's docstring states why it may run off the MLX worker thread.
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
1a803aa05f
commit
b1323ed0a8
@@ -595,16 +595,16 @@ def unload_model_by_config(config: ModelConfig) -> bool:
|
|||||||
backend = get_tts_backend_for_engine(config.engine)
|
backend = get_tts_backend_for_engine(config.engine)
|
||||||
loaded_size = getattr(backend, "_current_model_size", None) or getattr(backend, "model_size", None)
|
loaded_size = getattr(backend, "_current_model_size", None) or getattr(backend, "model_size", None)
|
||||||
if backend.is_loaded() and loaded_size == config.model_size:
|
if backend.is_loaded() and loaded_size == config.model_size:
|
||||||
backend.unload_model()
|
|
||||||
clear_voice_prompt_memory_cache()
|
clear_voice_prompt_memory_cache()
|
||||||
|
backend.unload_model()
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
# All other TTS engines
|
# All other TTS engines
|
||||||
backend = get_tts_backend_for_engine(config.engine)
|
backend = get_tts_backend_for_engine(config.engine)
|
||||||
if backend.is_loaded():
|
if backend.is_loaded():
|
||||||
backend.unload_model()
|
|
||||||
clear_voice_prompt_memory_cache()
|
clear_voice_prompt_memory_cache()
|
||||||
|
backend.unload_model()
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|||||||
@@ -232,6 +232,12 @@ def empty_mlx_cache() -> None:
|
|||||||
returning them to the OS. Backends must call this after unloading an
|
returning them to the OS. Backends must call this after unloading an
|
||||||
MLX model, or the process's memory footprint never shrinks even though
|
MLX model, or the process's memory footprint never shrinks even though
|
||||||
the model object itself was dropped.
|
the model object itself was dropped.
|
||||||
|
|
||||||
|
Safe from any thread: ``mx.clear_cache`` only drains the global
|
||||||
|
allocator pool and never touches the per-thread stream registry, so
|
||||||
|
unlike load/generate it does not have to run on the MLX worker thread
|
||||||
|
(verified from the FastAPI event loop with a generation in flight on
|
||||||
|
the worker).
|
||||||
"""
|
"""
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
|
|
||||||
|
|||||||
@@ -24,8 +24,8 @@ def get_tts_model() -> TTSBackend:
|
|||||||
def unload_tts_model():
|
def unload_tts_model():
|
||||||
"""Unload TTS model to free memory."""
|
"""Unload TTS model to free memory."""
|
||||||
backend = get_tts_backend()
|
backend = get_tts_backend()
|
||||||
backend.unload_model()
|
|
||||||
clear_voice_prompt_memory_cache()
|
clear_voice_prompt_memory_cache()
|
||||||
|
backend.unload_model()
|
||||||
|
|
||||||
|
|
||||||
def audio_to_wav_bytes(audio: np.ndarray, sample_rate: int) -> bytes:
|
def audio_to_wav_bytes(audio: np.ndarray, sample_rate: int) -> bytes:
|
||||||
|
|||||||
Reference in New Issue
Block a user