mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 09:05:17 -07:00
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:
@@ -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
@@ -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()
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
@@ -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()
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user