diff --git a/app/src/components/Generation/FloatingGenerateBox.tsx b/app/src/components/Generation/FloatingGenerateBox.tsx index d4ab574d..8a8512f9 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 { LANGUAGE_OPTIONS } from '@/lib/constants/languages'; +import { getLanguageOptionsForEngine } from '@/lib/constants/languages'; import { useGenerationForm } from '@/lib/hooks/useGenerationForm'; import { useProfile, useProfiles } from '@/lib/hooks/useProfiles'; import { useAddStoryItem, useStory } from '@/lib/hooks/useStories'; @@ -381,25 +381,30 @@ export function FloatingGenerateBox({ ( - - - - - )} + render={({ field }) => { + const engineLangs = getLanguageOptionsForEngine( + form.watch('engine') || 'qwen', + ); + return ( + + + + + ); + }} /> @@ -409,13 +414,19 @@ export function FloatingGenerateBox({ ? 'luxtts' : form.watch('engine') === 'chatterbox' ? 'chatterbox' - : `qwen:${form.watch('modelSize') || '1.7B'}` + : form.watch('engine') === 'chatterbox_turbo' + ? 'chatterbox_turbo' + : `qwen:${form.watch('modelSize') || '1.7B'}` } onValueChange={(value) => { if (value === 'luxtts') { form.setValue('engine', 'luxtts'); + form.setValue('language', 'en'); } else if (value === 'chatterbox') { form.setValue('engine', 'chatterbox'); + } else if (value === 'chatterbox_turbo') { + form.setValue('engine', 'chatterbox_turbo'); + form.setValue('language', 'en'); } else { const [, modelSize] = value.split(':'); form.setValue('engine', 'qwen'); @@ -441,6 +452,12 @@ export function FloatingGenerateBox({ Chatterbox + + Chatterbox Turbo + diff --git a/app/src/components/Generation/GenerationForm.tsx b/app/src/components/Generation/GenerationForm.tsx index 26fd13e3..a3c96cbc 100644 --- a/app/src/components/Generation/GenerationForm.tsx +++ b/app/src/components/Generation/GenerationForm.tsx @@ -19,7 +19,7 @@ import { SelectValue, } from '@/components/ui/select'; import { Textarea } from '@/components/ui/textarea'; -import { LANGUAGE_OPTIONS } from '@/lib/constants/languages'; +import { getLanguageOptionsForEngine } from '@/lib/constants/languages'; import { useGenerationForm } from '@/lib/hooks/useGenerationForm'; import { useProfile } from '@/lib/hooks/useProfiles'; import { useUIStore } from '@/stores/uiStore'; @@ -109,13 +109,19 @@ export function GenerationForm() { ? 'luxtts' : form.watch('engine') === 'chatterbox' ? 'chatterbox' - : `qwen:${form.watch('modelSize') || '1.7B'}` + : form.watch('engine') === 'chatterbox_turbo' + ? 'chatterbox_turbo' + : `qwen:${form.watch('modelSize') || '1.7B'}` } onValueChange={(value) => { if (value === 'luxtts') { form.setValue('engine', 'luxtts'); + form.setValue('language', 'en'); } else if (value === 'chatterbox') { form.setValue('engine', 'chatterbox'); + } else if (value === 'chatterbox_turbo') { + form.setValue('engine', 'chatterbox_turbo'); + form.setValue('language', 'en'); } else { const [, modelSize] = value.split(':'); form.setValue('engine', 'qwen'); @@ -133,40 +139,46 @@ export function GenerationForm() { Qwen3-TTS 0.6B LuxTTS Chatterbox + Chatterbox Turbo {form.watch('engine') === 'luxtts' ? 'Fast, English-focused' : form.watch('engine') === 'chatterbox' - ? 'Multilingual, incl. Hebrew' - : 'Multi-language, two sizes'} + ? '23 languages, incl. Hebrew' + : form.watch('engine') === 'chatterbox_turbo' + ? 'English, [laugh] [cough] tags' + : 'Multi-language, two sizes'} ( - - Language - - - - )} + render={({ field }) => { + const engineLangs = getLanguageOptionsForEngine(form.watch('engine') || 'qwen'); + return ( + + Language + + + + ); + }} /> = { + qwen: ['zh', 'en', 'ja', 'ko', 'de', 'fr', 'ru', 'pt', 'es', 'it'], + luxtts: ['en'], + chatterbox: [ + 'ar', + 'da', + 'de', + 'el', + 'en', + 'es', + 'fi', + 'fr', + 'he', + 'hi', + 'it', + 'ja', + 'ko', + 'ms', + 'nl', + 'no', + 'pl', + 'pt', + 'ru', + 'sv', + 'sw', + 'tr', + 'zh', + ], + chatterbox_turbo: ['en'], +} as const; +/** Helper: get language options for a given engine. */ +export function getLanguageOptionsForEngine(engine: string) { + const codes = ENGINE_LANGUAGES[engine] ?? ENGINE_LANGUAGES.qwen; + return codes.map((code) => ({ + value: code, + label: ALL_LANGUAGES[code], + })); +} + +// ── Backwards-compatible exports used elsewhere ────────────────────── +export const SUPPORTED_LANGUAGES = ALL_LANGUAGES; +export const LANGUAGE_CODES = Object.keys(ALL_LANGUAGES) as LanguageCode[]; export const LANGUAGE_OPTIONS = LANGUAGE_CODES.map((code) => ({ value: code, - label: SUPPORTED_LANGUAGES[code], + label: ALL_LANGUAGES[code], })); diff --git a/app/src/lib/hooks/useGenerationForm.ts b/app/src/lib/hooks/useGenerationForm.ts index ec5b9d6a..5a83ce41 100644 --- a/app/src/lib/hooks/useGenerationForm.ts +++ b/app/src/lib/hooks/useGenerationForm.ts @@ -16,7 +16,7 @@ const generationSchema = z.object({ seed: z.number().int().optional(), modelSize: z.enum(['1.7B', '0.6B']).optional(), instruct: z.string().max(500).optional(), - engine: z.enum(['qwen', 'luxtts', 'chatterbox']).optional(), + engine: z.enum(['qwen', 'luxtts', 'chatterbox', 'chatterbox_turbo']).optional(), }); export type GenerationFormValues = z.infer; @@ -75,15 +75,19 @@ export function useGenerationForm(options: UseGenerationFormOptions = {}) { ? 'luxtts' : engine === 'chatterbox' ? 'chatterbox-tts' - : `qwen-tts-${data.modelSize}`; + : engine === 'chatterbox_turbo' + ? 'chatterbox-turbo' + : `qwen-tts-${data.modelSize}`; const displayName = engine === 'luxtts' ? 'LuxTTS' : engine === 'chatterbox' ? 'Chatterbox TTS' - : data.modelSize === '1.7B' - ? 'Qwen TTS 1.7B' - : 'Qwen TTS 0.6B'; + : engine === 'chatterbox_turbo' + ? 'Chatterbox Turbo' + : data.modelSize === '1.7B' + ? 'Qwen TTS 1.7B' + : 'Qwen TTS 0.6B'; try { const modelStatus = await apiClient.getModelStatus(); diff --git a/backend/backends/__init__.py b/backend/backends/__init__.py index a7b4d54c..f120e6ec 100644 --- a/backend/backends/__init__.py +++ b/backend/backends/__init__.py @@ -122,6 +122,7 @@ TTS_ENGINES = { "qwen": "Qwen TTS", "luxtts": "LuxTTS", "chatterbox": "Chatterbox TTS", + "chatterbox_turbo": "Chatterbox Turbo", } @@ -171,6 +172,9 @@ def get_tts_backend_for_engine(engine: str) -> TTSBackend: elif engine == "chatterbox": from .chatterbox_backend import ChatterboxTTSBackend backend = ChatterboxTTSBackend() + elif engine == "chatterbox_turbo": + from .chatterbox_turbo_backend import ChatterboxTurboTTSBackend + backend = ChatterboxTurboTTSBackend() else: raise ValueError(f"Unknown TTS engine: {engine}. Supported: {list(TTS_ENGINES.keys())}") diff --git a/backend/backends/chatterbox_turbo_backend.py b/backend/backends/chatterbox_turbo_backend.py new file mode 100644 index 00000000..16bb5d70 --- /dev/null +++ b/backend/backends/chatterbox_turbo_backend.py @@ -0,0 +1,307 @@ +""" +Chatterbox Turbo TTS backend implementation. + +Wraps ChatterboxTurboTTS from chatterbox-tts for fast, English-only +voice cloning with paralinguistic tag support ([laugh], [cough], etc.). +Forces CPU on macOS due to known MPS tensor issues. +""" + +import asyncio +import logging +import platform +import threading +from pathlib import Path +from typing import ClassVar, List, Optional, Tuple + +import numpy as np + +from . import TTSBackend +from ..utils.audio import normalize_audio, load_audio +from ..utils.progress import get_progress_manager +from ..utils.tasks import get_task_manager + +logger = logging.getLogger(__name__) + +CHATTERBOX_TURBO_HF_REPO = "ResembleAI/chatterbox-turbo" + +# Files that must be present for the turbo model +_TURBO_WEIGHT_FILES = [ + "t3_turbo_v1.safetensors", + "s3gen_meanflow.safetensors", + "ve.safetensors", +] + + +class ChatterboxTurboTTSBackend: + """Chatterbox Turbo TTS backend — fast, English-only, with paralinguistic tags.""" + + # Class-level lock for torch.load monkey-patching + _load_lock: ClassVar[threading.Lock] = threading.Lock() + + def __init__(self): + self.model = None + self.model_size = "default" + self._device = None + self._model_load_lock = asyncio.Lock() + + def _get_device(self) -> str: + """Get the best available device. Forces CPU on macOS (MPS issue).""" + if platform.system() == "Darwin": + return "cpu" + try: + import torch + + if torch.cuda.is_available(): + return "cuda" + except ImportError: + pass + return "cpu" + + def is_loaded(self) -> bool: + return self.model is not None + + def _get_model_path(self, model_size: str = "default") -> str: + return CHATTERBOX_TURBO_HF_REPO + + def _is_model_cached(self, model_size: str = "default") -> bool: + """Check if the Chatterbox Turbo model is cached locally.""" + try: + from huggingface_hub import constants as hf_constants + + repo_cache = Path(hf_constants.HF_HUB_CACHE) / ( + "models--" + CHATTERBOX_TURBO_HF_REPO.replace("/", "--") + ) + + if not repo_cache.exists(): + return False + + blobs_dir = repo_cache / "blobs" + if blobs_dir.exists() and any(blobs_dir.glob("*.incomplete")): + return False + + # Check for turbo weight files + snapshots_dir = repo_cache / "snapshots" + if snapshots_dir.exists(): + for fname in _TURBO_WEIGHT_FILES: + if not any(snapshots_dir.rglob(fname)): + return False + return True + + return False + except Exception as e: + logger.warning(f"Error checking Chatterbox Turbo cache: {e}") + return False + + async def load_model(self, model_size: str = "default") -> None: + """Load the Chatterbox Turbo model.""" + if self.model is not None: + return + async with self._model_load_lock: + if self.model is not None: + return + await asyncio.to_thread(self._load_model_sync) + + def _load_model_sync(self): + """Synchronous model loading.""" + from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback + + progress_manager = get_progress_manager() + task_manager = get_task_manager() + model_name = "chatterbox-turbo" + + is_cached = self._is_model_cached() + + # Set up HF progress tracking (intercepts tqdm for file-level progress) + progress_callback = create_hf_progress_callback(model_name, progress_manager) + tracker = HFProgressTracker(progress_callback, filter_non_downloads=is_cached) + tracker_context = tracker.patch_download() + tracker_context.__enter__() + + if not is_cached: + task_manager.start_download(model_name) + progress_manager.update_progress( + model_name=model_name, + current=0, + total=0, + filename="Connecting to HuggingFace...", + status="downloading", + ) + + try: + device = self._get_device() + self._device = device + + logger.info(f"Loading Chatterbox Turbo TTS on {device}...") + + import torch + from huggingface_hub import snapshot_download + from chatterbox.tts_turbo import ChatterboxTurboTTS + + # Download model files ourselves so we can pass token=None + # (upstream from_pretrained passes token=True which requires + # a stored HF token even though the repo is public). + try: + local_path = snapshot_download( + repo_id=CHATTERBOX_TURBO_HF_REPO, + token=None, + allow_patterns=[ + "*.safetensors", "*.json", "*.txt", "*.pt", "*.model", + ], + ) + finally: + tracker_context.__exit__(None, None, None) + + # Monkey-patch torch.load for CPU loading. The model's .pt files + # were saved on CUDA; from_local() doesn't pass map_location + # so loading on CPU fails without this. + if device == "cpu": + _orig_torch_load = torch.load + + def _patched_load(*args, **kwargs): + kwargs.setdefault("map_location", "cpu") + return _orig_torch_load(*args, **kwargs) + + with ChatterboxTurboTTSBackend._load_lock: + torch.load = _patched_load + try: + self.model = ChatterboxTurboTTS.from_local( + local_path, device, + ) + finally: + torch.load = _orig_torch_load + else: + self.model = ChatterboxTurboTTS.from_local( + local_path, device, + ) + + if not is_cached: + progress_manager.mark_complete(model_name) + task_manager.complete_download(model_name) + + logger.info("Chatterbox Turbo TTS loaded successfully") + + except ImportError as e: + logger.error( + "chatterbox-tts package not found. " + "Install with: pip install chatterbox-tts" + ) + if not is_cached: + progress_manager.mark_error(model_name, str(e)) + task_manager.error_download(model_name, str(e)) + raise + except Exception as e: + logger.error(f"Failed to load Chatterbox Turbo: {e}") + if not is_cached: + progress_manager.mark_error(model_name, str(e)) + task_manager.error_download(model_name, str(e)) + raise + + def unload_model(self) -> None: + """Unload model to free memory.""" + if self.model is not None: + device = self._device + del self.model + self.model = None + self._device = None + if device == "cuda": + import torch + + torch.cuda.empty_cache() + logger.info("Chatterbox Turbo unloaded") + + async def create_voice_prompt( + self, + audio_path: str, + reference_text: str, + use_cache: bool = True, + ) -> Tuple[dict, bool]: + """ + Create voice prompt from reference audio. + + Chatterbox Turbo processes reference audio at generation time, so the + prompt just stores the file path. + """ + voice_prompt = { + "ref_audio": str(audio_path), + "ref_text": reference_text, + } + return voice_prompt, False + + async def combine_voice_prompts( + self, + audio_paths: List[str], + reference_texts: List[str], + ) -> Tuple[np.ndarray, str]: + """Combine multiple reference samples.""" + combined_audio = [] + for path in audio_paths: + audio, _sr = load_audio(path) + audio = normalize_audio(audio) + combined_audio.append(audio) + + mixed = np.concatenate(combined_audio) + mixed = normalize_audio(mixed) + combined_text = " ".join(reference_texts) + return mixed, combined_text + + async def generate( + self, + text: str, + voice_prompt: dict, + language: str = "en", + seed: Optional[int] = None, + instruct: Optional[str] = None, + ) -> Tuple[np.ndarray, int]: + """ + Generate audio using Chatterbox Turbo TTS. + + Supports paralinguistic tags in text: [laugh], [cough], [chuckle], etc. + + Args: + text: Text to synthesize (may include paralinguistic tags) + voice_prompt: Dict with ref_audio path + language: Ignored (Turbo is English-only) + seed: Random seed for reproducibility + instruct: Unused (protocol compatibility) + + Returns: + Tuple of (audio_array, sample_rate) + """ + await self.load_model() + + ref_audio = voice_prompt.get("ref_audio") + if ref_audio and not Path(ref_audio).exists(): + logger.warning(f"Reference audio not found: {ref_audio}") + ref_audio = None + + def _generate_sync(): + import torch + + if seed is not None: + torch.manual_seed(seed) + + logger.info("[Chatterbox Turbo] Generating (English)") + + wav = self.model.generate( + text, + audio_prompt_path=ref_audio, + temperature=0.8, + top_k=1000, + top_p=0.95, + repetition_penalty=1.2, + ) + + # Convert tensor -> numpy + if isinstance(wav, torch.Tensor): + audio = wav.squeeze().cpu().numpy().astype(np.float32) + else: + audio = np.asarray(wav, dtype=np.float32) + + sample_rate = ( + getattr(self.model, "sr", None) + or getattr(self.model, "sample_rate", 24000) + ) + + return audio, sample_rate + + return await asyncio.to_thread(_generate_sync) diff --git a/backend/main.py b/backend/main.py index 3d2ec359..d0be9a54 100644 --- a/backend/main.py +++ b/backend/main.py @@ -699,6 +699,29 @@ async def generate_speech( ) await tts_model.load_model() + elif engine == "chatterbox_turbo": + if not tts_model._is_model_cached(): + model_name = "chatterbox-turbo" + + async def download_chatterbox_turbo_background(): + try: + await tts_model.load_model() + except Exception as e: + task_manager.error_download(model_name, str(e)) + + task_manager.start_download(model_name) + asyncio.create_task(download_chatterbox_turbo_background()) + + raise HTTPException( + status_code=202, + detail={ + "message": "Chatterbox Turbo model is being downloaded. Please wait and try again.", + "model_name": model_name, + "downloading": True, + }, + ) + + await tts_model.load_model() # Create voice prompt from profile voice_prompt = await profiles.create_voice_prompt_for_profile( @@ -717,7 +740,7 @@ async def generate_speech( ) # Trim trailing silence/hallucination for Chatterbox output - if engine == "chatterbox": + if engine in ("chatterbox", "chatterbox_turbo"): from .utils.audio import trim_tts_output audio = trim_tts_output(audio, sample_rate) @@ -798,6 +821,13 @@ async def stream_speech( detail="Chatterbox model is not downloaded yet. Use /generate to trigger a download.", ) await tts_model.load_model() + elif engine == "chatterbox_turbo": + if not tts_model._is_model_cached(): + raise HTTPException( + status_code=400, + detail="Chatterbox Turbo model is not downloaded yet. Use /generate to trigger a download.", + ) + await tts_model.load_model() voice_prompt = await profiles.create_voice_prompt_for_profile( data.profile_id, db, engine=engine, @@ -812,7 +842,7 @@ async def stream_speech( ) # Trim trailing silence/hallucination for Chatterbox output - if engine == "chatterbox": + if engine in ("chatterbox", "chatterbox_turbo"): from .utils.audio import trim_tts_output audio = trim_tts_output(audio, sample_rate) @@ -1433,6 +1463,15 @@ async def get_model_status(): except Exception: return False + # Check if Chatterbox Turbo backend is loaded + def check_chatterbox_turbo_loaded(): + try: + from .backends import get_tts_backend_for_engine + backend = get_tts_backend_for_engine("chatterbox_turbo") + return backend.is_loaded() + except Exception: + return False + model_configs = [ { "model_name": "qwen-tts-1.7B", @@ -1462,6 +1501,13 @@ async def get_model_status(): "model_size": "default", "check_loaded": check_chatterbox_loaded, }, + { + "model_name": "chatterbox-turbo", + "display_name": "Chatterbox Turbo (English, Tags)", + "hf_repo_id": "ResembleAI/chatterbox-turbo", + "model_size": "default", + "check_loaded": check_chatterbox_turbo_loaded, + }, { "model_name": "whisper-base", "display_name": "Whisper Base", @@ -1668,6 +1714,10 @@ async def trigger_model_download(request: models.ModelDownloadRequest): "model_size": "default", "load_func": lambda: get_tts_backend_for_engine("chatterbox").load_model(), }, + "chatterbox-turbo": { + "model_size": "default", + "load_func": lambda: get_tts_backend_for_engine("chatterbox_turbo").load_model(), + }, "whisper-base": { "model_size": "base", "load_func": lambda: transcribe.get_whisper_model().load_model("base"), @@ -1790,6 +1840,11 @@ async def delete_model(model_name: str): "model_size": "default", "model_type": "chatterbox", }, + "chatterbox-turbo": { + "hf_repo_id": "ResembleAI/chatterbox-turbo", + "model_size": "default", + "model_type": "chatterbox_turbo", + }, "whisper-base": { "hf_repo_id": "openai/whisper-base", "model_size": "base", @@ -1834,6 +1889,11 @@ async def delete_model(model_name: str): chatterbox = get_tts_backend_for_engine("chatterbox") if chatterbox.is_loaded(): chatterbox.unload_model() + elif config["model_type"] == "chatterbox_turbo": + from .backends import get_tts_backend_for_engine + turbo = get_tts_backend_for_engine("chatterbox_turbo") + if turbo.is_loaded(): + turbo.unload_model() elif config["model_type"] == "whisper": whisper_model = transcribe.get_whisper_model() if whisper_model.is_loaded() and whisper_model.model_size == config["model_size"]: diff --git a/backend/models.py b/backend/models.py index c46ded58..80d69495 100644 --- a/backend/models.py +++ b/backend/models.py @@ -11,7 +11,7 @@ class VoiceProfileCreate(BaseModel): """Request model for creating a voice profile.""" name: str = Field(..., min_length=1, max_length=100) description: Optional[str] = Field(None, max_length=500) - language: str = Field(default="en", pattern="^(zh|en|ja|ko|de|fr|ru|pt|es|it|he)$") + language: str = Field(default="en", pattern="^(zh|en|ja|ko|de|fr|ru|pt|es|it|he|ar|da|el|fi|hi|ms|nl|no|pl|sv|sw|tr)$") class VoiceProfileResponse(BaseModel): @@ -57,7 +57,7 @@ class GenerationRequest(BaseModel): seed: Optional[int] = Field(None, ge=0) model_size: Optional[str] = Field(default="1.7B", pattern="^(1\\.7B|0\\.6B)$") instruct: Optional[str] = Field(None, max_length=500) - engine: Optional[str] = Field(default="qwen", pattern="^(qwen|luxtts|chatterbox)$") + engine: Optional[str] = Field(default="qwen", pattern="^(qwen|luxtts|chatterbox|chatterbox_turbo)$") class GenerationResponse(BaseModel):