mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -07:00
Unloading a TTS/Whisper/LLM model on the MLX backend only dropped the
Python reference (`del self.model`). MLX keeps freed array buffers in
its own allocator pool for reuse instead of returning them to the OS,
so the process's memory footprint never actually shrank after unload
on Apple Silicon (the default backend there) until the process exited.
Add empty_mlx_cache() (backend/backends/base.py), wrapping
mx.clear_cache(), and call it from the three MLX unload_model()
implementations: MLXTTSBackend, MLXSTTBackend, MLXQwenLLMBackend.
Separately, the voice-clone prompt cache (backend/utils/cache.py) is a
process-lifetime dict populated by create_voice_prompt() across every
TTS engine, but nothing ever cleared it on model unload — only the
unrelated /tasks/clear-cache endpoint touched it. Add
clear_voice_prompt_memory_cache() (memory only, disk cache untouched
so a later generation still reloads the prompt instead of recomputing
it) and wire it into every TTS unload path (services/tts.py and the
qwen_custom_voice / generic branches of unload_model_by_config).
Whisper and the LLM backends never produce voice prompts, so their
unload paths are left alone.
Testing:
- New unit tests: backend/tests/test_mlx_unload_clears_cache.py,
backend/tests/test_voice_prompt_cache_unload.py (8 tests, all pass).
- Verified end-to-end on Apple Silicon against real cached models
(Qwen TTS 1.7B, Whisper Turbo, Qwen3 0.6B): loaded each via the
running app, unloaded via the real /models/{name}/unload endpoint,
and confirmed via mx.get_cache_memory()/get_active_memory() that the
MLX allocator's cache drops to 0 on every cycle. Ran a real
voice-clone generation end to end and confirmed the in-memory prompt
cache goes from 1 entry to 0 on unload while the on-disk .prompt
file is left intact.
167 lines
4.3 KiB
Python
167 lines
4.3 KiB
Python
"""
|
|
Voice prompt caching utilities.
|
|
"""
|
|
|
|
import hashlib
|
|
import logging
|
|
import torch
|
|
from pathlib import Path
|
|
from typing import Optional, Union, Dict, Any
|
|
|
|
from .. import config
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def _get_cache_dir() -> Path:
|
|
"""Get cache directory from config."""
|
|
return config.get_cache_dir()
|
|
|
|
|
|
# In-memory cache - can store dict (voice prompt) or tensor (legacy)
|
|
_memory_cache: dict[str, Union[torch.Tensor, Dict[str, Any]]] = {}
|
|
|
|
|
|
def get_cache_key(audio_path: str, reference_text: str) -> str:
|
|
"""
|
|
Generate cache key from audio file and reference text.
|
|
|
|
Args:
|
|
audio_path: Path to audio file
|
|
reference_text: Reference text
|
|
|
|
Returns:
|
|
Cache key (MD5 hash)
|
|
"""
|
|
# Read audio file
|
|
with open(audio_path, "rb") as f:
|
|
audio_bytes = f.read()
|
|
|
|
# Combine audio bytes and text
|
|
combined = audio_bytes + reference_text.encode("utf-8")
|
|
|
|
# Generate hash
|
|
return hashlib.md5(combined).hexdigest()
|
|
|
|
|
|
def get_cached_voice_prompt(
|
|
cache_key: str,
|
|
) -> Optional[Union[torch.Tensor, Dict[str, Any]]]:
|
|
"""
|
|
Get cached voice prompt if available.
|
|
|
|
Args:
|
|
cache_key: Cache key
|
|
|
|
Returns:
|
|
Cached voice prompt (dict or tensor) or None
|
|
"""
|
|
# Check in-memory cache
|
|
if cache_key in _memory_cache:
|
|
return _memory_cache[cache_key]
|
|
|
|
# Check disk cache
|
|
cache_file = _get_cache_dir() / f"{cache_key}.prompt"
|
|
if cache_file.exists():
|
|
try:
|
|
prompt = torch.load(cache_file, weights_only=True)
|
|
_memory_cache[cache_key] = prompt
|
|
return prompt
|
|
except Exception:
|
|
# Cache file corrupted, delete it
|
|
cache_file.unlink()
|
|
|
|
return None
|
|
|
|
|
|
def cache_voice_prompt(
|
|
cache_key: str,
|
|
voice_prompt: Union[torch.Tensor, Dict[str, Any]],
|
|
) -> None:
|
|
"""
|
|
Cache voice prompt to memory and disk.
|
|
|
|
Args:
|
|
cache_key: Cache key
|
|
voice_prompt: Voice prompt (dict or tensor)
|
|
"""
|
|
# Store in memory
|
|
_memory_cache[cache_key] = voice_prompt
|
|
|
|
# Store on disk (torch.save can handle both dicts and tensors)
|
|
cache_file = _get_cache_dir() / f"{cache_key}.prompt"
|
|
torch.save(voice_prompt, cache_file)
|
|
|
|
|
|
def clear_voice_prompt_memory_cache() -> None:
|
|
"""
|
|
Drop the in-memory voice prompt cache without touching the disk cache.
|
|
|
|
Backends call this when a TTS model unloads: the cached prompts (tensors
|
|
or device-backed dicts produced by that model) would otherwise keep
|
|
referencing memory forever, since nothing else ever clears this
|
|
process-lifetime dict. The disk cache is left alone, so the next
|
|
generation just reloads the prompt from disk instead of recomputing it.
|
|
"""
|
|
_memory_cache.clear()
|
|
|
|
|
|
def clear_voice_prompt_cache() -> int:
|
|
"""
|
|
Clear all voice prompt caches (memory and disk).
|
|
|
|
Returns:
|
|
Number of cache files deleted
|
|
"""
|
|
# Clear memory cache
|
|
_memory_cache.clear()
|
|
|
|
# Clear disk cache
|
|
cache_dir = _get_cache_dir()
|
|
deleted_count = 0
|
|
|
|
if cache_dir.exists():
|
|
# Delete prompt cache files
|
|
for cache_file in cache_dir.glob("*.prompt"):
|
|
try:
|
|
cache_file.unlink()
|
|
deleted_count += 1
|
|
except Exception as e:
|
|
logger.warning("Failed to delete cache file %s: %s", cache_file, e)
|
|
|
|
# Delete combined audio files
|
|
for audio_file in cache_dir.glob("combined_*.wav"):
|
|
try:
|
|
audio_file.unlink()
|
|
deleted_count += 1
|
|
except Exception as e:
|
|
logger.warning("Failed to delete combined audio file %s: %s", audio_file, e)
|
|
|
|
return deleted_count
|
|
|
|
|
|
def clear_profile_cache(profile_id: str) -> int:
|
|
"""
|
|
Clear cache files for a specific profile.
|
|
|
|
Args:
|
|
profile_id: Profile ID
|
|
|
|
Returns:
|
|
Number of cache files deleted
|
|
"""
|
|
cache_dir = _get_cache_dir()
|
|
deleted_count = 0
|
|
|
|
if cache_dir.exists():
|
|
# Delete combined audio files for this profile
|
|
pattern = f"combined_{profile_id}_*.wav"
|
|
for audio_file in cache_dir.glob(pattern):
|
|
try:
|
|
audio_file.unlink()
|
|
deleted_count += 1
|
|
except Exception as e:
|
|
logger.warning("Failed to delete combined audio file %s: %s", audio_file, e)
|
|
|
|
return deleted_count
|