diff --git a/app/src/components/Generation/FloatingGenerateBox.tsx b/app/src/components/Generation/FloatingGenerateBox.tsx index f0dc2b7e..a9814432 100644 --- a/app/src/components/Generation/FloatingGenerateBox.tsx +++ b/app/src/components/Generation/FloatingGenerateBox.tsx @@ -13,7 +13,7 @@ import { } from '@/components/ui/select'; import { Textarea } from '@/components/ui/textarea'; 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 { useProfile, useProfiles } from '@/lib/hooks/useProfiles'; import { useAddStoryItem, useStory } from '@/lib/hooks/useStories'; @@ -112,6 +112,13 @@ export function FloatingGenerateBox({ } }, [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) useEffect(() => { if (!isExpanded) { @@ -344,9 +351,7 @@ export function FloatingGenerateBox({ : 'bg-card border border-border hover:bg-background/50', )} aria-label={ - isInstructMode - ? 'Fine tune instructions, on' - : 'Fine tune instructions' + isInstructMode ? 'Fine tune instructions, on' : 'Fine tune instructions' } > diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index 49b1b924..617e61b0 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -21,6 +21,12 @@ from ..utils.progress import get_progress_manager from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback 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: """MLX-based TTS backend using mlx-audio.""" @@ -343,7 +349,8 @@ class MLXTTSBackend: # MLX generate() returns a generator yielding GenerationResult objects audio_chunks = [] sample_rate = 24000 - + lang = LANGUAGE_CODE_TO_NAME.get(language, "auto") + # Set seed if provided (MLX uses numpy random) if seed is not None: import mlx.core as mx @@ -371,23 +378,23 @@ class MLXTTSBackend: sig = inspect.signature(self.model.generate) if "ref_audio" in sig.parameters: # 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)) sample_rate = result.sample_rate else: # 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)) sample_rate = result.sample_rate else: # 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)) sample_rate = result.sample_rate except Exception as e: # If voice cloning fails, try without it 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)) sample_rate = result.sample_rate diff --git a/backend/backends/pytorch_backend.py b/backend/backends/pytorch_backend.py index 729053b4..517403c7 100644 --- a/backend/backends/pytorch_backend.py +++ b/backend/backends/pytorch_backend.py @@ -15,6 +15,12 @@ from ..utils.progress import get_progress_manager from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback 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: """PyTorch-based TTS backend using Qwen3-TTS.""" @@ -359,6 +365,7 @@ class PyTorchTTSBackend: wavs, sample_rate = self.model.generate_voice_clone( text=text, voice_clone_prompt=voice_prompt, + language=LANGUAGE_CODE_TO_NAME.get(language, "auto"), instruct=instruct, ) return wavs[0], sample_rate