mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-15 12:50:42 -07:00
- Added support for MLX backend on Apple Silicon, enabling optimized performance for TTS and STT tasks. - Updated release workflow to include MLX-specific dependencies and configurations for macOS platforms. - Refactored backend code to dynamically select between MLX and PyTorch based on the runtime environment. - Enhanced model loading and inference logic to accommodate backend-specific requirements, including updated model IDs and hidden imports. - Improved health check and model status reporting to reflect the active backend type. - Streamlined caching mechanisms to support both backend types, ensuring compatibility and performance.
91 lines
2.1 KiB
Python
91 lines
2.1 KiB
Python
"""
|
|
Voice prompt caching utilities.
|
|
"""
|
|
|
|
import hashlib
|
|
import torch
|
|
from pathlib import Path
|
|
from typing import Optional, Union, Dict, Any
|
|
|
|
from .. import config
|
|
|
|
|
|
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)
|
|
_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)
|