mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-16 21:30:39 -07:00
ADDED MLX FOR SUPER FAST GENERATIONS ON APPLE SILICON
- 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.
This commit is contained in:
@@ -5,7 +5,7 @@ Voice prompt caching utilities.
|
||||
import hashlib
|
||||
import torch
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
from typing import Optional, Union, Dict, Any
|
||||
|
||||
from .. import config
|
||||
|
||||
@@ -15,8 +15,8 @@ def _get_cache_dir() -> Path:
|
||||
return config.get_cache_dir()
|
||||
|
||||
|
||||
# In-memory cache
|
||||
_memory_cache: dict[str, torch.Tensor] = {}
|
||||
# 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:
|
||||
@@ -43,7 +43,7 @@ def get_cache_key(audio_path: str, reference_text: str) -> str:
|
||||
|
||||
def get_cached_voice_prompt(
|
||||
cache_key: str,
|
||||
) -> Optional[torch.Tensor]:
|
||||
) -> Optional[Union[torch.Tensor, Dict[str, Any]]]:
|
||||
"""
|
||||
Get cached voice prompt if available.
|
||||
|
||||
@@ -51,7 +51,7 @@ def get_cached_voice_prompt(
|
||||
cache_key: Cache key
|
||||
|
||||
Returns:
|
||||
Cached voice prompt tensor or None
|
||||
Cached voice prompt (dict or tensor) or None
|
||||
"""
|
||||
# Check in-memory cache
|
||||
if cache_key in _memory_cache:
|
||||
@@ -73,18 +73,18 @@ def get_cached_voice_prompt(
|
||||
|
||||
def cache_voice_prompt(
|
||||
cache_key: str,
|
||||
voice_prompt: torch.Tensor,
|
||||
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 tensor
|
||||
voice_prompt: Voice prompt (dict or tensor)
|
||||
"""
|
||||
# Store in memory
|
||||
_memory_cache[cache_key] = voice_prompt
|
||||
|
||||
# Store on disk
|
||||
# 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)
|
||||
|
||||
Reference in New Issue
Block a user