mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 09:05:17 -07:00
fix: pass language parameter to Qwen TTS models and sync form with profile language
Both PyTorch and MLX backends silently dropped the language parameter — it was accepted by generate() but never forwarded to the underlying Qwen3-TTS model, causing it to default to auto-detection which frequently confuses similar languages (e.g. Portuguese for Spanish). - Add LANGUAGE_CODE_TO_NAME mapping (ISO 639-1 to full name) to both backends - PyTorch: pass language= to generate_voice_clone() - MLX: pass lang_code= to all 4 model.generate() call sites - Frontend: auto-sync generation form language with selected voice profile Closes #97
This commit is contained in:
@@ -13,7 +13,7 @@ import {
|
|||||||
} from '@/components/ui/select';
|
} from '@/components/ui/select';
|
||||||
import { Textarea } from '@/components/ui/textarea';
|
import { Textarea } from '@/components/ui/textarea';
|
||||||
import { useToast } from '@/components/ui/use-toast';
|
import { useToast } from '@/components/ui/use-toast';
|
||||||
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, useProfiles } from '@/lib/hooks/useProfiles';
|
import { useProfile, useProfiles } from '@/lib/hooks/useProfiles';
|
||||||
import { useAddStoryItem, useStory } from '@/lib/hooks/useStories';
|
import { useAddStoryItem, useStory } from '@/lib/hooks/useStories';
|
||||||
@@ -112,6 +112,13 @@ export function FloatingGenerateBox({
|
|||||||
}
|
}
|
||||||
}, [selectedProfileId, profiles, setSelectedProfileId]);
|
}, [selectedProfileId, profiles, setSelectedProfileId]);
|
||||||
|
|
||||||
|
// Sync generation form language with selected profile's language
|
||||||
|
useEffect(() => {
|
||||||
|
if (selectedProfile?.language) {
|
||||||
|
form.setValue('language', selectedProfile.language as LanguageCode);
|
||||||
|
}
|
||||||
|
}, [selectedProfile, form]);
|
||||||
|
|
||||||
// Auto-resize textarea based on content (only when expanded)
|
// Auto-resize textarea based on content (only when expanded)
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
if (!isExpanded) {
|
if (!isExpanded) {
|
||||||
@@ -344,9 +351,7 @@ export function FloatingGenerateBox({
|
|||||||
: 'bg-card border border-border hover:bg-background/50',
|
: 'bg-card border border-border hover:bg-background/50',
|
||||||
)}
|
)}
|
||||||
aria-label={
|
aria-label={
|
||||||
isInstructMode
|
isInstructMode ? 'Fine tune instructions, on' : 'Fine tune instructions'
|
||||||
? 'Fine tune instructions, on'
|
|
||||||
: 'Fine tune instructions'
|
|
||||||
}
|
}
|
||||||
>
|
>
|
||||||
<SlidersHorizontal className="h-4 w-4" />
|
<SlidersHorizontal className="h-4 w-4" />
|
||||||
|
|||||||
@@ -21,6 +21,12 @@ from ..utils.progress import get_progress_manager
|
|||||||
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
||||||
from ..utils.tasks import get_task_manager
|
from ..utils.tasks import get_task_manager
|
||||||
|
|
||||||
|
LANGUAGE_CODE_TO_NAME = {
|
||||||
|
"zh": "chinese", "en": "english", "ja": "japanese", "ko": "korean",
|
||||||
|
"de": "german", "fr": "french", "ru": "russian", "pt": "portuguese",
|
||||||
|
"es": "spanish", "it": "italian",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
class MLXTTSBackend:
|
class MLXTTSBackend:
|
||||||
"""MLX-based TTS backend using mlx-audio."""
|
"""MLX-based TTS backend using mlx-audio."""
|
||||||
@@ -343,6 +349,7 @@ class MLXTTSBackend:
|
|||||||
# MLX generate() returns a generator yielding GenerationResult objects
|
# MLX generate() returns a generator yielding GenerationResult objects
|
||||||
audio_chunks = []
|
audio_chunks = []
|
||||||
sample_rate = 24000
|
sample_rate = 24000
|
||||||
|
lang = LANGUAGE_CODE_TO_NAME.get(language, "auto")
|
||||||
|
|
||||||
# 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:
|
||||||
@@ -371,23 +378,23 @@ class MLXTTSBackend:
|
|||||||
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
|
||||||
for result in self.model.generate(text, ref_audio=ref_audio, ref_text=ref_text):
|
for result in self.model.generate(text, ref_audio=ref_audio, ref_text=ref_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
|
||||||
else:
|
else:
|
||||||
# Fallback: generate without voice cloning
|
# Fallback: generate without voice cloning
|
||||||
for result in self.model.generate(text):
|
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
|
||||||
else:
|
else:
|
||||||
# No voice prompt, generate normally
|
# No voice prompt, generate normally
|
||||||
for result in self.model.generate(text):
|
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
|
||||||
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}")
|
print(f"Warning: Voice cloning failed, generating without voice prompt: {e}")
|
||||||
for result in self.model.generate(text):
|
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
|
||||||
|
|
||||||
|
|||||||
@@ -15,6 +15,12 @@ from ..utils.progress import get_progress_manager
|
|||||||
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
||||||
from ..utils.tasks import get_task_manager
|
from ..utils.tasks import get_task_manager
|
||||||
|
|
||||||
|
LANGUAGE_CODE_TO_NAME = {
|
||||||
|
"zh": "chinese", "en": "english", "ja": "japanese", "ko": "korean",
|
||||||
|
"de": "german", "fr": "french", "ru": "russian", "pt": "portuguese",
|
||||||
|
"es": "spanish", "it": "italian",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
class PyTorchTTSBackend:
|
class PyTorchTTSBackend:
|
||||||
"""PyTorch-based TTS backend using Qwen3-TTS."""
|
"""PyTorch-based TTS backend using Qwen3-TTS."""
|
||||||
@@ -359,6 +365,7 @@ class PyTorchTTSBackend:
|
|||||||
wavs, sample_rate = self.model.generate_voice_clone(
|
wavs, sample_rate = self.model.generate_voice_clone(
|
||||||
text=text,
|
text=text,
|
||||||
voice_clone_prompt=voice_prompt,
|
voice_clone_prompt=voice_prompt,
|
||||||
|
language=LANGUAGE_CODE_TO_NAME.get(language, "auto"),
|
||||||
instruct=instruct,
|
instruct=instruct,
|
||||||
)
|
)
|
||||||
return wavs[0], sample_rate
|
return wavs[0], sample_rate
|
||||||
|
|||||||
Reference in New Issue
Block a user