fix: enforce preset profile engine compatibility

This commit is contained in:
James Pine
2026-03-19 19:51:53 -07:00
parent 4e0c731db8
commit 72c13fd3fc
10 changed files with 102 additions and 150 deletions
@@ -57,7 +57,7 @@ function getSelectValue(engine: string, modelSize?: string): string {
return engine; return engine;
} }
function handleEngineChange(form: UseFormReturn<GenerationFormValues>, value: string) { export function applyEngineSelection(form: UseFormReturn<GenerationFormValues>, value: string) {
if (value.startsWith('qwen_custom_voice:')) { if (value.startsWith('qwen_custom_voice:')) {
const [, modelSize] = value.split(':'); const [, modelSize] = value.split(':');
form.setValue('engine', 'qwen_custom_voice'); form.setValue('engine', 'qwen_custom_voice');
@@ -123,7 +123,7 @@ export function EngineModelSelector({ form, compact, selectedProfile }: EngineMo
useEffect(() => { useEffect(() => {
if (!currentEngineAvailable && availableOptions.length > 0) { if (!currentEngineAvailable && availableOptions.length > 0) {
handleEngineChange(form, availableOptions[0].value); applyEngineSelection(form, availableOptions[0].value);
} }
}, [availableOptions, currentEngineAvailable, form]); }, [availableOptions, currentEngineAvailable, form]);
@@ -133,7 +133,7 @@ export function EngineModelSelector({ form, compact, selectedProfile }: EngineMo
: undefined; : undefined;
return ( return (
<Select value={selectValue} onValueChange={(v) => handleEngineChange(form, v)}> <Select value={selectValue} onValueChange={(v) => applyEngineSelection(form, v)}>
<FormControl> <FormControl>
<SelectTrigger className={triggerClass}> <SelectTrigger className={triggerClass}>
<SelectValue /> <SelectValue />
@@ -1,3 +1,4 @@
import { useEffect } from 'react';
import { Loader2, Mic } from 'lucide-react'; import { Loader2, Mic } from 'lucide-react';
import { Button } from '@/components/ui/button'; import { Button } from '@/components/ui/button';
import { Card, CardContent, CardHeader, CardTitle } from '@/components/ui/card'; import { Card, CardContent, CardHeader, CardTitle } from '@/components/ui/card';
@@ -19,19 +20,41 @@ import {
SelectValue, SelectValue,
} from '@/components/ui/select'; } from '@/components/ui/select';
import { Textarea } from '@/components/ui/textarea'; 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 { useGenerationForm } from '@/lib/hooks/useGenerationForm';
import { useProfile } from '@/lib/hooks/useProfiles'; import { useProfile } from '@/lib/hooks/useProfiles';
import { useUIStore } from '@/stores/uiStore'; import { useUIStore } from '@/stores/uiStore';
import { EngineModelSelector, getEngineDescription } from './EngineModelSelector'; import { EngineModelSelector, applyEngineSelection, getEngineDescription } from './EngineModelSelector';
import { ParalinguisticInput } from './ParalinguisticInput'; 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() { export function GenerationForm() {
const selectedProfileId = useUIStore((state) => state.selectedProfileId); const selectedProfileId = useUIStore((state) => state.selectedProfileId);
const { data: selectedProfile } = useProfile(selectedProfileId || ''); const { data: selectedProfile } = useProfile(selectedProfileId || '');
const { form, handleSubmit, isPending } = useGenerationForm(); 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]) { async function onSubmit(data: Parameters<typeof handleSubmit>[0]) {
await handleSubmit(data, selectedProfileId); await handleSubmit(data, selectedProfileId);
} }
@@ -375,6 +375,15 @@ export function ProfileForm() {
} }
}, [availableDefaultEngines, defaultEngine]); }, [availableDefaultEngines, defaultEngine]);
useEffect(() => {
if (!selectedPresetVoiceId) {
return;
}
if (!presetVoices.some((voice: PresetVoice) => voice.voice_id === selectedPresetVoiceId)) {
setSelectedPresetVoiceId('');
}
}, [presetVoices, selectedPresetVoiceId]);
async function handleTranscribe() { async function handleTranscribe() {
const file = form.getValues('sampleFile'); const file = form.getValues('sampleFile');
if (!file) { if (!file) {
-6
View File
@@ -102,12 +102,6 @@ class ApiClient {
return this.request<{ engine: string; voices: PresetVoice[] }>(`/profiles/presets/${engine}`); 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> { async updateProfile(profileId: string, data: VoiceProfileCreate): Promise<VoiceProfileResponse> {
return this.request<VoiceProfileResponse>(`/profiles/${profileId}`, { return this.request<VoiceProfileResponse>(`/profiles/${profileId}`, {
method: 'PUT', method: 'PUT',
+5 -6
View File
@@ -194,6 +194,8 @@ def _resolve_relative_paths(engine, tables: set[str]) -> None:
configured data directory. If the path starts with "data/", strip that configured data directory. If the path starts with "data/", strip that
prefix and prepend get_data_dir(). Otherwise, join the relative path prefix and prepend get_data_dir(). Otherwise, join the relative path
directly under get_data_dir(). 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 from pathlib import Path
@@ -222,16 +224,13 @@ def _resolve_relative_paths(engine, tables: set[str]) -> None:
p = Path(path_val) p = Path(path_val)
if p.is_absolute(): if p.is_absolute():
continue continue
# Try rebasing: "data/generations/abc.wav" → data_dir / "generations/abc.wav"
parts = p.parts parts = p.parts
if parts and parts[0] == "data": if parts and parts[0] == "data":
rebased = data_dir / Path(*parts[1:]) rebased = (data_dir / Path(*parts[1:])).resolve()
else: else:
rebased = data_dir / p rebased = (data_dir / p).resolve()
resolved = rebased.resolve()
resolved = rebased if rebased.exists() else p.resolve()
if resolved.exists(): if resolved.exists():
conn.execute( conn.execute(
text(f"UPDATE {table} SET {column} = :path WHERE id = :id"), text(f"UPDATE {table} SET {column} = :path WHERE id = :id"),
+1 -1
View File
@@ -42,7 +42,7 @@ torchaudio
# Kokoro TTS (lightweight 82M-param engine) # Kokoro TTS (lightweight 82M-param engine)
kokoro>=0.9.4 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 # spacy model for misaki English G2P — must be pre-installed or misaki
# tries spacy.cli.download() at runtime which crashes frozen builds # 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 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
+15 -10
View File
@@ -20,6 +20,10 @@ from ..utils.tasks import get_task_manager
router = APIRouter() 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) @router.post("/generate", response_model=models.GenerationResponse)
async def generate_speech( async def generate_speech(
data: models.GenerationRequest, data: models.GenerationRequest,
@@ -35,7 +39,12 @@ async def generate_speech(
from ..backends import engine_has_model_sizes 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 model_size = (data.model_size or "1.7B") if engine_has_model_sizes(engine) else None
generation = await history.create_generation( generation = await history.create_generation(
@@ -230,15 +239,11 @@ async def stream_speech(
if not profile: if not profile:
raise HTTPException(status_code=404, detail="Profile not found") raise HTTPException(status_code=404, detail="Profile not found")
# Mirror the regular /generate endpoint behavior more closely: engine = _resolve_generation_engine(data, profile)
# if the caller doesn't specify an engine, prefer the profile's default try:
# engine (or preset engine) before falling back to qwen. profiles.validate_profile_engine(profile, engine)
engine = ( except ValueError as e:
data.engine raise HTTPException(status_code=400, detail=str(e))
or getattr(profile, "default_engine", None)
or getattr(profile, "preset_engine", None)
or "qwen"
)
tts_model = get_tts_backend_for_engine(engine) tts_model = get_tts_backend_for_engine(engine)
model_size = data.model_size or "1.7B" model_size = data.model_size or "1.7B"
-121
View File
@@ -4,8 +4,6 @@ import io
import json as _json import json as _json
import logging import logging
import tempfile import tempfile
import uuid
from datetime import datetime
from pathlib import Path from pathlib import Path
from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile
@@ -107,125 +105,6 @@ async def list_preset_voices(engine: str):
} }
return {"engine": engine, "voices": []} 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) @router.get("/profiles/{profile_id}", response_model=models.VoiceProfileResponse)
async def get_profile( async def get_profile(
profile_id: str, profile_id: str,
+43
View File
@@ -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( def _validate_profile_fields(
*, *,
voice_type: str, voice_type: str,
@@ -74,6 +88,10 @@ def _validate_profile_fields(
return "Preset profiles require both preset_engine and preset_voice_id" return "Preset profiles require both preset_engine and preset_voice_id"
if default_engine and default_engine != preset_engine: if default_engine and default_engine != preset_engine:
return "Preset profiles must use their preset_engine as default_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 return None
if voice_type == "designed": if voice_type == "designed":
@@ -92,6 +110,30 @@ def _validate_profile_fields(
return None 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( async def create_profile(
data: VoiceProfileCreate, data: VoiceProfileCreate,
db: Session, db: Session,
@@ -476,6 +518,7 @@ async def create_voice_prompt_for_profile(
raise ValueError(f"Profile not found: {profile_id}") raise ValueError(f"Profile not found: {profile_id}")
voice_type = getattr(profile, "voice_type", None) or "cloned" voice_type = getattr(profile, "voice_type", None) or "cloned"
validate_profile_engine(profile, engine)
# ── Preset profiles: return engine-specific voice reference ── # ── Preset profiles: return engine-specific voice reference ──
if voice_type == "preset": if voice_type == "preset":
+1 -1
View File
@@ -5,7 +5,7 @@ from PyInstaller.utils.hooks import copy_metadata
datas = [] datas = []
binaries = [] 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('qwen-tts')
datas += copy_metadata('requests') datas += copy_metadata('requests')
datas += copy_metadata('transformers') datas += copy_metadata('transformers')