fix take-label race in regeneration, add accessible focus to select

- Use DB COUNT query instead of list length for take-N label to avoid
  TOCTOU race between list_versions and create_version
- Add focus:bg-muted to SelectTrigger for keyboard focus visibility
This commit is contained in:
James Pine
2026-03-16 03:22:05 -07:00
parent 0d0b62ea93
commit 473bb3e9fb
6 changed files with 428 additions and 279 deletions
+1 -1
View File
@@ -16,7 +16,7 @@ const SelectTrigger = React.forwardRef<
<SelectPrimitive.Trigger <SelectPrimitive.Trigger
ref={ref} ref={ref}
className={cn( className={cn(
'flex h-10 w-full items-center justify-between rounded-md border border-input bg-background px-3 py-2 text-sm ring-offset-background placeholder:text-muted-foreground focus:outline-none disabled:cursor-not-allowed disabled:opacity-50 [&>span]:line-clamp-1', 'flex h-10 w-full items-center justify-between rounded-md border border-input bg-background px-3 py-2 text-sm ring-offset-background placeholder:text-muted-foreground focus:outline-none focus:bg-muted disabled:cursor-not-allowed disabled:opacity-50 [&>span]:line-clamp-1',
className, className,
)} )}
{...props} {...props}
+21 -20
View File
@@ -1,9 +1,12 @@
"""FastAPI application factory, middleware, and lifecycle events.""" """FastAPI application factory, middleware, and lifecycle events."""
import asyncio import asyncio
import logging
import os import os
from pathlib import Path from pathlib import Path
logger = logging.getLogger(__name__)
# AMD GPU environment variables must be set before torch import # AMD GPU environment variables must be set before torch import
if not os.environ.get("HSA_OVERRIDE_GFX_VERSION"): if not os.environ.get("HSA_OVERRIDE_GFX_VERSION"):
os.environ["HSA_OVERRIDE_GFX_VERSION"] = "10.3.0" os.environ["HSA_OVERRIDE_GFX_VERSION"] = "10.3.0"
@@ -30,11 +33,9 @@ def safe_content_disposition(disposition_type: str, filename: str) -> str:
Uses RFC 5987 ``filename*`` parameter so browsers can decode UTF-8 Uses RFC 5987 ``filename*`` parameter so browsers can decode UTF-8
filenames while the ``filename`` fallback stays ASCII-only. filenames while the ``filename`` fallback stays ASCII-only.
""" """
ascii_name = ( ascii_name = "".join(c for c in filename if c.isascii() and (c.isalnum() or c in " -_.")).strip() or "download"
"".join(c for c in filename if c.isascii() and (c.isalnum() or c in " -_.")).strip() or "download"
)
utf8_name = quote(filename, safe="") utf8_name = quote(filename, safe="")
return f'{disposition_type}; filename="{ascii_name}"; filename*=UTF-8\'\'{utf8_name}' return f"{disposition_type}; filename=\"{ascii_name}\"; filename*=UTF-8''{utf8_name}"
def create_app() -> FastAPI: def create_app() -> FastAPI:
@@ -55,13 +56,13 @@ def create_app() -> FastAPI:
def _configure_cors(application: FastAPI) -> None: def _configure_cors(application: FastAPI) -> None:
"""Set up CORS middleware with local-first defaults.""" """Set up CORS middleware with local-first defaults."""
default_origins = [ default_origins = [
"http://localhost:5173", # Vite dev server "http://localhost:5173", # Vite dev server
"http://127.0.0.1:5173", "http://127.0.0.1:5173",
"http://localhost:17493", "http://localhost:17493",
"http://127.0.0.1:17493", "http://127.0.0.1:17493",
"tauri://localhost", # Tauri webview (macOS) "tauri://localhost", # Tauri webview (macOS)
"https://tauri.localhost", # Tauri webview (Windows/Linux) "https://tauri.localhost", # Tauri webview (Windows/Linux)
"http://tauri.localhost", # Tauri webview (Windows, some builds) "http://tauri.localhost", # Tauri webview (Windows, some builds)
] ]
env_origins = os.environ.get("VOICEBOX_CORS_ORIGINS", "") env_origins = os.environ.get("VOICEBOX_CORS_ORIGINS", "")
all_origins = default_origins + [o.strip() for o in env_origins.split(",") if o.strip()] all_origins = default_origins + [o.strip() for o in env_origins.split(",") if o.strip()]
@@ -96,9 +97,9 @@ def _register_lifecycle(application: FastAPI) -> None:
@application.on_event("startup") @application.on_event("startup")
async def startup_event(): async def startup_event():
print("voicebox API starting up...") logger.info("Voicebox server starting up...")
database.init_db() database.init_db()
print(f"Database initialized at {database._db_path}") logger.info("Database initialized at %s", database._db_path)
init_queue() init_queue()
@@ -115,15 +116,15 @@ def _register_lifecycle(application: FastAPI) -> None:
) )
) )
if result.rowcount > 0: if result.rowcount > 0:
print(f"Marked {result.rowcount} stale generation(s) as failed") logger.info("Marked %d stale generation(s) as failed", result.rowcount)
db.commit() db.commit()
db.close() db.close()
except Exception as e: except Exception as e:
print(f"Warning: Could not clean up stale generations: {e}") logger.warning("Could not clean up stale generations: %s", e)
backend_type = get_backend_type() backend_type = get_backend_type()
print(f"Backend: {backend_type.upper()}") logger.info("Backend: %s", backend_type.upper())
print(f"GPU available: {_get_gpu_status()}") logger.info("GPU available: %s", _get_gpu_status())
from .services.cuda import check_and_update_cuda_binary from .services.cuda import check_and_update_cuda_binary
@@ -132,23 +133,23 @@ def _register_lifecycle(application: FastAPI) -> None:
try: try:
progress_manager = get_progress_manager() progress_manager = get_progress_manager()
progress_manager._set_main_loop(asyncio.get_running_loop()) progress_manager._set_main_loop(asyncio.get_running_loop())
print("Progress manager initialized with event loop") logger.info("Progress manager initialized with event loop")
except Exception as e: except Exception as e:
print(f"Warning: Could not initialize progress manager event loop: {e}") logger.warning("Could not initialize progress manager event loop: %s", e)
try: try:
from huggingface_hub import constants as hf_constants from huggingface_hub import constants as hf_constants
cache_dir = Path(hf_constants.HF_HUB_CACHE) cache_dir = Path(hf_constants.HF_HUB_CACHE)
cache_dir.mkdir(parents=True, exist_ok=True) cache_dir.mkdir(parents=True, exist_ok=True)
print(f"HuggingFace cache directory: {cache_dir}") logger.info("HuggingFace cache directory: %s", cache_dir)
except Exception as e: except Exception as e:
print(f"Warning: Could not create HuggingFace cache directory: {e}") logger.warning("Could not create HuggingFace cache directory: %s", e)
print("Model downloads may fail. Please ensure the directory exists and has write permissions.") logger.warning("Model downloads may fail. Please ensure the directory exists and has write permissions.")
@application.on_event("shutdown") @application.on_event("shutdown")
async def shutdown_event(): async def shutdown_event():
print("voicebox API shutting down...") logger.info("Voicebox server shutting down...")
tts.unload_tts_model() tts.unload_tts_model()
transcribe.unload_whisper_model() transcribe.unload_whisper_model()
+23 -15
View File
@@ -4,13 +4,17 @@ MLX backend implementation for TTS and STT using mlx-audio.
from typing import Optional, List, Tuple from typing import Optional, List, Tuple
import asyncio import asyncio
import logging
import numpy as np import numpy as np
import os import os
from pathlib import Path from pathlib import Path
logger = logging.getLogger(__name__)
# PATCH: Import and apply offline patch BEFORE any huggingface_hub usage # PATCH: Import and apply offline patch BEFORE any huggingface_hub usage
# This prevents mlx_audio from making network requests when models are cached # This prevents mlx_audio from making network requests when models are cached
from ..utils.hf_offline_patch import patch_huggingface_hub_offline, ensure_original_qwen_config_cached from ..utils.hf_offline_patch import patch_huggingface_hub_offline, ensure_original_qwen_config_cached
patch_huggingface_hub_offline() patch_huggingface_hub_offline()
ensure_original_qwen_config_cached() ensure_original_qwen_config_cached()
@@ -52,7 +56,7 @@ class MLXTTSBackend:
raise ValueError(f"Unknown model size: {model_size}") raise ValueError(f"Unknown model size: {model_size}")
hf_model_id = mlx_model_map[model_size] hf_model_id = mlx_model_map[model_size]
print(f"Will download MLX model from HuggingFace Hub: {hf_model_id}") logger.info("Will download MLX model from HuggingFace Hub: %s", hf_model_id)
return hf_model_id return hf_model_id
@@ -96,18 +100,19 @@ class MLXTTSBackend:
original_hf_hub_offline = os.environ.get("HF_HUB_OFFLINE") original_hf_hub_offline = os.environ.get("HF_HUB_OFFLINE")
if is_cached: if is_cached:
os.environ["HF_HUB_OFFLINE"] = "1" os.environ["HF_HUB_OFFLINE"] = "1"
print(f"[PATCH] Model {model_size} is cached, forcing HF_HUB_OFFLINE=1 to avoid network requests") logger.info("[PATCH] Model %s is cached, forcing HF_HUB_OFFLINE=1 to avoid network requests", model_size)
try: try:
with model_load_progress(model_name, is_cached): with model_load_progress(model_name, is_cached):
from mlx_audio.tts import load from mlx_audio.tts import load
print(f"Loading MLX TTS model {model_size}...")
logger.info("Loading MLX TTS model %s...", model_size)
try: try:
self.model = load(model_path) self.model = load(model_path)
except Exception as load_error: except Exception as load_error:
if is_cached and "offline" in str(load_error).lower(): if is_cached and "offline" in str(load_error).lower():
print(f"[PATCH] Offline load failed, trying with network: {load_error}") logger.warning("[PATCH] Offline load failed, trying with network: %s", load_error)
os.environ.pop("HF_HUB_OFFLINE", None) os.environ.pop("HF_HUB_OFFLINE", None)
self.model = load(model_path) self.model = load(model_path)
else: else:
@@ -120,7 +125,7 @@ class MLXTTSBackend:
self._current_model_size = model_size self._current_model_size = model_size
self.model_size = model_size self.model_size = model_size
print(f"MLX TTS model {model_size} loaded successfully") logger.info("MLX TTS model %s loaded successfully", model_size)
def unload_model(self): def unload_model(self):
"""Unload the model to free memory.""" """Unload the model to free memory."""
@@ -128,7 +133,7 @@ class MLXTTSBackend:
del self.model del self.model
self.model = None self.model = None
self._current_model_size = None self._current_model_size = None
print("MLX TTS model unloaded") logger.info("MLX TTS model unloaded")
async def create_voice_prompt( async def create_voice_prompt(
self, self,
@@ -165,7 +170,7 @@ class MLXTTSBackend:
return cached_prompt, True return cached_prompt, True
else: else:
# Cached file no longer exists, invalidate cache # Cached file no longer exists, invalidate cache
print(f"Cached audio file not found: {cached_audio_path}, regenerating prompt") logger.warning("Cached audio file not found: %s, regenerating prompt", cached_audio_path)
# MLX voice prompt format - store audio path and text # MLX voice prompt format - store audio path and text
# The model will process this during generation # The model will process this during generation
@@ -207,7 +212,7 @@ class MLXTTSBackend:
""" """
await self.load_model_async(None) await self.load_model_async(None)
print(f"Generating audio for text: {text}") logger.info("Generating audio for text: %s", text)
def _generate_sync(): def _generate_sync():
"""Run synchronous generation in thread pool.""" """Run synchronous generation in thread pool."""
@@ -219,6 +224,7 @@ class MLXTTSBackend:
# Set seed if provided (MLX uses numpy random) # Set seed if provided (MLX uses numpy random)
if seed is not None: if seed is not None:
import mlx.core as mx import mlx.core as mx
np.random.seed(seed) np.random.seed(seed)
mx.random.seed(seed) mx.random.seed(seed)
@@ -228,9 +234,9 @@ class MLXTTSBackend:
# Validate that the audio file exists # Validate that the audio file exists
if ref_audio and not Path(ref_audio).exists(): if ref_audio and not Path(ref_audio).exists():
print(f"Warning: Audio file not found: {ref_audio}") logger.warning("Audio file not found: %s", ref_audio)
print("This may be due to a cached voice prompt referencing a deleted temp file.") logger.warning("This may be due to a cached voice prompt referencing a deleted temp file.")
print("Regenerating without voice prompt.") logger.warning("Regenerating without voice prompt.")
ref_audio = None ref_audio = None
# Check if model supports voice cloning via generate method # Check if model supports voice cloning via generate method
@@ -240,6 +246,7 @@ class MLXTTSBackend:
if ref_audio: if ref_audio:
# Check if generate accepts ref_audio parameter # Check if generate accepts ref_audio parameter
import inspect import inspect
sig = inspect.signature(self.model.generate) sig = inspect.signature(self.model.generate)
if "ref_audio" in sig.parameters: if "ref_audio" in sig.parameters:
# Generate with voice cloning # Generate with voice cloning
@@ -258,7 +265,7 @@ class MLXTTSBackend:
sample_rate = result.sample_rate sample_rate = result.sample_rate
except Exception as e: except Exception as e:
# If voice cloning fails, try without it # If voice cloning fails, try without it
print(f"Warning: Voice cloning failed, generating without voice prompt: {e}") logger.warning("Voice cloning failed, generating without voice prompt: %s", e)
for result in self.model.generate(text, lang_code=lang): for result in self.model.generate(text, lang_code=lang):
audio_chunks.append(np.array(result.audio)) audio_chunks.append(np.array(result.audio))
sample_rate = result.sample_rate sample_rate = result.sample_rate
@@ -319,19 +326,20 @@ class MLXSTTBackend:
with model_load_progress(progress_model_name, is_cached): with model_load_progress(progress_model_name, is_cached):
from mlx_audio.stt import load from mlx_audio.stt import load
model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}") model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
print(f"Loading MLX Whisper model {model_size}...") logger.info("Loading MLX Whisper model %s...", model_size)
self.model = load(model_name) self.model = load(model_name)
self.model_size = model_size self.model_size = model_size
print(f"MLX Whisper model {model_size} loaded successfully") logger.info("MLX Whisper model %s loaded successfully", model_size)
def unload_model(self): def unload_model(self):
"""Unload the model to free memory.""" """Unload the model to free memory."""
if self.model is not None: if self.model is not None:
del self.model del self.model
self.model = None self.model = None
print("MLX Whisper model unloaded") logger.info("MLX Whisper model unloaded")
async def transcribe( async def transcribe(
self, self,
+17 -7
View File
@@ -4,11 +4,19 @@ PyTorch backend implementation for TTS and STT.
from typing import Optional, List, Tuple from typing import Optional, List, Tuple
import asyncio import asyncio
import logging
import torch import torch
import numpy as np import numpy as np
logger = logging.getLogger(__name__)
from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS
from .base import is_model_cached, get_torch_device, combine_voice_prompts as _combine_voice_prompts, model_load_progress from .base import (
is_model_cached,
get_torch_device,
combine_voice_prompts as _combine_voice_prompts,
model_load_progress,
)
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
from ..utils.audio import load_audio from ..utils.audio import load_audio
@@ -84,8 +92,9 @@ class PyTorchTTSBackend:
with model_load_progress(model_name, is_cached): with model_load_progress(model_name, is_cached):
from qwen_tts import Qwen3TTSModel from qwen_tts import Qwen3TTSModel
model_path = self._get_model_path(model_size) model_path = self._get_model_path(model_size)
print(f"Loading TTS model {model_size} on {self.device}...") logger.info("Loading TTS model %s on %s...", model_size, self.device)
if self.device == "cpu": if self.device == "cpu":
self.model = Qwen3TTSModel.from_pretrained( self.model = Qwen3TTSModel.from_pretrained(
@@ -102,7 +111,7 @@ class PyTorchTTSBackend:
self._current_model_size = model_size self._current_model_size = model_size
self.model_size = model_size self.model_size = model_size
print(f"TTS model {model_size} loaded successfully") logger.info("TTS model %s loaded successfully", model_size)
def unload_model(self): def unload_model(self):
"""Unload the model to free memory.""" """Unload the model to free memory."""
@@ -114,7 +123,7 @@ class PyTorchTTSBackend:
if torch.cuda.is_available(): if torch.cuda.is_available():
torch.cuda.empty_cache() torch.cuda.empty_cache()
print("TTS model unloaded") logger.info("TTS model unloaded")
async def create_voice_prompt( async def create_voice_prompt(
self, self,
@@ -269,15 +278,16 @@ class PyTorchSTTBackend:
with model_load_progress(progress_model_name, is_cached): with model_load_progress(progress_model_name, is_cached):
from transformers import WhisperProcessor, WhisperForConditionalGeneration from transformers import WhisperProcessor, WhisperForConditionalGeneration
model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}") model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
print(f"Loading Whisper model {model_size} on {self.device}...") logger.info("Loading Whisper model %s on %s...", model_size, self.device)
self.processor = WhisperProcessor.from_pretrained(model_name) self.processor = WhisperProcessor.from_pretrained(model_name)
self.model = WhisperForConditionalGeneration.from_pretrained(model_name) self.model = WhisperForConditionalGeneration.from_pretrained(model_name)
self.model.to(self.device) self.model.to(self.device)
self.model_size = model_size self.model_size = model_size
print(f"Whisper model {model_size} loaded successfully") logger.info("Whisper model %s loaded successfully", model_size)
def unload_model(self): def unload_model(self):
"""Unload the model to free memory.""" """Unload the model to free memory."""
@@ -290,7 +300,7 @@ class PyTorchSTTBackend:
if torch.cuda.is_available(): if torch.cuda.is_available():
torch.cuda.empty_cache() torch.cuda.empty_cache()
print("Whisper model unloaded") logger.info("Whisper model unloaded")
async def transcribe( async def transcribe(
self, self,
+259 -132
View File
@@ -8,11 +8,14 @@ Usage:
import PyInstaller.__main__ import PyInstaller.__main__
import argparse import argparse
import logging
import os import os
import platform import platform
import sys import sys
from pathlib import Path from pathlib import Path
logger = logging.getLogger(__name__)
def is_apple_silicon(): def is_apple_silicon():
"""Check if running on Apple Silicon.""" """Check if running on Apple Silicon."""
@@ -28,151 +31,242 @@ def build_server(cuda=False):
""" """
backend_dir = Path(__file__).parent backend_dir = Path(__file__).parent
binary_name = 'voicebox-server-cuda' if cuda else 'voicebox-server' binary_name = "voicebox-server-cuda" if cuda else "voicebox-server"
# PyInstaller arguments # PyInstaller arguments
args = [ args = [
'server.py', # Use server.py as entry point instead of main.py "server.py", # Use server.py as entry point instead of main.py
'--onefile', "--onefile",
'--name', binary_name, "--name",
binary_name,
] ]
# Hide console window on Windows only. On macOS/Linux the sidecar needs # Hide console window on Windows only. On macOS/Linux the sidecar needs
# stdout/stderr for Tauri to capture logs. # stdout/stderr for Tauri to capture logs.
if platform.system() == "Windows": if platform.system() == "Windows":
args.append('--noconsole') args.append("--noconsole")
# Add local qwen_tts path if specified (for editable installs) # Add local qwen_tts path if specified (for editable installs)
qwen_tts_path = os.getenv('QWEN_TTS_PATH') qwen_tts_path = os.getenv("QWEN_TTS_PATH")
if qwen_tts_path and Path(qwen_tts_path).exists(): if qwen_tts_path and Path(qwen_tts_path).exists():
args.extend(['--paths', str(qwen_tts_path)]) args.extend(["--paths", str(qwen_tts_path)])
print(f"Using local qwen_tts source from: {qwen_tts_path}") print(f"Using local qwen_tts source from: {qwen_tts_path}")
# Add common hidden imports # Add common hidden imports
args.extend([ args.extend(
'--hidden-import', 'backend', [
'--hidden-import', 'backend.main', "--hidden-import",
'--hidden-import', 'backend.config', "backend",
'--hidden-import', 'backend.database', "--hidden-import",
'--hidden-import', 'backend.models', "backend.main",
'--hidden-import', 'backend.services.profiles', "--hidden-import",
'--hidden-import', 'backend.services.history', "backend.config",
'--hidden-import', 'backend.services.tts', "--hidden-import",
'--hidden-import', 'backend.services.transcribe', "backend.database",
'--hidden-import', 'backend.utils.platform_detect', "--hidden-import",
'--hidden-import', 'backend.backends', "backend.models",
'--hidden-import', 'backend.backends.pytorch_backend', "--hidden-import",
'--hidden-import', 'backend.utils.audio', "backend.services.profiles",
'--hidden-import', 'backend.utils.cache', "--hidden-import",
'--hidden-import', 'backend.utils.progress', "backend.services.history",
'--hidden-import', 'backend.utils.hf_progress', "--hidden-import",
"backend.services.tts",
'--hidden-import', 'backend.services.cuda', "--hidden-import",
'--hidden-import', 'backend.services.effects', "backend.services.transcribe",
'--hidden-import', 'backend.utils.effects', "--hidden-import",
'--hidden-import', 'backend.services.versions', "backend.utils.platform_detect",
'--hidden-import', 'pedalboard', "--hidden-import",
'--hidden-import', 'chatterbox', "backend.backends",
'--hidden-import', 'chatterbox.tts_turbo', "--hidden-import",
'--hidden-import', 'chatterbox.mtl_tts', "backend.backends.pytorch_backend",
'--hidden-import', 'backend.backends.chatterbox_backend', "--hidden-import",
'--hidden-import', 'backend.backends.chatterbox_turbo_backend', "backend.utils.audio",
'--hidden-import', 'backend.backends.luxtts_backend', "--hidden-import",
'--hidden-import', 'zipvoice', "backend.utils.cache",
'--hidden-import', 'zipvoice.luxvoice', "--hidden-import",
'--collect-all', 'zipvoice', "backend.utils.progress",
'--collect-all', 'linacodec', "--hidden-import",
'--hidden-import', 'torch', "backend.utils.hf_progress",
'--hidden-import', 'transformers', "--hidden-import",
'--hidden-import', 'fastapi', "backend.services.cuda",
'--hidden-import', 'uvicorn', "--hidden-import",
'--hidden-import', 'sqlalchemy', "backend.services.effects",
'--hidden-import', 'librosa', "--hidden-import",
'--hidden-import', 'soundfile', "backend.utils.effects",
'--hidden-import', 'qwen_tts', "--hidden-import",
'--hidden-import', 'qwen_tts.inference', "backend.services.versions",
'--hidden-import', 'qwen_tts.inference.qwen3_tts_model', "--hidden-import",
'--hidden-import', 'qwen_tts.inference.qwen3_tts_tokenizer', "pedalboard",
'--hidden-import', 'qwen_tts.core', "--hidden-import",
'--hidden-import', 'qwen_tts.cli', "chatterbox",
'--copy-metadata', 'qwen-tts', "--hidden-import",
'--copy-metadata', 'requests', "chatterbox.tts_turbo",
'--copy-metadata', 'transformers', "--hidden-import",
'--copy-metadata', 'huggingface-hub', "chatterbox.mtl_tts",
'--copy-metadata', 'tokenizers', "--hidden-import",
'--copy-metadata', 'safetensors', "backend.backends.chatterbox_backend",
'--copy-metadata', 'tqdm', "--hidden-import",
'--hidden-import', 'requests', "backend.backends.chatterbox_turbo_backend",
'--collect-submodules', 'qwen_tts', "--hidden-import",
'--collect-data', 'qwen_tts', "backend.backends.luxtts_backend",
# Fix for pkg_resources and jaraco namespace packages "--hidden-import",
'--hidden-import', 'pkg_resources.extern', "zipvoice",
'--collect-submodules', 'jaraco', "--hidden-import",
# inflect uses typeguard @typechecked which calls inspect.getsource() "zipvoice.luxvoice",
# at import time — needs .py source files, not just .pyc bytecode "--collect-all",
'--collect-all', 'inflect', "zipvoice",
# perth ships pretrained watermark model files (hparams.yaml, .pth.tar) "--collect-all",
# in perth/perth_net/pretrained/ — needed by chatterbox at runtime "linacodec",
'--collect-all', 'perth', "--hidden-import",
# piper_phonemize ships espeak-ng-data/ (phoneme tables, language dicts) "torch",
# needed by LuxTTS for text-to-phoneme conversion "--hidden-import",
'--collect-all', 'piper_phonemize', "transformers",
]) "--hidden-import",
"fastapi",
"--hidden-import",
"uvicorn",
"--hidden-import",
"sqlalchemy",
"--hidden-import",
"librosa",
"--hidden-import",
"soundfile",
"--hidden-import",
"qwen_tts",
"--hidden-import",
"qwen_tts.inference",
"--hidden-import",
"qwen_tts.inference.qwen3_tts_model",
"--hidden-import",
"qwen_tts.inference.qwen3_tts_tokenizer",
"--hidden-import",
"qwen_tts.core",
"--hidden-import",
"qwen_tts.cli",
"--copy-metadata",
"qwen-tts",
"--copy-metadata",
"requests",
"--copy-metadata",
"transformers",
"--copy-metadata",
"huggingface-hub",
"--copy-metadata",
"tokenizers",
"--copy-metadata",
"safetensors",
"--copy-metadata",
"tqdm",
"--hidden-import",
"requests",
"--collect-submodules",
"qwen_tts",
"--collect-data",
"qwen_tts",
# Fix for pkg_resources and jaraco namespace packages
"--hidden-import",
"pkg_resources.extern",
"--collect-submodules",
"jaraco",
# inflect uses typeguard @typechecked which calls inspect.getsource()
# at import time — needs .py source files, not just .pyc bytecode
"--collect-all",
"inflect",
# perth ships pretrained watermark model files (hparams.yaml, .pth.tar)
# in perth/perth_net/pretrained/ — needed by chatterbox at runtime
"--collect-all",
"perth",
# piper_phonemize ships espeak-ng-data/ (phoneme tables, language dicts)
# needed by LuxTTS for text-to-phoneme conversion
"--collect-all",
"piper_phonemize",
]
)
# Add CUDA-specific hidden imports # Add CUDA-specific hidden imports
if cuda: if cuda:
print("Building with CUDA support") print("Building with CUDA support")
args.extend([ args.extend(
'--hidden-import', 'torch.cuda', [
'--hidden-import', 'torch.backends.cudnn', "--hidden-import",
]) "torch.cuda",
"--hidden-import",
"torch.backends.cudnn",
]
)
else: else:
# Exclude NVIDIA CUDA packages from CPU-only builds to keep binary small. # Exclude NVIDIA CUDA packages from CPU-only builds to keep binary small.
# When building from a venv with CUDA torch installed, PyInstaller would # When building from a venv with CUDA torch installed, PyInstaller would
# bundle ~3GB of NVIDIA shared libraries. We exclude both the Python # bundle ~3GB of NVIDIA shared libraries. We exclude both the Python
# modules and the binary DLLs. # modules and the binary DLLs.
nvidia_packages = [ nvidia_packages = [
'nvidia', 'nvidia.cublas', 'nvidia.cuda_cupti', 'nvidia.cuda_nvrtc', "nvidia",
'nvidia.cuda_runtime', 'nvidia.cudnn', 'nvidia.cufft', 'nvidia.curand', "nvidia.cublas",
'nvidia.cusolver', 'nvidia.cusparse', 'nvidia.nccl', 'nvidia.nvjitlink', "nvidia.cuda_cupti",
'nvidia.nvtx', "nvidia.cuda_nvrtc",
"nvidia.cuda_runtime",
"nvidia.cudnn",
"nvidia.cufft",
"nvidia.curand",
"nvidia.cusolver",
"nvidia.cusparse",
"nvidia.nccl",
"nvidia.nvjitlink",
"nvidia.nvtx",
] ]
for pkg in nvidia_packages: for pkg in nvidia_packages:
args.extend(['--exclude-module', pkg]) args.extend(["--exclude-module", pkg])
# Add MLX-specific imports if building on Apple Silicon (never for CUDA builds) # Add MLX-specific imports if building on Apple Silicon (never for CUDA builds)
if is_apple_silicon() and not cuda: if is_apple_silicon() and not cuda:
print("Building for Apple Silicon - including MLX dependencies") print("Building for Apple Silicon - including MLX dependencies")
args.extend([ args.extend(
'--hidden-import', 'backend.backends.mlx_backend', [
'--hidden-import', 'mlx', "--hidden-import",
'--hidden-import', 'mlx.core', "backend.backends.mlx_backend",
'--hidden-import', 'mlx.nn', "--hidden-import",
'--hidden-import', 'mlx_audio', "mlx",
'--hidden-import', 'mlx_audio.tts', "--hidden-import",
'--hidden-import', 'mlx_audio.stt', "mlx.core",
'--collect-submodules', 'mlx', "--hidden-import",
'--collect-submodules', 'mlx_audio', "mlx.nn",
# Use --collect-all so PyInstaller bundles both data files AND "--hidden-import",
# native shared libraries (.dylib, .metallib) for MLX. "mlx_audio",
# Previously only --collect-data was used, which caused MLX to "--hidden-import",
# raise OSError at runtime inside the bundled binary because "mlx_audio.tts",
# the Metal shader libraries were missing. "--hidden-import",
'--collect-all', 'mlx', "mlx_audio.stt",
'--collect-all', 'mlx_audio', "--collect-submodules",
]) "mlx",
"--collect-submodules",
"mlx_audio",
# Use --collect-all so PyInstaller bundles both data files AND
# native shared libraries (.dylib, .metallib) for MLX.
# Previously only --collect-data was used, which caused MLX to
# raise OSError at runtime inside the bundled binary because
# the Metal shader libraries were missing.
"--collect-all",
"mlx",
"--collect-all",
"mlx_audio",
]
)
elif not cuda: elif not cuda:
print("Building for non-Apple Silicon platform - PyTorch only") print("Building for non-Apple Silicon platform - PyTorch only")
dist_dir = str(backend_dir / 'dist') dist_dir = str(backend_dir / "dist")
build_dir = str(backend_dir / 'build') build_dir = str(backend_dir / "build")
args.extend([ args.extend(
'--distpath', dist_dir, [
'--workpath', build_dir, "--distpath",
'--noconfirm', dist_dir,
'--clean', "--workpath",
]) build_dir,
"--noconfirm",
"--clean",
]
)
# Change to backend directory # Change to backend directory
os.chdir(backend_dir) os.chdir(backend_dir)
@@ -184,17 +278,28 @@ def build_server(cuda=False):
restore_cuda = False restore_cuda = False
if not cuda and platform.system() == "Windows": if not cuda and platform.system() == "Windows":
import subprocess import subprocess
result = subprocess.run( result = subprocess.run(
[sys.executable, "-c", "import torch; print(torch.version.cuda or '')"], [sys.executable, "-c", "import torch; print(torch.version.cuda or '')"], capture_output=True, text=True
capture_output=True, text=True
) )
has_cuda_torch = bool(result.stdout.strip()) has_cuda_torch = bool(result.stdout.strip())
if has_cuda_torch: if has_cuda_torch:
print("CUDA torch detected — installing CPU torch for CPU build...") print("CUDA torch detected — installing CPU torch for CPU build...")
subprocess.run( subprocess.run(
[sys.executable, "-m", "pip", "install", "torch", "torchvision", "torchaudio", [
"--index-url", "https://download.pytorch.org/whl/cpu", "--force-reinstall", "-q"], sys.executable,
check=True "-m",
"pip",
"install",
"torch",
"torchvision",
"torchaudio",
"--index-url",
"https://download.pytorch.org/whl/cpu",
"--force-reinstall",
"-q",
],
check=True,
) )
restore_cuda = True restore_cuda = True
@@ -206,10 +311,22 @@ def build_server(cuda=False):
if restore_cuda: if restore_cuda:
print("Restoring CUDA torch...") print("Restoring CUDA torch...")
import subprocess import subprocess
subprocess.run( subprocess.run(
[sys.executable, "-m", "pip", "install", "torch", "torchvision", "torchaudio", [
"--index-url", "https://download.pytorch.org/whl/cu126", "--force-reinstall", "-q"], sys.executable,
check=True "-m",
"pip",
"install",
"torch",
"torchvision",
"torchaudio",
"--index-url",
"https://download.pytorch.org/whl/cu126",
"--force-reinstall",
"-q",
],
check=True,
) )
print(f"Binary built in {backend_dir / 'dist' / binary_name}") print(f"Binary built in {backend_dir / 'dist' / binary_name}")
@@ -223,38 +340,48 @@ def _get_cuda_dll_excludes():
""" """
try: try:
import torch import torch
torch_lib = Path(torch.__file__).parent / 'lib'
torch_lib = Path(torch.__file__).parent / "lib"
except ImportError: except ImportError:
return [] return []
cuda_prefixes = ( cuda_prefixes = (
'torch_cuda', 'cublas', 'cublasLt', 'cudnn', 'cusparse', 'cufft', "torch_cuda",
'cusolver', 'cusolverMg', 'curand', 'nvrtc', 'nvJitLink', 'nccl', "cublas",
'nvperf', 'nvrtc-builtins', "cublasLt",
"cudnn",
"cusparse",
"cufft",
"cusolver",
"cusolverMg",
"curand",
"nvrtc",
"nvJitLink",
"nccl",
"nvperf",
"nvrtc-builtins",
) )
exclude_dlls = [] exclude_dlls = []
if torch_lib.exists(): if torch_lib.exists():
for f in torch_lib.iterdir(): for f in torch_lib.iterdir():
if f.suffix == '.dll' and any(f.name.startswith(p) for p in cuda_prefixes): if f.suffix == ".dll" and any(f.name.startswith(p) for p in cuda_prefixes):
exclude_dlls.append(f.name) exclude_dlls.append(f.name)
if exclude_dlls: if exclude_dlls:
total_mb = sum( total_mb = (
(torch_lib / dll).stat().st_size sum((torch_lib / dll).stat().st_size for dll in exclude_dlls if (torch_lib / dll).exists()) / 1024 / 1024
for dll in exclude_dlls )
if (torch_lib / dll).exists()
) / 1024 / 1024
print(f"CPU build: will exclude {len(exclude_dlls)} CUDA DLLs ({total_mb:.0f} MB)") print(f"CPU build: will exclude {len(exclude_dlls)} CUDA DLLs ({total_mb:.0f} MB)")
return exclude_dlls return exclude_dlls
if __name__ == '__main__': if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Build voicebox-server binary") parser = argparse.ArgumentParser(description="Build voicebox-server binary")
parser.add_argument( parser.add_argument(
'--cuda', "--cuda",
action='store_true', action="store_true",
help="Build CUDA-enabled binary (voicebox-server-cuda)", help="Build CUDA-enabled binary (voicebox-server-cuda)",
) )
cli_args = parser.parse_args() cli_args = parser.parse_args()
+5 -2
View File
@@ -235,8 +235,11 @@ def _save_regenerate(
audio_path = config.get_generations_dir() / f"{generation_id}_{suffix}.wav" audio_path = config.get_generations_dir() / f"{generation_id}_{suffix}.wav"
save_audio(audio, str(audio_path), sample_rate) save_audio(audio, str(audio_path), sample_rate)
existing = versions_mod.list_versions(generation_id, db) # Count via DB query rather than list length to avoid TOCTOU race
label = f"take-{len(existing) + 1}" from ..database import GenerationVersion as DBGenerationVersion
count = db.query(DBGenerationVersion).filter_by(generation_id=generation_id).count()
label = f"take-{count + 1}"
versions_mod.create_version( versions_mod.create_version(
generation_id=generation_id, generation_id=generation_id,