mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 09:05:17 -07:00
feat: add Chatterbox Turbo engine and per-engine language lists
- New ChatterboxTurboTTSBackend wrapping ChatterboxTurboTTS (ResembleAI/chatterbox-turbo) - English-only 350M model with paralinguistic tag support ([laugh], [cough], [chuckle]) - Bypasses upstream token=True bug by calling snapshot_download(token=None) + from_local() - Same CPU-on-macOS forcing and torch.load monkey-patching as multilingual backend - Full engine integration: generate, stream, model status/download/delete endpoints - Language dropdown now shows only languages supported by the selected engine - Per-engine language maps: Qwen (10), LuxTTS (en), Chatterbox (23), Turbo (en) - Auto-switches to English when selecting English-only engines - Backend language regex expanded to accept all 23 Chatterbox languages
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 { LANGUAGE_OPTIONS } from '@/lib/constants/languages';
|
import { getLanguageOptionsForEngine } 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';
|
||||||
@@ -381,25 +381,30 @@ export function FloatingGenerateBox({
|
|||||||
<FormField
|
<FormField
|
||||||
control={form.control}
|
control={form.control}
|
||||||
name="language"
|
name="language"
|
||||||
render={({ field }) => (
|
render={({ field }) => {
|
||||||
<FormItem className="flex-1 space-y-0">
|
const engineLangs = getLanguageOptionsForEngine(
|
||||||
<Select onValueChange={field.onChange} defaultValue={field.value}>
|
form.watch('engine') || 'qwen',
|
||||||
<FormControl>
|
);
|
||||||
<SelectTrigger className="h-8 text-xs bg-card border-border rounded-full hover:bg-background/50 transition-all">
|
return (
|
||||||
<SelectValue />
|
<FormItem className="flex-1 space-y-0">
|
||||||
</SelectTrigger>
|
<Select onValueChange={field.onChange} value={field.value}>
|
||||||
</FormControl>
|
<FormControl>
|
||||||
<SelectContent>
|
<SelectTrigger className="h-8 text-xs bg-card border-border rounded-full hover:bg-background/50 transition-all">
|
||||||
{LANGUAGE_OPTIONS.map((lang) => (
|
<SelectValue />
|
||||||
<SelectItem key={lang.value} value={lang.value} className="text-xs">
|
</SelectTrigger>
|
||||||
{lang.label}
|
</FormControl>
|
||||||
</SelectItem>
|
<SelectContent>
|
||||||
))}
|
{engineLangs.map((lang) => (
|
||||||
</SelectContent>
|
<SelectItem key={lang.value} value={lang.value} className="text-xs">
|
||||||
</Select>
|
{lang.label}
|
||||||
<FormMessage className="text-xs" />
|
</SelectItem>
|
||||||
</FormItem>
|
))}
|
||||||
)}
|
</SelectContent>
|
||||||
|
</Select>
|
||||||
|
<FormMessage className="text-xs" />
|
||||||
|
</FormItem>
|
||||||
|
);
|
||||||
|
}}
|
||||||
/>
|
/>
|
||||||
|
|
||||||
<FormItem className="flex-1 space-y-0">
|
<FormItem className="flex-1 space-y-0">
|
||||||
@@ -409,13 +414,19 @@ export function FloatingGenerateBox({
|
|||||||
? 'luxtts'
|
? 'luxtts'
|
||||||
: form.watch('engine') === 'chatterbox'
|
: form.watch('engine') === 'chatterbox'
|
||||||
? 'chatterbox'
|
? 'chatterbox'
|
||||||
: `qwen:${form.watch('modelSize') || '1.7B'}`
|
: form.watch('engine') === 'chatterbox_turbo'
|
||||||
|
? 'chatterbox_turbo'
|
||||||
|
: `qwen:${form.watch('modelSize') || '1.7B'}`
|
||||||
}
|
}
|
||||||
onValueChange={(value) => {
|
onValueChange={(value) => {
|
||||||
if (value === 'luxtts') {
|
if (value === 'luxtts') {
|
||||||
form.setValue('engine', 'luxtts');
|
form.setValue('engine', 'luxtts');
|
||||||
|
form.setValue('language', 'en');
|
||||||
} else if (value === 'chatterbox') {
|
} else if (value === 'chatterbox') {
|
||||||
form.setValue('engine', 'chatterbox');
|
form.setValue('engine', 'chatterbox');
|
||||||
|
} else if (value === 'chatterbox_turbo') {
|
||||||
|
form.setValue('engine', 'chatterbox_turbo');
|
||||||
|
form.setValue('language', 'en');
|
||||||
} else {
|
} else {
|
||||||
const [, modelSize] = value.split(':');
|
const [, modelSize] = value.split(':');
|
||||||
form.setValue('engine', 'qwen');
|
form.setValue('engine', 'qwen');
|
||||||
@@ -441,6 +452,12 @@ export function FloatingGenerateBox({
|
|||||||
<SelectItem value="chatterbox" className="text-xs text-muted-foreground">
|
<SelectItem value="chatterbox" className="text-xs text-muted-foreground">
|
||||||
Chatterbox
|
Chatterbox
|
||||||
</SelectItem>
|
</SelectItem>
|
||||||
|
<SelectItem
|
||||||
|
value="chatterbox_turbo"
|
||||||
|
className="text-xs text-muted-foreground"
|
||||||
|
>
|
||||||
|
Chatterbox Turbo
|
||||||
|
</SelectItem>
|
||||||
</SelectContent>
|
</SelectContent>
|
||||||
</Select>
|
</Select>
|
||||||
</FormItem>
|
</FormItem>
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ import {
|
|||||||
SelectValue,
|
SelectValue,
|
||||||
} from '@/components/ui/select';
|
} from '@/components/ui/select';
|
||||||
import { Textarea } from '@/components/ui/textarea';
|
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 { 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';
|
||||||
@@ -109,13 +109,19 @@ export function GenerationForm() {
|
|||||||
? 'luxtts'
|
? 'luxtts'
|
||||||
: form.watch('engine') === 'chatterbox'
|
: form.watch('engine') === 'chatterbox'
|
||||||
? 'chatterbox'
|
? 'chatterbox'
|
||||||
: `qwen:${form.watch('modelSize') || '1.7B'}`
|
: form.watch('engine') === 'chatterbox_turbo'
|
||||||
|
? 'chatterbox_turbo'
|
||||||
|
: `qwen:${form.watch('modelSize') || '1.7B'}`
|
||||||
}
|
}
|
||||||
onValueChange={(value) => {
|
onValueChange={(value) => {
|
||||||
if (value === 'luxtts') {
|
if (value === 'luxtts') {
|
||||||
form.setValue('engine', 'luxtts');
|
form.setValue('engine', 'luxtts');
|
||||||
|
form.setValue('language', 'en');
|
||||||
} else if (value === 'chatterbox') {
|
} else if (value === 'chatterbox') {
|
||||||
form.setValue('engine', 'chatterbox');
|
form.setValue('engine', 'chatterbox');
|
||||||
|
} else if (value === 'chatterbox_turbo') {
|
||||||
|
form.setValue('engine', 'chatterbox_turbo');
|
||||||
|
form.setValue('language', 'en');
|
||||||
} else {
|
} else {
|
||||||
const [, modelSize] = value.split(':');
|
const [, modelSize] = value.split(':');
|
||||||
form.setValue('engine', 'qwen');
|
form.setValue('engine', 'qwen');
|
||||||
@@ -133,40 +139,46 @@ export function GenerationForm() {
|
|||||||
<SelectItem value="qwen:0.6B">Qwen3-TTS 0.6B</SelectItem>
|
<SelectItem value="qwen:0.6B">Qwen3-TTS 0.6B</SelectItem>
|
||||||
<SelectItem value="luxtts">LuxTTS</SelectItem>
|
<SelectItem value="luxtts">LuxTTS</SelectItem>
|
||||||
<SelectItem value="chatterbox">Chatterbox</SelectItem>
|
<SelectItem value="chatterbox">Chatterbox</SelectItem>
|
||||||
|
<SelectItem value="chatterbox_turbo">Chatterbox Turbo</SelectItem>
|
||||||
</SelectContent>
|
</SelectContent>
|
||||||
</Select>
|
</Select>
|
||||||
<FormDescription>
|
<FormDescription>
|
||||||
{form.watch('engine') === 'luxtts'
|
{form.watch('engine') === 'luxtts'
|
||||||
? 'Fast, English-focused'
|
? 'Fast, English-focused'
|
||||||
: form.watch('engine') === 'chatterbox'
|
: form.watch('engine') === 'chatterbox'
|
||||||
? 'Multilingual, incl. Hebrew'
|
? '23 languages, incl. Hebrew'
|
||||||
: 'Multi-language, two sizes'}
|
: form.watch('engine') === 'chatterbox_turbo'
|
||||||
|
? 'English, [laugh] [cough] tags'
|
||||||
|
: 'Multi-language, two sizes'}
|
||||||
</FormDescription>
|
</FormDescription>
|
||||||
</FormItem>
|
</FormItem>
|
||||||
|
|
||||||
<FormField
|
<FormField
|
||||||
control={form.control}
|
control={form.control}
|
||||||
name="language"
|
name="language"
|
||||||
render={({ field }) => (
|
render={({ field }) => {
|
||||||
<FormItem>
|
const engineLangs = getLanguageOptionsForEngine(form.watch('engine') || 'qwen');
|
||||||
<FormLabel>Language</FormLabel>
|
return (
|
||||||
<Select onValueChange={field.onChange} defaultValue={field.value}>
|
<FormItem>
|
||||||
<FormControl>
|
<FormLabel>Language</FormLabel>
|
||||||
<SelectTrigger>
|
<Select onValueChange={field.onChange} value={field.value}>
|
||||||
<SelectValue />
|
<FormControl>
|
||||||
</SelectTrigger>
|
<SelectTrigger>
|
||||||
</FormControl>
|
<SelectValue />
|
||||||
<SelectContent>
|
</SelectTrigger>
|
||||||
{LANGUAGE_OPTIONS.map((lang) => (
|
</FormControl>
|
||||||
<SelectItem key={lang.value} value={lang.value}>
|
<SelectContent>
|
||||||
{lang.label}
|
{engineLangs.map((lang) => (
|
||||||
</SelectItem>
|
<SelectItem key={lang.value} value={lang.value}>
|
||||||
))}
|
{lang.label}
|
||||||
</SelectContent>
|
</SelectItem>
|
||||||
</Select>
|
))}
|
||||||
<FormMessage />
|
</SelectContent>
|
||||||
</FormItem>
|
</Select>
|
||||||
)}
|
<FormMessage />
|
||||||
|
</FormItem>
|
||||||
|
);
|
||||||
|
}}
|
||||||
/>
|
/>
|
||||||
|
|
||||||
<FormField
|
<FormField
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ export interface GenerationRequest {
|
|||||||
language: LanguageCode;
|
language: LanguageCode;
|
||||||
seed?: number;
|
seed?: number;
|
||||||
model_size?: '1.7B' | '0.6B';
|
model_size?: '1.7B' | '0.6B';
|
||||||
engine?: 'qwen' | 'luxtts' | 'chatterbox';
|
engine?: 'qwen' | 'luxtts' | 'chatterbox' | 'chatterbox_turbo';
|
||||||
instruct?: string;
|
instruct?: string;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,27 +1,86 @@
|
|||||||
/**
|
/**
|
||||||
* Supported languages for voice generation.
|
* Supported languages for voice generation, per engine.
|
||||||
* Most languages use Qwen3-TTS; Hebrew uses Chatterbox TTS.
|
*
|
||||||
|
* Qwen3-TTS supports 10 languages.
|
||||||
|
* LuxTTS is English-only.
|
||||||
|
* Chatterbox Multilingual supports 23 languages.
|
||||||
|
* Chatterbox Turbo is English-only.
|
||||||
*/
|
*/
|
||||||
|
|
||||||
export const SUPPORTED_LANGUAGES = {
|
/** All languages that any engine supports. */
|
||||||
zh: 'Chinese',
|
export const ALL_LANGUAGES = {
|
||||||
|
ar: 'Arabic',
|
||||||
|
da: 'Danish',
|
||||||
|
de: 'German',
|
||||||
|
el: 'Greek',
|
||||||
en: 'English',
|
en: 'English',
|
||||||
|
es: 'Spanish',
|
||||||
|
fi: 'Finnish',
|
||||||
|
fr: 'French',
|
||||||
|
he: 'Hebrew',
|
||||||
|
hi: 'Hindi',
|
||||||
|
it: 'Italian',
|
||||||
ja: 'Japanese',
|
ja: 'Japanese',
|
||||||
ko: 'Korean',
|
ko: 'Korean',
|
||||||
de: 'German',
|
ms: 'Malay',
|
||||||
fr: 'French',
|
nl: 'Dutch',
|
||||||
ru: 'Russian',
|
no: 'Norwegian',
|
||||||
|
pl: 'Polish',
|
||||||
pt: 'Portuguese',
|
pt: 'Portuguese',
|
||||||
es: 'Spanish',
|
ru: 'Russian',
|
||||||
it: 'Italian',
|
sv: 'Swedish',
|
||||||
he: 'Hebrew',
|
sw: 'Swahili',
|
||||||
|
tr: 'Turkish',
|
||||||
|
zh: 'Chinese',
|
||||||
} as const;
|
} as const;
|
||||||
|
|
||||||
export type LanguageCode = keyof typeof SUPPORTED_LANGUAGES;
|
export type LanguageCode = keyof typeof ALL_LANGUAGES;
|
||||||
|
|
||||||
export const LANGUAGE_CODES = Object.keys(SUPPORTED_LANGUAGES) as LanguageCode[];
|
/** Per-engine supported language codes. */
|
||||||
|
export const ENGINE_LANGUAGES: Record<string, readonly LanguageCode[]> = {
|
||||||
|
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) => ({
|
export const LANGUAGE_OPTIONS = LANGUAGE_CODES.map((code) => ({
|
||||||
value: code,
|
value: code,
|
||||||
label: SUPPORTED_LANGUAGES[code],
|
label: ALL_LANGUAGES[code],
|
||||||
}));
|
}));
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ const generationSchema = z.object({
|
|||||||
seed: z.number().int().optional(),
|
seed: z.number().int().optional(),
|
||||||
modelSize: z.enum(['1.7B', '0.6B']).optional(),
|
modelSize: z.enum(['1.7B', '0.6B']).optional(),
|
||||||
instruct: z.string().max(500).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<typeof generationSchema>;
|
export type GenerationFormValues = z.infer<typeof generationSchema>;
|
||||||
@@ -75,15 +75,19 @@ export function useGenerationForm(options: UseGenerationFormOptions = {}) {
|
|||||||
? 'luxtts'
|
? 'luxtts'
|
||||||
: engine === 'chatterbox'
|
: engine === 'chatterbox'
|
||||||
? 'chatterbox-tts'
|
? 'chatterbox-tts'
|
||||||
: `qwen-tts-${data.modelSize}`;
|
: engine === 'chatterbox_turbo'
|
||||||
|
? 'chatterbox-turbo'
|
||||||
|
: `qwen-tts-${data.modelSize}`;
|
||||||
const displayName =
|
const displayName =
|
||||||
engine === 'luxtts'
|
engine === 'luxtts'
|
||||||
? 'LuxTTS'
|
? 'LuxTTS'
|
||||||
: engine === 'chatterbox'
|
: engine === 'chatterbox'
|
||||||
? 'Chatterbox TTS'
|
? 'Chatterbox TTS'
|
||||||
: data.modelSize === '1.7B'
|
: engine === 'chatterbox_turbo'
|
||||||
? 'Qwen TTS 1.7B'
|
? 'Chatterbox Turbo'
|
||||||
: 'Qwen TTS 0.6B';
|
: data.modelSize === '1.7B'
|
||||||
|
? 'Qwen TTS 1.7B'
|
||||||
|
: 'Qwen TTS 0.6B';
|
||||||
|
|
||||||
try {
|
try {
|
||||||
const modelStatus = await apiClient.getModelStatus();
|
const modelStatus = await apiClient.getModelStatus();
|
||||||
|
|||||||
@@ -122,6 +122,7 @@ TTS_ENGINES = {
|
|||||||
"qwen": "Qwen TTS",
|
"qwen": "Qwen TTS",
|
||||||
"luxtts": "LuxTTS",
|
"luxtts": "LuxTTS",
|
||||||
"chatterbox": "Chatterbox TTS",
|
"chatterbox": "Chatterbox TTS",
|
||||||
|
"chatterbox_turbo": "Chatterbox Turbo",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -171,6 +172,9 @@ def get_tts_backend_for_engine(engine: str) -> TTSBackend:
|
|||||||
elif engine == "chatterbox":
|
elif engine == "chatterbox":
|
||||||
from .chatterbox_backend import ChatterboxTTSBackend
|
from .chatterbox_backend import ChatterboxTTSBackend
|
||||||
backend = ChatterboxTTSBackend()
|
backend = ChatterboxTTSBackend()
|
||||||
|
elif engine == "chatterbox_turbo":
|
||||||
|
from .chatterbox_turbo_backend import ChatterboxTurboTTSBackend
|
||||||
|
backend = ChatterboxTurboTTSBackend()
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unknown TTS engine: {engine}. Supported: {list(TTS_ENGINES.keys())}")
|
raise ValueError(f"Unknown TTS engine: {engine}. Supported: {list(TTS_ENGINES.keys())}")
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
+62
-2
@@ -699,6 +699,29 @@ async def generate_speech(
|
|||||||
)
|
)
|
||||||
|
|
||||||
await tts_model.load_model()
|
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
|
# Create voice prompt from profile
|
||||||
voice_prompt = await profiles.create_voice_prompt_for_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
|
# Trim trailing silence/hallucination for Chatterbox output
|
||||||
if engine == "chatterbox":
|
if engine in ("chatterbox", "chatterbox_turbo"):
|
||||||
from .utils.audio import trim_tts_output
|
from .utils.audio import trim_tts_output
|
||||||
audio = trim_tts_output(audio, sample_rate)
|
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.",
|
detail="Chatterbox model is not downloaded yet. Use /generate to trigger a download.",
|
||||||
)
|
)
|
||||||
await tts_model.load_model()
|
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(
|
voice_prompt = await profiles.create_voice_prompt_for_profile(
|
||||||
data.profile_id, db, engine=engine,
|
data.profile_id, db, engine=engine,
|
||||||
@@ -812,7 +842,7 @@ async def stream_speech(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Trim trailing silence/hallucination for Chatterbox output
|
# Trim trailing silence/hallucination for Chatterbox output
|
||||||
if engine == "chatterbox":
|
if engine in ("chatterbox", "chatterbox_turbo"):
|
||||||
from .utils.audio import trim_tts_output
|
from .utils.audio import trim_tts_output
|
||||||
audio = trim_tts_output(audio, sample_rate)
|
audio = trim_tts_output(audio, sample_rate)
|
||||||
|
|
||||||
@@ -1433,6 +1463,15 @@ async def get_model_status():
|
|||||||
except Exception:
|
except Exception:
|
||||||
return False
|
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_configs = [
|
||||||
{
|
{
|
||||||
"model_name": "qwen-tts-1.7B",
|
"model_name": "qwen-tts-1.7B",
|
||||||
@@ -1462,6 +1501,13 @@ async def get_model_status():
|
|||||||
"model_size": "default",
|
"model_size": "default",
|
||||||
"check_loaded": check_chatterbox_loaded,
|
"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",
|
"model_name": "whisper-base",
|
||||||
"display_name": "Whisper Base",
|
"display_name": "Whisper Base",
|
||||||
@@ -1668,6 +1714,10 @@ async def trigger_model_download(request: models.ModelDownloadRequest):
|
|||||||
"model_size": "default",
|
"model_size": "default",
|
||||||
"load_func": lambda: get_tts_backend_for_engine("chatterbox").load_model(),
|
"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": {
|
"whisper-base": {
|
||||||
"model_size": "base",
|
"model_size": "base",
|
||||||
"load_func": lambda: transcribe.get_whisper_model().load_model("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_size": "default",
|
||||||
"model_type": "chatterbox",
|
"model_type": "chatterbox",
|
||||||
},
|
},
|
||||||
|
"chatterbox-turbo": {
|
||||||
|
"hf_repo_id": "ResembleAI/chatterbox-turbo",
|
||||||
|
"model_size": "default",
|
||||||
|
"model_type": "chatterbox_turbo",
|
||||||
|
},
|
||||||
"whisper-base": {
|
"whisper-base": {
|
||||||
"hf_repo_id": "openai/whisper-base",
|
"hf_repo_id": "openai/whisper-base",
|
||||||
"model_size": "base",
|
"model_size": "base",
|
||||||
@@ -1834,6 +1889,11 @@ async def delete_model(model_name: str):
|
|||||||
chatterbox = get_tts_backend_for_engine("chatterbox")
|
chatterbox = get_tts_backend_for_engine("chatterbox")
|
||||||
if chatterbox.is_loaded():
|
if chatterbox.is_loaded():
|
||||||
chatterbox.unload_model()
|
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":
|
elif config["model_type"] == "whisper":
|
||||||
whisper_model = transcribe.get_whisper_model()
|
whisper_model = transcribe.get_whisper_model()
|
||||||
if whisper_model.is_loaded() and whisper_model.model_size == config["model_size"]:
|
if whisper_model.is_loaded() and whisper_model.model_size == config["model_size"]:
|
||||||
|
|||||||
+2
-2
@@ -11,7 +11,7 @@ class VoiceProfileCreate(BaseModel):
|
|||||||
"""Request model for creating a voice profile."""
|
"""Request model for creating a voice profile."""
|
||||||
name: str = Field(..., min_length=1, max_length=100)
|
name: str = Field(..., min_length=1, max_length=100)
|
||||||
description: Optional[str] = Field(None, max_length=500)
|
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):
|
class VoiceProfileResponse(BaseModel):
|
||||||
@@ -57,7 +57,7 @@ class GenerationRequest(BaseModel):
|
|||||||
seed: Optional[int] = Field(None, ge=0)
|
seed: Optional[int] = Field(None, ge=0)
|
||||||
model_size: Optional[str] = Field(default="1.7B", pattern="^(1\\.7B|0\\.6B)$")
|
model_size: Optional[str] = Field(default="1.7B", pattern="^(1\\.7B|0\\.6B)$")
|
||||||
instruct: Optional[str] = Field(None, max_length=500)
|
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):
|
class GenerationResponse(BaseModel):
|
||||||
|
|||||||
Reference in New Issue
Block a user