mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-02 16:45:15 -07:00
fix(backend): close model load races
Kokoro and LuxTTS load_model had no lock, so two concurrent requests could both observe an unloaded model and double-load; they now use the same double-checked asyncio.Lock pattern as the Chatterbox backends. get_stt_backend gets the threading.Lock treatment the TTS and LLM factories already had. Includes the ruff-era typing cleanup for backends/__init__.
This commit is contained in:
@@ -10,13 +10,14 @@ and a model config registry that eliminates per-engine dispatch maps.
|
|||||||
# import time, which wraps transformers' tokenizer load against the
|
# import time, which wraps transformers' tokenizer load against the
|
||||||
# unconditional HuggingFace metadata call that otherwise raises on
|
# unconditional HuggingFace metadata call that otherwise raises on
|
||||||
# HF_HUB_OFFLINE=1 and on network failures.
|
# HF_HUB_OFFLINE=1 and on network failures.
|
||||||
from ..utils import hf_offline_patch # noqa: F401
|
|
||||||
|
|
||||||
import threading
|
import threading
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Protocol, Optional, Tuple, List
|
from typing import Protocol
|
||||||
from typing_extensions import runtime_checkable
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
from typing_extensions import runtime_checkable
|
||||||
|
|
||||||
|
from ..utils import hf_offline_patch
|
||||||
|
|
||||||
DEFAULT_LLM_MAX_TOKENS = 512
|
DEFAULT_LLM_MAX_TOKENS = 512
|
||||||
DEFAULT_LLM_TEMPERATURE = 0.7
|
DEFAULT_LLM_TEMPERATURE = 0.7
|
||||||
@@ -76,7 +77,7 @@ class TTSBackend(Protocol):
|
|||||||
audio_path: str,
|
audio_path: str,
|
||||||
reference_text: str,
|
reference_text: str,
|
||||||
use_cache: bool = True,
|
use_cache: bool = True,
|
||||||
) -> Tuple[dict, bool]:
|
) -> tuple[dict, bool]:
|
||||||
"""
|
"""
|
||||||
Create voice prompt from reference audio.
|
Create voice prompt from reference audio.
|
||||||
|
|
||||||
@@ -87,9 +88,9 @@ class TTSBackend(Protocol):
|
|||||||
|
|
||||||
async def combine_voice_prompts(
|
async def combine_voice_prompts(
|
||||||
self,
|
self,
|
||||||
audio_paths: List[str],
|
audio_paths: list[str],
|
||||||
reference_texts: List[str],
|
reference_texts: list[str],
|
||||||
) -> Tuple[np.ndarray, str]:
|
) -> tuple[np.ndarray, str]:
|
||||||
"""
|
"""
|
||||||
Combine multiple voice prompts.
|
Combine multiple voice prompts.
|
||||||
|
|
||||||
@@ -103,9 +104,9 @@ class TTSBackend(Protocol):
|
|||||||
text: str,
|
text: str,
|
||||||
voice_prompt: dict,
|
voice_prompt: dict,
|
||||||
language: str = "en",
|
language: str = "en",
|
||||||
seed: Optional[int] = None,
|
seed: int | None = None,
|
||||||
instruct: Optional[str] = None,
|
instruct: str | None = None,
|
||||||
) -> Tuple[np.ndarray, int]:
|
) -> tuple[np.ndarray, int]:
|
||||||
"""
|
"""
|
||||||
Generate audio from text.
|
Generate audio from text.
|
||||||
|
|
||||||
@@ -143,8 +144,8 @@ class STTBackend(Protocol):
|
|||||||
async def transcribe(
|
async def transcribe(
|
||||||
self,
|
self,
|
||||||
audio_path: str,
|
audio_path: str,
|
||||||
language: Optional[str] = None,
|
language: str | None = None,
|
||||||
model_size: Optional[str] = None,
|
model_size: str | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
Transcribe audio to text.
|
Transcribe audio to text.
|
||||||
@@ -174,11 +175,11 @@ class LLMBackend(Protocol):
|
|||||||
async def generate(
|
async def generate(
|
||||||
self,
|
self,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
system: Optional[str] = None,
|
system: str | None = None,
|
||||||
max_tokens: int = DEFAULT_LLM_MAX_TOKENS,
|
max_tokens: int = DEFAULT_LLM_MAX_TOKENS,
|
||||||
temperature: float = DEFAULT_LLM_TEMPERATURE,
|
temperature: float = DEFAULT_LLM_TEMPERATURE,
|
||||||
model_size: Optional[str] = None,
|
model_size: str | None = None,
|
||||||
examples: Optional[list[tuple[str, str]]] = None,
|
examples: list[tuple[str, str]] | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Run a single-turn chat completion and return the assistant reply.
|
"""Run a single-turn chat completion and return the assistant reply.
|
||||||
|
|
||||||
@@ -198,10 +199,11 @@ class LLMBackend(Protocol):
|
|||||||
|
|
||||||
|
|
||||||
# Global backend instances
|
# Global backend instances
|
||||||
_tts_backend: Optional[TTSBackend] = None
|
_tts_backend: TTSBackend | None = None
|
||||||
_tts_backends: dict[str, TTSBackend] = {}
|
_tts_backends: dict[str, TTSBackend] = {}
|
||||||
_tts_backends_lock = threading.Lock()
|
_tts_backends_lock = threading.Lock()
|
||||||
_stt_backend: Optional[STTBackend] = None
|
_stt_backend: STTBackend | None = None
|
||||||
|
_stt_backend_lock = threading.Lock()
|
||||||
_llm_backends: dict[str, LLMBackend] = {}
|
_llm_backends: dict[str, LLMBackend] = {}
|
||||||
_llm_backends_lock = threading.Lock()
|
_llm_backends_lock = threading.Lock()
|
||||||
|
|
||||||
@@ -488,7 +490,7 @@ def get_stt_model_configs() -> list[ModelConfig]:
|
|||||||
# Lookup helpers — these replace the if/elif chains in main.py
|
# Lookup helpers — these replace the if/elif chains in main.py
|
||||||
|
|
||||||
|
|
||||||
def get_model_config(model_name: str) -> Optional[ModelConfig]:
|
def get_model_config(model_name: str) -> ModelConfig | None:
|
||||||
"""Look up a model config by model_name."""
|
"""Look up a model config by model_name."""
|
||||||
for cfg in get_all_model_configs():
|
for cfg in get_all_model_configs():
|
||||||
if cfg.model_name == model_name:
|
if cfg.model_name == model_name:
|
||||||
@@ -563,8 +565,8 @@ async def unload_backend(backend) -> None:
|
|||||||
|
|
||||||
async def unload_model_by_config(config: ModelConfig) -> bool:
|
async def unload_model_by_config(config: ModelConfig) -> bool:
|
||||||
"""Unload a model given its config. Returns True if it was loaded, False otherwise."""
|
"""Unload a model given its config. Returns True if it was loaded, False otherwise."""
|
||||||
|
from ..services import llm as llm_service, transcribe, tts
|
||||||
from . import get_tts_backend_for_engine
|
from . import get_tts_backend_for_engine
|
||||||
from ..services import tts, transcribe, llm as llm_service
|
|
||||||
|
|
||||||
if config.engine == "whisper":
|
if config.engine == "whisper":
|
||||||
whisper_model = transcribe.get_whisper_model()
|
whisper_model = transcribe.get_whisper_model()
|
||||||
@@ -607,8 +609,8 @@ async def unload_model_by_config(config: ModelConfig) -> bool:
|
|||||||
|
|
||||||
def check_model_loaded(config: ModelConfig) -> bool:
|
def check_model_loaded(config: ModelConfig) -> bool:
|
||||||
"""Check if a model is currently loaded."""
|
"""Check if a model is currently loaded."""
|
||||||
|
from ..services import llm as llm_service, transcribe, tts
|
||||||
from . import get_tts_backend_for_engine
|
from . import get_tts_backend_for_engine
|
||||||
from ..services import tts, transcribe, llm as llm_service
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if config.engine == "whisper":
|
if config.engine == "whisper":
|
||||||
@@ -638,8 +640,8 @@ def check_model_loaded(config: ModelConfig) -> bool:
|
|||||||
|
|
||||||
def get_model_load_func(config: ModelConfig):
|
def get_model_load_func(config: ModelConfig):
|
||||||
"""Return a callable that loads/downloads the model."""
|
"""Return a callable that loads/downloads the model."""
|
||||||
|
from ..services import llm as llm_service, transcribe, tts
|
||||||
from . import get_tts_backend_for_engine
|
from . import get_tts_backend_for_engine
|
||||||
from ..services import tts, transcribe, llm as llm_service
|
|
||||||
|
|
||||||
if config.engine == "whisper":
|
if config.engine == "whisper":
|
||||||
return lambda: transcribe.get_whisper_model().load_model(config.model_size)
|
return lambda: transcribe.get_whisper_model().load_model(config.model_size)
|
||||||
@@ -738,7 +740,13 @@ def get_stt_backend() -> STTBackend:
|
|||||||
"""
|
"""
|
||||||
global _stt_backend
|
global _stt_backend
|
||||||
|
|
||||||
if _stt_backend is None:
|
if _stt_backend is not None:
|
||||||
|
return _stt_backend
|
||||||
|
|
||||||
|
with _stt_backend_lock:
|
||||||
|
if _stt_backend is not None:
|
||||||
|
return _stt_backend
|
||||||
|
|
||||||
backend_type = get_backend_type()
|
backend_type = get_backend_type()
|
||||||
|
|
||||||
if backend_type == "mlx":
|
if backend_type == "mlx":
|
||||||
@@ -750,7 +758,7 @@ def get_stt_backend() -> STTBackend:
|
|||||||
|
|
||||||
_stt_backend = PyTorchSTTBackend()
|
_stt_backend = PyTorchSTTBackend()
|
||||||
|
|
||||||
return _stt_backend
|
return _stt_backend
|
||||||
|
|
||||||
|
|
||||||
def get_llm_backend() -> LLMBackend:
|
def get_llm_backend() -> LLMBackend:
|
||||||
|
|||||||
@@ -17,15 +17,12 @@ Languages supported (via misaki G2P):
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import os
|
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from . import TTSBackend
|
|
||||||
from .base import (
|
from .base import (
|
||||||
get_torch_device,
|
|
||||||
combine_voice_prompts as _combine_voice_prompts,
|
combine_voice_prompts as _combine_voice_prompts,
|
||||||
|
get_torch_device,
|
||||||
model_load_progress,
|
model_load_progress,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -122,8 +119,9 @@ class KokoroTTSBackend:
|
|||||||
def __init__(self):
|
def __init__(self):
|
||||||
self._model = None
|
self._model = None
|
||||||
self._pipelines: dict = {} # lang_code -> KPipeline
|
self._pipelines: dict = {} # lang_code -> KPipeline
|
||||||
self._device: Optional[str] = None
|
self._device: str | None = None
|
||||||
self.model_size = "default"
|
self.model_size = "default"
|
||||||
|
self._model_load_lock = asyncio.Lock()
|
||||||
|
|
||||||
def _get_device(self) -> str:
|
def _get_device(self) -> str:
|
||||||
"""Select device. Kokoro supports CUDA and CPU. MPS needs fallback env var."""
|
"""Select device. Kokoro supports CUDA and CPU. MPS needs fallback env var."""
|
||||||
@@ -157,7 +155,10 @@ class KokoroTTSBackend:
|
|||||||
"""Load the Kokoro model."""
|
"""Load the Kokoro model."""
|
||||||
if self._model is not None:
|
if self._model is not None:
|
||||||
return
|
return
|
||||||
await asyncio.to_thread(self._load_model_sync)
|
async with self._model_load_lock:
|
||||||
|
if self._model is not None:
|
||||||
|
return
|
||||||
|
await asyncio.to_thread(self._load_model_sync)
|
||||||
|
|
||||||
def _load_model_sync(self):
|
def _load_model_sync(self):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
@@ -239,8 +240,8 @@ class KokoroTTSBackend:
|
|||||||
text: str,
|
text: str,
|
||||||
voice_prompt: dict,
|
voice_prompt: dict,
|
||||||
language: str = "en",
|
language: str = "en",
|
||||||
seed: Optional[int] = None,
|
seed: int | None = None,
|
||||||
instruct: Optional[str] = None,
|
instruct: str | None = None,
|
||||||
) -> tuple[np.ndarray, int]:
|
) -> tuple[np.ndarray, int]:
|
||||||
"""
|
"""
|
||||||
Generate audio from text using Kokoro.
|
Generate audio from text using Kokoro.
|
||||||
|
|||||||
@@ -7,20 +7,18 @@ Wraps the LuxTTS (ZipVoice) model for zero-shot voice cloning.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
from typing import Optional, Tuple
|
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from . import TTSBackend
|
from ..utils.cache import cache_voice_prompt, get_cache_key, get_cached_voice_prompt
|
||||||
from .base import (
|
from .base import (
|
||||||
is_model_cached,
|
|
||||||
get_torch_device,
|
|
||||||
empty_device_cache,
|
|
||||||
manual_seed,
|
|
||||||
combine_voice_prompts as _combine_voice_prompts,
|
combine_voice_prompts as _combine_voice_prompts,
|
||||||
|
empty_device_cache,
|
||||||
|
get_torch_device,
|
||||||
|
is_model_cached,
|
||||||
|
manual_seed,
|
||||||
model_load_progress,
|
model_load_progress,
|
||||||
)
|
)
|
||||||
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -35,6 +33,7 @@ class LuxTTSBackend:
|
|||||||
self.model = None
|
self.model = None
|
||||||
self.model_size = "default" # LuxTTS has only one model size
|
self.model_size = "default" # LuxTTS has only one model size
|
||||||
self._device = None
|
self._device = None
|
||||||
|
self._model_load_lock = asyncio.Lock()
|
||||||
|
|
||||||
def _get_device(self) -> str:
|
def _get_device(self) -> str:
|
||||||
return get_torch_device(allow_mps=True, allow_xpu=True)
|
return get_torch_device(allow_mps=True, allow_xpu=True)
|
||||||
@@ -61,8 +60,10 @@ class LuxTTSBackend:
|
|||||||
"""Load the LuxTTS model."""
|
"""Load the LuxTTS model."""
|
||||||
if self.model is not None:
|
if self.model is not None:
|
||||||
return
|
return
|
||||||
|
async with self._model_load_lock:
|
||||||
await asyncio.to_thread(self._load_model_sync)
|
if self.model is not None:
|
||||||
|
return
|
||||||
|
await asyncio.to_thread(self._load_model_sync)
|
||||||
|
|
||||||
def _load_model_sync(self):
|
def _load_model_sync(self):
|
||||||
model_name = "luxtts"
|
model_name = "luxtts"
|
||||||
@@ -105,7 +106,7 @@ class LuxTTSBackend:
|
|||||||
audio_path: str,
|
audio_path: str,
|
||||||
reference_text: str,
|
reference_text: str,
|
||||||
use_cache: bool = True,
|
use_cache: bool = True,
|
||||||
) -> Tuple[dict, bool]:
|
) -> tuple[dict, bool]:
|
||||||
"""
|
"""
|
||||||
Create voice prompt from reference audio.
|
Create voice prompt from reference audio.
|
||||||
|
|
||||||
@@ -145,9 +146,9 @@ class LuxTTSBackend:
|
|||||||
text: str,
|
text: str,
|
||||||
voice_prompt: dict,
|
voice_prompt: dict,
|
||||||
language: str = "en",
|
language: str = "en",
|
||||||
seed: Optional[int] = None,
|
seed: int | None = None,
|
||||||
instruct: Optional[str] = None,
|
instruct: str | None = None,
|
||||||
) -> Tuple[np.ndarray, int]:
|
) -> tuple[np.ndarray, int]:
|
||||||
"""
|
"""
|
||||||
Generate audio from text using LuxTTS.
|
Generate audio from text using LuxTTS.
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user