mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-12 19:30:38 -07:00
fix: enforce preset profile engine compatibility
This commit is contained in:
@@ -57,7 +57,7 @@ function getSelectValue(engine: string, modelSize?: string): string {
|
||||
return engine;
|
||||
}
|
||||
|
||||
function handleEngineChange(form: UseFormReturn<GenerationFormValues>, value: string) {
|
||||
export function applyEngineSelection(form: UseFormReturn<GenerationFormValues>, value: string) {
|
||||
if (value.startsWith('qwen_custom_voice:')) {
|
||||
const [, modelSize] = value.split(':');
|
||||
form.setValue('engine', 'qwen_custom_voice');
|
||||
@@ -123,7 +123,7 @@ export function EngineModelSelector({ form, compact, selectedProfile }: EngineMo
|
||||
|
||||
useEffect(() => {
|
||||
if (!currentEngineAvailable && availableOptions.length > 0) {
|
||||
handleEngineChange(form, availableOptions[0].value);
|
||||
applyEngineSelection(form, availableOptions[0].value);
|
||||
}
|
||||
}, [availableOptions, currentEngineAvailable, form]);
|
||||
|
||||
@@ -133,7 +133,7 @@ export function EngineModelSelector({ form, compact, selectedProfile }: EngineMo
|
||||
: undefined;
|
||||
|
||||
return (
|
||||
<Select value={selectValue} onValueChange={(v) => handleEngineChange(form, v)}>
|
||||
<Select value={selectValue} onValueChange={(v) => applyEngineSelection(form, v)}>
|
||||
<FormControl>
|
||||
<SelectTrigger className={triggerClass}>
|
||||
<SelectValue />
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import { useEffect } from 'react';
|
||||
import { Loader2, Mic } from 'lucide-react';
|
||||
import { Button } from '@/components/ui/button';
|
||||
import { Card, CardContent, CardHeader, CardTitle } from '@/components/ui/card';
|
||||
@@ -19,19 +20,41 @@ import {
|
||||
SelectValue,
|
||||
} from '@/components/ui/select';
|
||||
import { Textarea } from '@/components/ui/textarea';
|
||||
import { getLanguageOptionsForEngine } from '@/lib/constants/languages';
|
||||
import { getLanguageOptionsForEngine, type LanguageCode } from '@/lib/constants/languages';
|
||||
import { useGenerationForm } from '@/lib/hooks/useGenerationForm';
|
||||
import { useProfile } from '@/lib/hooks/useProfiles';
|
||||
import { useUIStore } from '@/stores/uiStore';
|
||||
import { EngineModelSelector, getEngineDescription } from './EngineModelSelector';
|
||||
import { EngineModelSelector, applyEngineSelection, getEngineDescription } from './EngineModelSelector';
|
||||
import { ParalinguisticInput } from './ParalinguisticInput';
|
||||
|
||||
function getEngineSelectValue(engine: string): string {
|
||||
if (engine === 'qwen') return 'qwen:1.7B';
|
||||
if (engine === 'qwen_custom_voice') return 'qwen_custom_voice:1.7B';
|
||||
if (engine === 'tada') return 'tada:1B';
|
||||
return engine;
|
||||
}
|
||||
|
||||
export function GenerationForm() {
|
||||
const selectedProfileId = useUIStore((state) => state.selectedProfileId);
|
||||
const { data: selectedProfile } = useProfile(selectedProfileId || '');
|
||||
|
||||
const { form, handleSubmit, isPending } = useGenerationForm();
|
||||
|
||||
useEffect(() => {
|
||||
if (!selectedProfile) {
|
||||
return;
|
||||
}
|
||||
|
||||
if (selectedProfile.language) {
|
||||
form.setValue('language', selectedProfile.language as LanguageCode);
|
||||
}
|
||||
|
||||
const preferredEngine = selectedProfile.default_engine || selectedProfile.preset_engine;
|
||||
if (preferredEngine) {
|
||||
applyEngineSelection(form, getEngineSelectValue(preferredEngine));
|
||||
}
|
||||
}, [form, selectedProfile]);
|
||||
|
||||
async function onSubmit(data: Parameters<typeof handleSubmit>[0]) {
|
||||
await handleSubmit(data, selectedProfileId);
|
||||
}
|
||||
|
||||
@@ -375,6 +375,15 @@ export function ProfileForm() {
|
||||
}
|
||||
}, [availableDefaultEngines, defaultEngine]);
|
||||
|
||||
useEffect(() => {
|
||||
if (!selectedPresetVoiceId) {
|
||||
return;
|
||||
}
|
||||
|
||||
if (!presetVoices.some((voice: PresetVoice) => voice.voice_id === selectedPresetVoiceId)) {
|
||||
setSelectedPresetVoiceId('');
|
||||
}
|
||||
}, [presetVoices, selectedPresetVoiceId]);
|
||||
async function handleTranscribe() {
|
||||
const file = form.getValues('sampleFile');
|
||||
if (!file) {
|
||||
|
||||
@@ -102,12 +102,6 @@ class ApiClient {
|
||||
return this.request<{ engine: string; voices: PresetVoice[] }>(`/profiles/presets/${engine}`);
|
||||
}
|
||||
|
||||
async seedPresetProfiles(
|
||||
engine: string,
|
||||
): Promise<{ engine: string; created: number; total_available: number }> {
|
||||
return this.request(`/profiles/presets/${engine}/seed`, { method: 'POST' });
|
||||
}
|
||||
|
||||
async updateProfile(profileId: string, data: VoiceProfileCreate): Promise<VoiceProfileResponse> {
|
||||
return this.request<VoiceProfileResponse>(`/profiles/${profileId}`, {
|
||||
method: 'PUT',
|
||||
|
||||
@@ -194,6 +194,8 @@ def _resolve_relative_paths(engine, tables: set[str]) -> None:
|
||||
configured data directory. If the path starts with "data/", strip that
|
||||
prefix and prepend get_data_dir(). Otherwise, join the relative path
|
||||
directly under get_data_dir().
|
||||
directly under get_data_dir(). If the rebased path still does not exist,
|
||||
fall back to resolving relative to CWD.
|
||||
"""
|
||||
from pathlib import Path
|
||||
|
||||
@@ -222,16 +224,13 @@ def _resolve_relative_paths(engine, tables: set[str]) -> None:
|
||||
p = Path(path_val)
|
||||
if p.is_absolute():
|
||||
continue
|
||||
|
||||
# Try rebasing: "data/generations/abc.wav" → data_dir / "generations/abc.wav"
|
||||
parts = p.parts
|
||||
if parts and parts[0] == "data":
|
||||
rebased = data_dir / Path(*parts[1:])
|
||||
rebased = (data_dir / Path(*parts[1:])).resolve()
|
||||
else:
|
||||
rebased = data_dir / p
|
||||
|
||||
resolved = rebased.resolve()
|
||||
rebased = (data_dir / p).resolve()
|
||||
|
||||
resolved = rebased if rebased.exists() else p.resolve()
|
||||
if resolved.exists():
|
||||
conn.execute(
|
||||
text(f"UPDATE {table} SET {column} = :path WHERE id = :id"),
|
||||
|
||||
@@ -42,7 +42,7 @@ torchaudio
|
||||
|
||||
# Kokoro TTS (lightweight 82M-param engine)
|
||||
kokoro>=0.9.4
|
||||
misaki[en]>=0.9.4
|
||||
misaki[en,ja,zh]>=0.9.4
|
||||
# spacy model for misaki English G2P — must be pre-installed or misaki
|
||||
# tries spacy.cli.download() at runtime which crashes frozen builds
|
||||
en_core_web_sm @ https://github.com/explosion/spacy-models/releases/download/en_core_web_sm-3.8.0/en_core_web_sm-3.8.0-py3-none-any.whl
|
||||
|
||||
@@ -20,6 +20,10 @@ from ..utils.tasks import get_task_manager
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _resolve_generation_engine(data: models.GenerationRequest, profile) -> str:
|
||||
return data.engine or getattr(profile, "default_engine", None) or getattr(profile, "preset_engine", None) or "qwen"
|
||||
|
||||
|
||||
@router.post("/generate", response_model=models.GenerationResponse)
|
||||
async def generate_speech(
|
||||
data: models.GenerationRequest,
|
||||
@@ -35,7 +39,12 @@ async def generate_speech(
|
||||
|
||||
from ..backends import engine_has_model_sizes
|
||||
|
||||
engine = data.engine or "qwen"
|
||||
engine = _resolve_generation_engine(data, profile)
|
||||
try:
|
||||
profiles.validate_profile_engine(profile, engine)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
model_size = (data.model_size or "1.7B") if engine_has_model_sizes(engine) else None
|
||||
|
||||
generation = await history.create_generation(
|
||||
@@ -230,15 +239,11 @@ async def stream_speech(
|
||||
if not profile:
|
||||
raise HTTPException(status_code=404, detail="Profile not found")
|
||||
|
||||
# Mirror the regular /generate endpoint behavior more closely:
|
||||
# if the caller doesn't specify an engine, prefer the profile's default
|
||||
# engine (or preset engine) before falling back to qwen.
|
||||
engine = (
|
||||
data.engine
|
||||
or getattr(profile, "default_engine", None)
|
||||
or getattr(profile, "preset_engine", None)
|
||||
or "qwen"
|
||||
)
|
||||
engine = _resolve_generation_engine(data, profile)
|
||||
try:
|
||||
profiles.validate_profile_engine(profile, engine)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
tts_model = get_tts_backend_for_engine(engine)
|
||||
model_size = data.model_size or "1.7B"
|
||||
|
||||
|
||||
@@ -4,8 +4,6 @@ import io
|
||||
import json as _json
|
||||
import logging
|
||||
import tempfile
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile
|
||||
@@ -107,125 +105,6 @@ async def list_preset_voices(engine: str):
|
||||
}
|
||||
return {"engine": engine, "voices": []}
|
||||
|
||||
|
||||
@router.post("/profiles/presets/{engine}/seed")
|
||||
async def seed_preset_profiles_route(
|
||||
engine: str,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""Seed preset voice profiles for an engine.
|
||||
|
||||
Creates profiles for all available preset voices that don't already exist.
|
||||
Returns the count of newly created profiles.
|
||||
"""
|
||||
if engine == "kokoro":
|
||||
return _seed_kokoro_presets(db)
|
||||
if engine == "qwen_custom_voice":
|
||||
return _seed_qwen_custom_voice_presets(db)
|
||||
raise HTTPException(status_code=400, detail=f"No presets available for engine: {engine}")
|
||||
|
||||
|
||||
def _seed_kokoro_presets(db: Session):
|
||||
"""Seed Kokoro preset profiles."""
|
||||
try:
|
||||
from ..backends.kokoro_backend import KOKORO_VOICES
|
||||
|
||||
created = 0
|
||||
for voice_id, display_name, gender, lang in KOKORO_VOICES:
|
||||
profile_name = display_name
|
||||
|
||||
# Disambiguate duplicate display names across languages
|
||||
# (e.g. "Alpha" exists in Hindi and Japanese, "Dora" in Spanish and Portuguese)
|
||||
dupes = [v for v in KOKORO_VOICES if v[1] == display_name]
|
||||
if len(dupes) > 1:
|
||||
lang_labels = {"en": "English", "es": "Spanish", "fr": "French", "hi": "Hindi",
|
||||
"it": "Italian", "pt": "Portuguese", "ja": "Japanese", "zh": "Chinese"}
|
||||
profile_name = f"{display_name} {lang_labels.get(lang, lang)}"
|
||||
|
||||
# Skip if preset already exists
|
||||
existing = (
|
||||
db.query(DBVoiceProfile)
|
||||
.filter_by(preset_engine="kokoro", preset_voice_id=voice_id)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
continue
|
||||
|
||||
unique_name = profile_name
|
||||
suffix = 2
|
||||
while db.query(DBVoiceProfile).filter_by(name=unique_name).first():
|
||||
unique_name = f"{profile_name} {suffix}"
|
||||
suffix += 1
|
||||
|
||||
profile = DBVoiceProfile(
|
||||
id=str(uuid.uuid4()),
|
||||
name=unique_name,
|
||||
description=f"Kokoro preset voice — {display_name} ({gender})",
|
||||
language=lang,
|
||||
voice_type="preset",
|
||||
preset_engine="kokoro",
|
||||
preset_voice_id=voice_id,
|
||||
created_at=datetime.utcnow(),
|
||||
updated_at=datetime.utcnow(),
|
||||
)
|
||||
db.add(profile)
|
||||
created += 1
|
||||
|
||||
if created > 0:
|
||||
db.commit()
|
||||
logger.info(f"Seeded {created} Kokoro preset profiles")
|
||||
|
||||
return {"engine": "kokoro", "created": created, "total_available": len(KOKORO_VOICES)}
|
||||
except Exception as e:
|
||||
logger.exception(f"Failed to seed Kokoro profiles: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
def _seed_qwen_custom_voice_presets(db: Session):
|
||||
"""Seed Qwen CustomVoice preset profiles."""
|
||||
try:
|
||||
from ..backends.qwen_custom_voice_backend import QWEN_CUSTOM_VOICES
|
||||
|
||||
created = 0
|
||||
for speaker_id, display_name, gender, lang, description in QWEN_CUSTOM_VOICES:
|
||||
# Skip if preset already exists
|
||||
existing = (
|
||||
db.query(DBVoiceProfile)
|
||||
.filter_by(preset_engine="qwen_custom_voice", preset_voice_id=speaker_id)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
continue
|
||||
|
||||
# Skip name collisions
|
||||
if db.query(DBVoiceProfile).filter_by(name=display_name).first():
|
||||
continue
|
||||
|
||||
profile = DBVoiceProfile(
|
||||
id=str(uuid.uuid4()),
|
||||
name=display_name,
|
||||
description=f"Qwen CustomVoice — {description}",
|
||||
language=lang,
|
||||
voice_type="preset",
|
||||
preset_engine="qwen_custom_voice",
|
||||
preset_voice_id=speaker_id,
|
||||
default_engine="qwen_custom_voice",
|
||||
created_at=datetime.utcnow(),
|
||||
updated_at=datetime.utcnow(),
|
||||
)
|
||||
db.add(profile)
|
||||
created += 1
|
||||
|
||||
if created > 0:
|
||||
db.commit()
|
||||
logger.info(f"Seeded {created} Qwen CustomVoice preset profiles")
|
||||
|
||||
return {"engine": "qwen_custom_voice", "created": created, "total_available": len(QWEN_CUSTOM_VOICES)}
|
||||
except Exception as e:
|
||||
logger.exception(f"Failed to seed Qwen CustomVoice profiles: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.get("/profiles/{profile_id}", response_model=models.VoiceProfileResponse)
|
||||
async def get_profile(
|
||||
profile_id: str,
|
||||
|
||||
@@ -61,6 +61,20 @@ def _profile_to_response(
|
||||
)
|
||||
|
||||
|
||||
def _get_preset_voice_ids(engine: str) -> set[str]:
|
||||
if engine == "kokoro":
|
||||
from ..backends.kokoro_backend import KOKORO_VOICES
|
||||
|
||||
return {voice_id for voice_id, _name, _gender, _lang in KOKORO_VOICES}
|
||||
|
||||
if engine == "qwen_custom_voice":
|
||||
from ..backends.qwen_custom_voice_backend import QWEN_CUSTOM_VOICES
|
||||
|
||||
return {voice_id for voice_id, _name, _gender, _lang, _desc in QWEN_CUSTOM_VOICES}
|
||||
|
||||
return set()
|
||||
|
||||
|
||||
def _validate_profile_fields(
|
||||
*,
|
||||
voice_type: str,
|
||||
@@ -74,6 +88,10 @@ def _validate_profile_fields(
|
||||
return "Preset profiles require both preset_engine and preset_voice_id"
|
||||
if default_engine and default_engine != preset_engine:
|
||||
return "Preset profiles must use their preset_engine as default_engine"
|
||||
|
||||
available_voice_ids = _get_preset_voice_ids(preset_engine)
|
||||
if available_voice_ids and preset_voice_id not in available_voice_ids:
|
||||
return f"Preset voice '{preset_voice_id}' is not valid for engine '{preset_engine}'"
|
||||
return None
|
||||
|
||||
if voice_type == "designed":
|
||||
@@ -92,6 +110,30 @@ def _validate_profile_fields(
|
||||
return None
|
||||
|
||||
|
||||
def validate_profile_engine(profile, engine: str) -> None:
|
||||
voice_type = getattr(profile, "voice_type", None) or "cloned"
|
||||
|
||||
if voice_type == "preset":
|
||||
preset_engine = getattr(profile, "preset_engine", None)
|
||||
preset_voice_id = getattr(profile, "preset_voice_id", None)
|
||||
if not preset_engine or not preset_voice_id:
|
||||
raise ValueError(f"Preset profile {profile.id} is missing preset engine metadata")
|
||||
if preset_engine != engine:
|
||||
raise ValueError(
|
||||
f"Preset profile {profile.id} only supports engine '{preset_engine}', not '{engine}'"
|
||||
)
|
||||
return
|
||||
|
||||
if voice_type == "designed":
|
||||
design_prompt = getattr(profile, "design_prompt", None)
|
||||
if not design_prompt or not design_prompt.strip():
|
||||
raise ValueError(f"Designed profile {profile.id} is missing design_prompt")
|
||||
return
|
||||
|
||||
if engine not in CLONING_ENGINES:
|
||||
raise ValueError(f"Engine '{engine}' does not support cloned voice profiles")
|
||||
|
||||
|
||||
async def create_profile(
|
||||
data: VoiceProfileCreate,
|
||||
db: Session,
|
||||
@@ -476,6 +518,7 @@ async def create_voice_prompt_for_profile(
|
||||
raise ValueError(f"Profile not found: {profile_id}")
|
||||
|
||||
voice_type = getattr(profile, "voice_type", None) or "cloned"
|
||||
validate_profile_engine(profile, engine)
|
||||
|
||||
# ── Preset profiles: return engine-specific voice reference ──
|
||||
if voice_type == "preset":
|
||||
|
||||
@@ -5,7 +5,7 @@ from PyInstaller.utils.hooks import copy_metadata
|
||||
|
||||
datas = []
|
||||
binaries = []
|
||||
hiddenimports = ['backend', 'backend.main', 'backend.config', 'backend.database', 'backend.models', 'backend.services.profiles', 'backend.services.history', 'backend.services.tts', 'backend.services.transcribe', 'backend.utils.platform_detect', 'backend.backends', 'backend.backends.pytorch_backend', 'backend.utils.audio', 'backend.utils.cache', 'backend.utils.progress', 'backend.utils.hf_progress', 'backend.services.cuda', 'backend.services.effects', 'backend.utils.effects', 'backend.services.versions', 'pedalboard', 'chatterbox', 'chatterbox.tts_turbo', 'chatterbox.mtl_tts', 'backend.backends.chatterbox_backend', 'backend.backends.chatterbox_turbo_backend', 'backend.backends.luxtts_backend', 'zipvoice', 'zipvoice.luxvoice', 'torch', 'transformers', 'fastapi', 'uvicorn', 'sqlalchemy', 'soundfile', 'qwen_tts', 'qwen_tts.inference', 'qwen_tts.inference.qwen3_tts_model', 'qwen_tts.inference.qwen3_tts_tokenizer', 'qwen_tts.core', 'qwen_tts.cli', 'requests', 'pkg_resources.extern', 'backend.backends.hume_backend', 'tada', 'tada.modules', 'tada.modules.tada', 'tada.modules.encoder', 'tada.modules.decoder', 'tada.modules.aligner', 'tada.modules.acoustic_spkr_verf', 'tada.nn', 'tada.nn.vibevoice', 'tada.utils', 'tada.utils.gray_code', 'tada.utils.text', 'backend.utils.dac_shim', 'torchaudio', 'backend.backends.kokoro_backend', 'kokoro', 'kokoro.pipeline', 'kokoro.model', 'kokoro.istftnet', 'kokoro.modules', 'kokoro.custom_stft', 'en_core_web_sm', 'loguru', 'backend.backends.mlx_backend', 'mlx', 'mlx.core', 'mlx.nn', 'mlx_audio', 'mlx_audio.tts', 'mlx_audio.stt']
|
||||
hiddenimports = ['backend', 'backend.main', 'backend.config', 'backend.database', 'backend.models', 'backend.services.profiles', 'backend.services.history', 'backend.services.tts', 'backend.services.transcribe', 'backend.utils.platform_detect', 'backend.backends', 'backend.backends.pytorch_backend', 'backend.backends.qwen_custom_voice_backend', 'backend.utils.audio', 'backend.utils.cache', 'backend.utils.progress', 'backend.utils.hf_progress', 'backend.services.cuda', 'backend.services.effects', 'backend.utils.effects', 'backend.services.versions', 'pedalboard', 'chatterbox', 'chatterbox.tts_turbo', 'chatterbox.mtl_tts', 'backend.backends.chatterbox_backend', 'backend.backends.chatterbox_turbo_backend', 'backend.backends.luxtts_backend', 'zipvoice', 'zipvoice.luxvoice', 'torch', 'transformers', 'fastapi', 'uvicorn', 'sqlalchemy', 'soundfile', 'qwen_tts', 'qwen_tts.inference', 'qwen_tts.inference.qwen3_tts_model', 'qwen_tts.inference.qwen3_tts_tokenizer', 'qwen_tts.core', 'qwen_tts.cli', 'requests', 'pkg_resources.extern', 'backend.backends.hume_backend', 'tada', 'tada.modules', 'tada.modules.tada', 'tada.modules.encoder', 'tada.modules.decoder', 'tada.modules.aligner', 'tada.modules.acoustic_spkr_verf', 'tada.nn', 'tada.nn.vibevoice', 'tada.utils', 'tada.utils.gray_code', 'tada.utils.text', 'backend.utils.dac_shim', 'torchaudio', 'backend.backends.kokoro_backend', 'kokoro', 'kokoro.pipeline', 'kokoro.model', 'kokoro.istftnet', 'kokoro.modules', 'kokoro.custom_stft', 'en_core_web_sm', 'loguru', 'backend.backends.mlx_backend', 'mlx', 'mlx.core', 'mlx.nn', 'mlx_audio', 'mlx_audio.tts', 'mlx_audio.stt']
|
||||
datas += copy_metadata('qwen-tts')
|
||||
datas += copy_metadata('requests')
|
||||
datas += copy_metadata('transformers')
|
||||
|
||||
Reference in New Issue
Block a user