mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-19 06:40:38 -07:00
feat: add LuxTTS as second TTS engine with multi-engine support
Introduce LuxTTS (ZipVoice) alongside Qwen TTS, enabling users to choose between engines at generation time. LuxTTS offers fast, English-focused voice cloning at 48kHz with ~1GB VRAM. Backend: - Add LuxTTSBackend with encode_prompt/generate_speech integration - Multi-engine registry (get_tts_backend_for_engine) replacing singleton - Engine-prefixed voice prompt cache keys to avoid collisions - Engine field on GenerationRequest (default 'qwen' for backward compat) - Engine dispatch in /generate and /generate/stream endpoints - LuxTTS in model status, download, and delete maps Frontend: - TTS Engine selector dropdown in GenerationForm (Qwen TTS / LuxTTS) - Conditionally hide Model Size and Delivery Instructions for LuxTTS - Engine field added to TypeScript types and Zod schema - LuxTTS section in Model Management page
This commit is contained in:
@@ -76,29 +76,56 @@ export function GenerationForm() {
|
|||||||
)}
|
)}
|
||||||
/>
|
/>
|
||||||
|
|
||||||
<FormField
|
{form.watch('engine') !== 'luxtts' && (
|
||||||
control={form.control}
|
<FormField
|
||||||
name="instruct"
|
control={form.control}
|
||||||
render={({ field }) => (
|
name="instruct"
|
||||||
<FormItem>
|
render={({ field }) => (
|
||||||
<FormLabel>Delivery Instructions (optional)</FormLabel>
|
<FormItem>
|
||||||
<FormControl>
|
<FormLabel>Delivery Instructions (optional)</FormLabel>
|
||||||
<Textarea
|
<FormControl>
|
||||||
placeholder="e.g. Speak slowly with emphasis, Warm and friendly tone, Professional and authoritative..."
|
<Textarea
|
||||||
className="min-h-[80px]"
|
placeholder="e.g. Speak slowly with emphasis, Warm and friendly tone, Professional and authoritative..."
|
||||||
{...field}
|
className="min-h-[80px]"
|
||||||
/>
|
{...field}
|
||||||
</FormControl>
|
/>
|
||||||
<FormDescription>
|
</FormControl>
|
||||||
Natural language instructions to control speech delivery (tone, emotion, pace).
|
<FormDescription>
|
||||||
Max 500 characters
|
Natural language instructions to control speech delivery (tone, emotion,
|
||||||
</FormDescription>
|
pace). Max 500 characters
|
||||||
<FormMessage />
|
</FormDescription>
|
||||||
</FormItem>
|
<FormMessage />
|
||||||
)}
|
</FormItem>
|
||||||
/>
|
)}
|
||||||
|
/>
|
||||||
|
)}
|
||||||
|
|
||||||
<div className="grid gap-4 md:grid-cols-3">
|
<div className="grid gap-4 md:grid-cols-3">
|
||||||
|
<FormField
|
||||||
|
control={form.control}
|
||||||
|
name="engine"
|
||||||
|
render={({ field }) => (
|
||||||
|
<FormItem>
|
||||||
|
<FormLabel>TTS Engine</FormLabel>
|
||||||
|
<Select onValueChange={field.onChange} defaultValue={field.value}>
|
||||||
|
<FormControl>
|
||||||
|
<SelectTrigger>
|
||||||
|
<SelectValue />
|
||||||
|
</SelectTrigger>
|
||||||
|
</FormControl>
|
||||||
|
<SelectContent>
|
||||||
|
<SelectItem value="qwen">Qwen TTS</SelectItem>
|
||||||
|
<SelectItem value="luxtts">LuxTTS</SelectItem>
|
||||||
|
</SelectContent>
|
||||||
|
</Select>
|
||||||
|
<FormDescription>
|
||||||
|
{field.value === 'luxtts' ? 'Fast, English-focused' : 'Multi-language'}
|
||||||
|
</FormDescription>
|
||||||
|
<FormMessage />
|
||||||
|
</FormItem>
|
||||||
|
)}
|
||||||
|
/>
|
||||||
|
|
||||||
<FormField
|
<FormField
|
||||||
control={form.control}
|
control={form.control}
|
||||||
name="language"
|
name="language"
|
||||||
@@ -124,29 +151,6 @@ export function GenerationForm() {
|
|||||||
)}
|
)}
|
||||||
/>
|
/>
|
||||||
|
|
||||||
<FormField
|
|
||||||
control={form.control}
|
|
||||||
name="modelSize"
|
|
||||||
render={({ field }) => (
|
|
||||||
<FormItem>
|
|
||||||
<FormLabel>Model Size</FormLabel>
|
|
||||||
<Select onValueChange={field.onChange} defaultValue={field.value}>
|
|
||||||
<FormControl>
|
|
||||||
<SelectTrigger>
|
|
||||||
<SelectValue />
|
|
||||||
</SelectTrigger>
|
|
||||||
</FormControl>
|
|
||||||
<SelectContent>
|
|
||||||
<SelectItem value="1.7B">Qwen TTS 1.7B (Higher Quality)</SelectItem>
|
|
||||||
<SelectItem value="0.6B">Qwen TTS 0.6B (Faster)</SelectItem>
|
|
||||||
</SelectContent>
|
|
||||||
</Select>
|
|
||||||
<FormDescription>Larger models produce better quality</FormDescription>
|
|
||||||
<FormMessage />
|
|
||||||
</FormItem>
|
|
||||||
)}
|
|
||||||
/>
|
|
||||||
|
|
||||||
<FormField
|
<FormField
|
||||||
control={form.control}
|
control={form.control}
|
||||||
name="seed"
|
name="seed"
|
||||||
@@ -170,11 +174,32 @@ export function GenerationForm() {
|
|||||||
/>
|
/>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<Button
|
{form.watch('engine') !== 'luxtts' && (
|
||||||
type="submit"
|
<FormField
|
||||||
className="w-full"
|
control={form.control}
|
||||||
disabled={isPending || !selectedProfileId}
|
name="modelSize"
|
||||||
>
|
render={({ field }) => (
|
||||||
|
<FormItem>
|
||||||
|
<FormLabel>Model Size</FormLabel>
|
||||||
|
<Select onValueChange={field.onChange} defaultValue={field.value}>
|
||||||
|
<FormControl>
|
||||||
|
<SelectTrigger>
|
||||||
|
<SelectValue />
|
||||||
|
</SelectTrigger>
|
||||||
|
</FormControl>
|
||||||
|
<SelectContent>
|
||||||
|
<SelectItem value="1.7B">Qwen TTS 1.7B (Higher Quality)</SelectItem>
|
||||||
|
<SelectItem value="0.6B">Qwen TTS 0.6B (Faster)</SelectItem>
|
||||||
|
</SelectContent>
|
||||||
|
</Select>
|
||||||
|
<FormDescription>Larger models produce better quality</FormDescription>
|
||||||
|
<FormMessage />
|
||||||
|
</FormItem>
|
||||||
|
)}
|
||||||
|
/>
|
||||||
|
)}
|
||||||
|
|
||||||
|
<Button type="submit" className="w-full" disabled={isPending || !selectedProfileId}>
|
||||||
{isPending ? (
|
{isPending ? (
|
||||||
<>
|
<>
|
||||||
<Loader2 className="mr-2 h-4 w-4 animate-spin" />
|
<Loader2 className="mr-2 h-4 w-4 animate-spin" />
|
||||||
|
|||||||
@@ -80,16 +80,19 @@ export function ModelManagement() {
|
|||||||
queryClient.invalidateQueries({ queryKey: ['activeTasks'] });
|
queryClient.invalidateQueries({ queryKey: ['activeTasks'] });
|
||||||
}, [queryClient]);
|
}, [queryClient]);
|
||||||
|
|
||||||
const handleDownloadError = useCallback((error: string) => {
|
const handleDownloadError = useCallback(
|
||||||
console.log('[ModelManagement] Download error, clearing state');
|
(error: string) => {
|
||||||
if (downloadingModel) {
|
console.log('[ModelManagement] Download error, clearing state');
|
||||||
setLocalErrors((prev) => new Map(prev).set(downloadingModel, error));
|
if (downloadingModel) {
|
||||||
setConsoleOpen(true);
|
setLocalErrors((prev) => new Map(prev).set(downloadingModel, error));
|
||||||
}
|
setConsoleOpen(true);
|
||||||
setDownloadingModel(null);
|
}
|
||||||
setDownloadingDisplayName(null);
|
setDownloadingModel(null);
|
||||||
queryClient.invalidateQueries({ queryKey: ['activeTasks'] });
|
setDownloadingDisplayName(null);
|
||||||
}, [queryClient, downloadingModel]);
|
queryClient.invalidateQueries({ queryKey: ['activeTasks'] });
|
||||||
|
},
|
||||||
|
[queryClient, downloadingModel],
|
||||||
|
);
|
||||||
|
|
||||||
// Use progress toast hook for the downloading model
|
// Use progress toast hook for the downloading model
|
||||||
useModelDownloadToast({
|
useModelDownloadToast({
|
||||||
@@ -165,7 +168,11 @@ export function ModelManagement() {
|
|||||||
|
|
||||||
// Optimistically hide the error and suppress downloading state in UI
|
// Optimistically hide the error and suppress downloading state in UI
|
||||||
setDismissedErrors((prev) => new Set(prev).add(modelName));
|
setDismissedErrors((prev) => new Set(prev).add(modelName));
|
||||||
setLocalErrors((prev) => { const next = new Map(prev); next.delete(modelName); return next; });
|
setLocalErrors((prev) => {
|
||||||
|
const next = new Map(prev);
|
||||||
|
next.delete(modelName);
|
||||||
|
return next;
|
||||||
|
});
|
||||||
if (downloadingModel === modelName) {
|
if (downloadingModel === modelName) {
|
||||||
setDownloadingModel(null);
|
setDownloadingModel(null);
|
||||||
setDownloadingDisplayName(null);
|
setDownloadingDisplayName(null);
|
||||||
@@ -178,7 +185,11 @@ export function ModelManagement() {
|
|||||||
setLocalErrors(prevLocalErrors);
|
setLocalErrors(prevLocalErrors);
|
||||||
setDownloadingModel(prevDownloadingModel);
|
setDownloadingModel(prevDownloadingModel);
|
||||||
setDownloadingDisplayName(prevDownloadingDisplayName);
|
setDownloadingDisplayName(prevDownloadingDisplayName);
|
||||||
toast({ title: 'Cancel failed', description: 'Could not cancel the download task.', variant: 'destructive' });
|
toast({
|
||||||
|
title: 'Cancel failed',
|
||||||
|
description: 'Could not cancel the download task.',
|
||||||
|
variant: 'destructive',
|
||||||
|
});
|
||||||
},
|
},
|
||||||
});
|
});
|
||||||
};
|
};
|
||||||
@@ -273,7 +284,9 @@ export function ModelManagement() {
|
|||||||
}}
|
}}
|
||||||
onCancel={() => handleCancel(model.model_name)}
|
onCancel={() => handleCancel(model.model_name)}
|
||||||
isDownloading={downloadingModel === model.model_name}
|
isDownloading={downloadingModel === model.model_name}
|
||||||
isCancelling={cancelMutation.isPending && cancelMutation.variables === model.model_name}
|
isCancelling={
|
||||||
|
cancelMutation.isPending && cancelMutation.variables === model.model_name
|
||||||
|
}
|
||||||
isDismissed={dismissedErrors.has(model.model_name)}
|
isDismissed={dismissedErrors.has(model.model_name)}
|
||||||
erroredDownload={erroredDownloads.get(model.model_name)}
|
erroredDownload={erroredDownloads.get(model.model_name)}
|
||||||
formatSize={formatSize}
|
formatSize={formatSize}
|
||||||
@@ -282,6 +295,34 @@ export function ModelManagement() {
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
{/* LuxTTS Models */}
|
||||||
|
{modelStatus.models.some((m) => m.model_name.startsWith('luxtts')) && (
|
||||||
|
<div>
|
||||||
|
<h3 className="text-sm font-semibold mb-3 text-muted-foreground">LuxTTS Models</h3>
|
||||||
|
<div className="space-y-2">
|
||||||
|
{modelStatus.models
|
||||||
|
.filter((m) => m.model_name.startsWith('luxtts'))
|
||||||
|
.map((model) => (
|
||||||
|
<ModelItem
|
||||||
|
key={model.model_name}
|
||||||
|
model={model}
|
||||||
|
onDownload={() => handleDownload(model.model_name)}
|
||||||
|
onDelete={() => {
|
||||||
|
setModelToDelete({
|
||||||
|
name: model.model_name,
|
||||||
|
displayName: model.display_name,
|
||||||
|
sizeMb: model.size_mb,
|
||||||
|
});
|
||||||
|
setDeleteDialogOpen(true);
|
||||||
|
}}
|
||||||
|
isDownloading={downloadingModel === model.model_name}
|
||||||
|
formatSize={formatSize}
|
||||||
|
/>
|
||||||
|
))}
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
)}
|
||||||
|
|
||||||
{/* Whisper Models */}
|
{/* Whisper Models */}
|
||||||
<div>
|
<div>
|
||||||
<h3 className="text-sm font-semibold mb-3 text-muted-foreground">
|
<h3 className="text-sm font-semibold mb-3 text-muted-foreground">
|
||||||
@@ -305,7 +346,9 @@ export function ModelManagement() {
|
|||||||
}}
|
}}
|
||||||
onCancel={() => handleCancel(model.model_name)}
|
onCancel={() => handleCancel(model.model_name)}
|
||||||
isDownloading={downloadingModel === model.model_name}
|
isDownloading={downloadingModel === model.model_name}
|
||||||
isCancelling={cancelMutation.isPending && cancelMutation.variables === model.model_name}
|
isCancelling={
|
||||||
|
cancelMutation.isPending && cancelMutation.variables === model.model_name
|
||||||
|
}
|
||||||
isDismissed={dismissedErrors.has(model.model_name)}
|
isDismissed={dismissedErrors.has(model.model_name)}
|
||||||
erroredDownload={erroredDownloads.get(model.model_name)}
|
erroredDownload={erroredDownloads.get(model.model_name)}
|
||||||
formatSize={formatSize}
|
formatSize={formatSize}
|
||||||
@@ -353,12 +396,16 @@ export function ModelManagement() {
|
|||||||
{dl.error ? (
|
{dl.error ? (
|
||||||
<>
|
<>
|
||||||
{': '}
|
{': '}
|
||||||
<span className="text-[#ce9178] whitespace-pre-wrap break-all">{dl.error}</span>
|
<span className="text-[#ce9178] whitespace-pre-wrap break-all">
|
||||||
|
{dl.error}
|
||||||
|
</span>
|
||||||
</>
|
</>
|
||||||
) : (
|
) : (
|
||||||
<>
|
<>
|
||||||
{': '}
|
{': '}
|
||||||
<span className="text-[#808080]">No error details available. Try downloading again.</span>
|
<span className="text-[#808080]">
|
||||||
|
No error details available. Try downloading again.
|
||||||
|
</span>
|
||||||
</>
|
</>
|
||||||
)}
|
)}
|
||||||
<div className="text-[#6a9955] mt-0.5">
|
<div className="text-[#6a9955] mt-0.5">
|
||||||
@@ -422,21 +469,31 @@ interface ModelItemProps {
|
|||||||
model_name: string;
|
model_name: string;
|
||||||
display_name: string;
|
display_name: string;
|
||||||
downloaded: boolean;
|
downloaded: boolean;
|
||||||
downloading?: boolean; // From server - true if download in progress
|
downloading?: boolean; // From server - true if download in progress
|
||||||
size_mb?: number;
|
size_mb?: number;
|
||||||
loaded: boolean;
|
loaded: boolean;
|
||||||
};
|
};
|
||||||
onDownload: () => void;
|
onDownload: () => void;
|
||||||
onDelete: () => void;
|
onDelete: () => void;
|
||||||
onCancel: () => void;
|
onCancel: () => void;
|
||||||
isDownloading: boolean; // Local state - true if user just clicked download
|
isDownloading: boolean; // Local state - true if user just clicked download
|
||||||
isCancelling: boolean;
|
isCancelling: boolean;
|
||||||
isDismissed: boolean;
|
isDismissed: boolean;
|
||||||
erroredDownload?: ActiveDownloadTask;
|
erroredDownload?: ActiveDownloadTask;
|
||||||
formatSize: (sizeMb?: number) => string;
|
formatSize: (sizeMb?: number) => string;
|
||||||
}
|
}
|
||||||
|
|
||||||
function ModelItem({ model, onDownload, onDelete, onCancel, isDownloading, isCancelling, isDismissed, erroredDownload, formatSize }: ModelItemProps) {
|
function ModelItem({
|
||||||
|
model,
|
||||||
|
onDownload,
|
||||||
|
onDelete,
|
||||||
|
onCancel,
|
||||||
|
isDownloading,
|
||||||
|
isCancelling,
|
||||||
|
isDismissed,
|
||||||
|
erroredDownload,
|
||||||
|
formatSize,
|
||||||
|
}: ModelItemProps) {
|
||||||
// Use server's downloading state OR local state (for immediate feedback before server updates)
|
// Use server's downloading state OR local state (for immediate feedback before server updates)
|
||||||
// Suppress downloading if user just dismissed/cancelled this model
|
// Suppress downloading if user just dismissed/cancelled this model
|
||||||
const showDownloading = (model.downloading || isDownloading) && !erroredDownload && !isDismissed;
|
const showDownloading = (model.downloading || isDownloading) && !erroredDownload && !isDismissed;
|
||||||
|
|||||||
@@ -34,6 +34,8 @@ 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';
|
||||||
|
instruct?: string;
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface GenerationResponse {
|
export interface GenerationResponse {
|
||||||
|
|||||||
@@ -16,6 +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']).optional(),
|
||||||
});
|
});
|
||||||
|
|
||||||
export type GenerationFormValues = z.infer<typeof generationSchema>;
|
export type GenerationFormValues = z.infer<typeof generationSchema>;
|
||||||
@@ -47,6 +48,7 @@ export function useGenerationForm(options: UseGenerationFormOptions = {}) {
|
|||||||
seed: undefined,
|
seed: undefined,
|
||||||
modelSize: '1.7B',
|
modelSize: '1.7B',
|
||||||
instruct: '',
|
instruct: '',
|
||||||
|
engine: 'qwen',
|
||||||
...options.defaultValues,
|
...options.defaultValues,
|
||||||
},
|
},
|
||||||
});
|
});
|
||||||
@@ -67,8 +69,14 @@ export function useGenerationForm(options: UseGenerationFormOptions = {}) {
|
|||||||
try {
|
try {
|
||||||
setIsGenerating(true);
|
setIsGenerating(true);
|
||||||
|
|
||||||
const modelName = `qwen-tts-${data.modelSize}`;
|
const engine = data.engine || 'qwen';
|
||||||
const displayName = data.modelSize === '1.7B' ? 'Qwen TTS 1.7B' : 'Qwen TTS 0.6B';
|
const modelName = engine === 'luxtts' ? 'luxtts' : `qwen-tts-${data.modelSize}`;
|
||||||
|
const displayName =
|
||||||
|
engine === 'luxtts'
|
||||||
|
? 'LuxTTS'
|
||||||
|
: data.modelSize === '1.7B'
|
||||||
|
? 'Qwen TTS 1.7B'
|
||||||
|
: 'Qwen TTS 0.6B';
|
||||||
|
|
||||||
try {
|
try {
|
||||||
const modelStatus = await apiClient.getModelStatus();
|
const modelStatus = await apiClient.getModelStatus();
|
||||||
@@ -87,8 +95,9 @@ export function useGenerationForm(options: UseGenerationFormOptions = {}) {
|
|||||||
text: data.text,
|
text: data.text,
|
||||||
language: data.language,
|
language: data.language,
|
||||||
seed: data.seed,
|
seed: data.seed,
|
||||||
model_size: data.modelSize,
|
model_size: engine === 'luxtts' ? undefined : data.modelSize,
|
||||||
instruct: data.instruct || undefined,
|
engine,
|
||||||
|
instruct: engine === 'luxtts' ? undefined : data.instruct || undefined,
|
||||||
});
|
});
|
||||||
|
|
||||||
toast({
|
toast({
|
||||||
|
|||||||
@@ -112,29 +112,57 @@ class STTBackend(Protocol):
|
|||||||
|
|
||||||
# Global backend instances
|
# Global backend instances
|
||||||
_tts_backend: Optional[TTSBackend] = None
|
_tts_backend: Optional[TTSBackend] = None
|
||||||
|
_tts_backends: dict[str, TTSBackend] = {}
|
||||||
_stt_backend: Optional[STTBackend] = None
|
_stt_backend: Optional[STTBackend] = None
|
||||||
|
|
||||||
|
# Supported TTS engines
|
||||||
|
TTS_ENGINES = {
|
||||||
|
"qwen": "Qwen TTS",
|
||||||
|
"luxtts": "LuxTTS",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def get_tts_backend() -> TTSBackend:
|
def get_tts_backend() -> TTSBackend:
|
||||||
"""
|
"""
|
||||||
Get or create TTS backend instance based on platform.
|
Get or create the default (Qwen) TTS backend instance based on platform.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
TTS backend instance (MLX or PyTorch)
|
TTS backend instance (MLX or PyTorch)
|
||||||
"""
|
"""
|
||||||
global _tts_backend
|
return get_tts_backend_for_engine("qwen")
|
||||||
|
|
||||||
if _tts_backend is None:
|
|
||||||
|
def get_tts_backend_for_engine(engine: str) -> TTSBackend:
|
||||||
|
"""
|
||||||
|
Get or create a TTS backend for the given engine.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
engine: Engine name ("qwen" or "luxtts")
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
TTS backend instance
|
||||||
|
"""
|
||||||
|
global _tts_backends
|
||||||
|
|
||||||
|
if engine in _tts_backends:
|
||||||
|
return _tts_backends[engine]
|
||||||
|
|
||||||
|
if engine == "qwen":
|
||||||
backend_type = get_backend_type()
|
backend_type = get_backend_type()
|
||||||
|
|
||||||
if backend_type == "mlx":
|
if backend_type == "mlx":
|
||||||
from .mlx_backend import MLXTTSBackend
|
from .mlx_backend import MLXTTSBackend
|
||||||
_tts_backend = MLXTTSBackend()
|
backend = MLXTTSBackend()
|
||||||
else:
|
else:
|
||||||
from .pytorch_backend import PyTorchTTSBackend
|
from .pytorch_backend import PyTorchTTSBackend
|
||||||
_tts_backend = PyTorchTTSBackend()
|
backend = PyTorchTTSBackend()
|
||||||
|
elif engine == "luxtts":
|
||||||
|
from .luxtts_backend import LuxTTSBackend
|
||||||
|
backend = LuxTTSBackend()
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown TTS engine: {engine}. Supported: {list(TTS_ENGINES.keys())}")
|
||||||
|
|
||||||
return _tts_backend
|
_tts_backends[engine] = backend
|
||||||
|
return backend
|
||||||
|
|
||||||
|
|
||||||
def get_stt_backend() -> STTBackend:
|
def get_stt_backend() -> STTBackend:
|
||||||
@@ -161,6 +189,7 @@ def get_stt_backend() -> STTBackend:
|
|||||||
|
|
||||||
def reset_backends():
|
def reset_backends():
|
||||||
"""Reset backend instances (useful for testing)."""
|
"""Reset backend instances (useful for testing)."""
|
||||||
global _tts_backend, _stt_backend
|
global _tts_backend, _tts_backends, _stt_backend
|
||||||
_tts_backend = None
|
_tts_backend = None
|
||||||
|
_tts_backends.clear()
|
||||||
_stt_backend = None
|
_stt_backend = None
|
||||||
|
|||||||
@@ -0,0 +1,264 @@
|
|||||||
|
"""
|
||||||
|
LuxTTS backend implementation.
|
||||||
|
|
||||||
|
Wraps the LuxTTS (ZipVoice) model for zero-shot voice cloning.
|
||||||
|
~1GB VRAM, 48kHz output, 150x realtime on CPU.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import List, Optional, Tuple
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from . import TTSBackend
|
||||||
|
from ..utils.audio import normalize_audio, load_audio
|
||||||
|
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
|
||||||
|
from ..utils.progress import get_progress_manager
|
||||||
|
from ..utils.tasks import get_task_manager
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# HuggingFace repo for model weight detection
|
||||||
|
LUXTTS_HF_REPO = "YatharthS/LuxTTS"
|
||||||
|
|
||||||
|
|
||||||
|
class LuxTTSBackend:
|
||||||
|
"""LuxTTS backend for zero-shot voice cloning."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.model = None
|
||||||
|
self.model_size = "default" # LuxTTS has only one model size
|
||||||
|
self._device = None
|
||||||
|
|
||||||
|
def _get_device(self) -> str:
|
||||||
|
"""Get the best available device."""
|
||||||
|
import torch
|
||||||
|
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
return "cuda"
|
||||||
|
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||||
|
return "mps"
|
||||||
|
return "cpu"
|
||||||
|
|
||||||
|
def is_loaded(self) -> bool:
|
||||||
|
return self.model is not None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def device(self) -> str:
|
||||||
|
if self._device is None:
|
||||||
|
self._device = self._get_device()
|
||||||
|
return self._device
|
||||||
|
|
||||||
|
def _get_model_path(self, model_size: str) -> str:
|
||||||
|
return LUXTTS_HF_REPO
|
||||||
|
|
||||||
|
def _is_model_cached(self, model_size: str = "default") -> bool:
|
||||||
|
"""Check if LuxTTS model weights are cached locally."""
|
||||||
|
try:
|
||||||
|
from huggingface_hub import constants as hf_constants
|
||||||
|
|
||||||
|
repo_cache = (
|
||||||
|
Path(hf_constants.HF_HUB_CACHE)
|
||||||
|
/ ("models--" + LUXTTS_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
|
||||||
|
|
||||||
|
snapshots_dir = repo_cache / "snapshots"
|
||||||
|
if snapshots_dir.exists():
|
||||||
|
has_weights = any(snapshots_dir.rglob("*.pt")) or any(
|
||||||
|
snapshots_dir.rglob("*.safetensors")
|
||||||
|
) or any(snapshots_dir.rglob("*.onnx")) or any(
|
||||||
|
snapshots_dir.rglob("*.bin")
|
||||||
|
)
|
||||||
|
return has_weights
|
||||||
|
|
||||||
|
return False
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Error checking LuxTTS cache: {e}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def load_model(self, model_size: str = "default") -> None:
|
||||||
|
"""Load the LuxTTS model."""
|
||||||
|
if self.model is not None:
|
||||||
|
return
|
||||||
|
|
||||||
|
await asyncio.to_thread(self._load_model_sync)
|
||||||
|
|
||||||
|
def _load_model_sync(self):
|
||||||
|
"""Synchronous model loading."""
|
||||||
|
progress_manager = get_progress_manager()
|
||||||
|
task_manager = get_task_manager()
|
||||||
|
model_name = "luxtts"
|
||||||
|
|
||||||
|
is_cached = self._is_model_cached()
|
||||||
|
|
||||||
|
if not is_cached:
|
||||||
|
task_manager.start_download(model_name)
|
||||||
|
progress_manager.update_progress(
|
||||||
|
model_name=model_name,
|
||||||
|
current=0,
|
||||||
|
total=0,
|
||||||
|
filename="Downloading LuxTTS model...",
|
||||||
|
status="downloading",
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
from zipvoice.luxvoice import LuxTTS
|
||||||
|
|
||||||
|
device = self.device
|
||||||
|
logger.info(f"Loading LuxTTS on {device}...")
|
||||||
|
|
||||||
|
# LuxTTS constructor downloads model and loads everything
|
||||||
|
if device == "cpu":
|
||||||
|
import os
|
||||||
|
threads = os.cpu_count() or 4
|
||||||
|
self.model = LuxTTS(
|
||||||
|
model_path=LUXTTS_HF_REPO,
|
||||||
|
device="cpu",
|
||||||
|
threads=min(threads, 8),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.model = LuxTTS(
|
||||||
|
model_path=LUXTTS_HF_REPO,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not is_cached:
|
||||||
|
progress_manager.mark_complete(model_name)
|
||||||
|
task_manager.complete_download(model_name)
|
||||||
|
|
||||||
|
logger.info("LuxTTS loaded successfully")
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to load LuxTTS: {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:
|
||||||
|
del self.model
|
||||||
|
self.model = None
|
||||||
|
|
||||||
|
import torch
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
logger.info("LuxTTS 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.
|
||||||
|
|
||||||
|
LuxTTS uses its own encode_prompt() which runs Whisper ASR internally
|
||||||
|
to transcribe the reference. The reference_text parameter is not used
|
||||||
|
by LuxTTS itself, but we include it in the cache key for consistency.
|
||||||
|
"""
|
||||||
|
await self.load_model()
|
||||||
|
|
||||||
|
if use_cache:
|
||||||
|
# Include "luxtts" in the cache key so it doesn't collide with Qwen prompts
|
||||||
|
cache_key = "luxtts_" + get_cache_key(audio_path, reference_text)
|
||||||
|
cached = get_cached_voice_prompt(cache_key)
|
||||||
|
if cached is not None and isinstance(cached, dict):
|
||||||
|
return cached, True
|
||||||
|
|
||||||
|
def _encode_sync():
|
||||||
|
return self.model.encode_prompt(
|
||||||
|
prompt_audio=str(audio_path),
|
||||||
|
duration=5,
|
||||||
|
rms=0.01,
|
||||||
|
)
|
||||||
|
|
||||||
|
encoded = await asyncio.to_thread(_encode_sync)
|
||||||
|
|
||||||
|
if use_cache:
|
||||||
|
cache_key = "luxtts_" + get_cache_key(audio_path, reference_text)
|
||||||
|
cache_voice_prompt(cache_key, encoded)
|
||||||
|
|
||||||
|
return encoded, False
|
||||||
|
|
||||||
|
async def combine_voice_prompts(
|
||||||
|
self,
|
||||||
|
audio_paths: List[str],
|
||||||
|
reference_texts: List[str],
|
||||||
|
) -> Tuple[np.ndarray, str]:
|
||||||
|
"""
|
||||||
|
Combine multiple reference samples.
|
||||||
|
|
||||||
|
LuxTTS doesn't have native multi-prompt support, so we concatenate
|
||||||
|
the audio and let encode_prompt handle the combined clip.
|
||||||
|
"""
|
||||||
|
combined_audio = []
|
||||||
|
for path in audio_paths:
|
||||||
|
audio, sr = load_audio(path, sample_rate=24000)
|
||||||
|
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 from text using LuxTTS.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: Text to synthesize
|
||||||
|
voice_prompt: Encoded prompt dict from encode_prompt()
|
||||||
|
language: Language code (LuxTTS is English-focused)
|
||||||
|
seed: Random seed for reproducibility
|
||||||
|
instruct: Not supported by LuxTTS (ignored)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (audio_array, sample_rate)
|
||||||
|
"""
|
||||||
|
await self.load_model()
|
||||||
|
|
||||||
|
def _generate_sync():
|
||||||
|
import torch
|
||||||
|
|
||||||
|
if seed is not None:
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.manual_seed(seed)
|
||||||
|
|
||||||
|
wav = self.model.generate_speech(
|
||||||
|
text=text,
|
||||||
|
encode_dict=voice_prompt,
|
||||||
|
num_steps=4,
|
||||||
|
guidance_scale=3.0,
|
||||||
|
t_shift=0.5,
|
||||||
|
speed=1.0,
|
||||||
|
return_smooth=False, # 48kHz output
|
||||||
|
)
|
||||||
|
|
||||||
|
# LuxTTS returns a tensor, convert to numpy
|
||||||
|
audio = wav.numpy().squeeze()
|
||||||
|
return audio, 48000
|
||||||
|
|
||||||
|
return await asyncio.to_thread(_generate_sync)
|
||||||
+103
-39
@@ -602,47 +602,69 @@ async def generate_speech(
|
|||||||
raise HTTPException(status_code=404, detail="Profile not found")
|
raise HTTPException(status_code=404, detail="Profile not found")
|
||||||
|
|
||||||
# Generate audio
|
# Generate audio
|
||||||
|
from .backends import get_tts_backend_for_engine
|
||||||
|
|
||||||
# Resolve model size and load the correct model FIRST.
|
engine = data.engine or "qwen"
|
||||||
# This must happen before create_voice_prompt_for_profile because that
|
tts_model = get_tts_backend_for_engine(engine)
|
||||||
# function calls load_model_async(None), which falls back to self.model_size.
|
|
||||||
# If the model is already loaded with the right size at that point, it
|
# Resolve model size (only relevant for Qwen engine)
|
||||||
# returns immediately and the voice prompt is created by the correct model.
|
|
||||||
tts_model = tts.get_tts_model()
|
|
||||||
model_size = data.model_size or "1.7B"
|
model_size = data.model_size or "1.7B"
|
||||||
|
|
||||||
# Check if model needs to be downloaded first
|
# Check if model needs to be downloaded first
|
||||||
model_path = tts_model._get_model_path(model_size)
|
if engine == "qwen":
|
||||||
if not tts_model._is_model_cached(model_size):
|
if not tts_model._is_model_cached(model_size):
|
||||||
# Model is not fully cached — kick off a background download and tell
|
model_name = f"qwen-tts-{model_size}"
|
||||||
# the client to retry once it's ready.
|
|
||||||
model_name = f"qwen-tts-{model_size}"
|
|
||||||
|
|
||||||
async def download_model_background():
|
async def download_model_background():
|
||||||
try:
|
try:
|
||||||
await tts_model.load_model_async(model_size)
|
await tts_model.load_model_async(model_size)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
task_manager.error_download(model_name, str(e))
|
task_manager.error_download(model_name, str(e))
|
||||||
|
|
||||||
task_manager.start_download(model_name)
|
task_manager.start_download(model_name)
|
||||||
asyncio.create_task(download_model_background())
|
asyncio.create_task(download_model_background())
|
||||||
|
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=202,
|
status_code=202,
|
||||||
detail={
|
detail={
|
||||||
"message": f"Model {model_size} is being downloaded. Please wait and try again.",
|
"message": f"Model {model_size} is being downloaded. Please wait and try again.",
|
||||||
"model_name": model_name,
|
"model_name": model_name,
|
||||||
"downloading": True,
|
"downloading": True,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
# Load (or switch to) the requested model before building the voice prompt
|
# Load (or switch to) the requested model
|
||||||
await tts_model.load_model_async(model_size)
|
await tts_model.load_model_async(model_size)
|
||||||
|
elif engine == "luxtts":
|
||||||
|
if not tts_model._is_model_cached():
|
||||||
|
model_name = "luxtts"
|
||||||
|
|
||||||
# Create voice prompt from profile (model is already loaded with correct size)
|
async def download_luxtts_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_luxtts_background())
|
||||||
|
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=202,
|
||||||
|
detail={
|
||||||
|
"message": "LuxTTS 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(
|
voice_prompt = await profiles.create_voice_prompt_for_profile(
|
||||||
data.profile_id,
|
data.profile_id,
|
||||||
db,
|
db,
|
||||||
|
use_cache=True,
|
||||||
|
engine=engine,
|
||||||
)
|
)
|
||||||
|
|
||||||
audio, sample_rate = await tts_model.generate(
|
audio, sample_rate = await tts_model.generate(
|
||||||
@@ -699,23 +721,34 @@ async def stream_speech(
|
|||||||
playing audio before the entire file has been received. This endpoint
|
playing audio before the entire file has been received. This endpoint
|
||||||
does NOT create a history entry — use /generate for that.
|
does NOT create a history entry — use /generate for that.
|
||||||
"""
|
"""
|
||||||
|
from .backends import get_tts_backend_for_engine
|
||||||
|
|
||||||
profile = await profiles.get_profile(data.profile_id, db)
|
profile = await profiles.get_profile(data.profile_id, db)
|
||||||
if not profile:
|
if not profile:
|
||||||
raise HTTPException(status_code=404, detail="Profile not found")
|
raise HTTPException(status_code=404, detail="Profile not found")
|
||||||
|
|
||||||
tts_model = tts.get_tts_model()
|
engine = data.engine or "qwen"
|
||||||
|
tts_model = get_tts_backend_for_engine(engine)
|
||||||
model_size = data.model_size or "1.7B"
|
model_size = data.model_size or "1.7B"
|
||||||
|
|
||||||
if not tts_model._is_model_cached(model_size):
|
if engine == "qwen":
|
||||||
raise HTTPException(
|
if not tts_model._is_model_cached(model_size):
|
||||||
status_code=400,
|
raise HTTPException(
|
||||||
detail=f"Model {model_size} is not downloaded yet. Use /generate to trigger a download.",
|
status_code=400,
|
||||||
)
|
detail=f"Model {model_size} is not downloaded yet. Use /generate to trigger a download.",
|
||||||
|
)
|
||||||
|
await tts_model.load_model_async(model_size)
|
||||||
|
elif engine == "luxtts":
|
||||||
|
if not tts_model._is_model_cached():
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail="LuxTTS model is not downloaded yet. Use /generate to trigger a download.",
|
||||||
|
)
|
||||||
|
await tts_model.load_model()
|
||||||
|
|
||||||
# Load the correct model before building the voice prompt (fixes issue #96)
|
voice_prompt = await profiles.create_voice_prompt_for_profile(
|
||||||
await tts_model.load_model_async(model_size)
|
data.profile_id, db, engine=engine,
|
||||||
|
)
|
||||||
voice_prompt = await profiles.create_voice_prompt_for_profile(data.profile_id, db)
|
|
||||||
|
|
||||||
audio, sample_rate = await tts_model.generate(
|
audio, sample_rate = await tts_model.generate(
|
||||||
data.text,
|
data.text,
|
||||||
@@ -1324,6 +1357,15 @@ async def get_model_status():
|
|||||||
whisper_medium_id = "openai/whisper-medium"
|
whisper_medium_id = "openai/whisper-medium"
|
||||||
whisper_large_id = "openai/whisper-large-v3"
|
whisper_large_id = "openai/whisper-large-v3"
|
||||||
|
|
||||||
|
# Check if LuxTTS backend is loaded
|
||||||
|
def check_luxtts_loaded():
|
||||||
|
try:
|
||||||
|
from .backends import get_tts_backend_for_engine
|
||||||
|
backend = get_tts_backend_for_engine("luxtts")
|
||||||
|
return backend.is_loaded()
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
model_configs = [
|
model_configs = [
|
||||||
{
|
{
|
||||||
"model_name": "qwen-tts-1.7B",
|
"model_name": "qwen-tts-1.7B",
|
||||||
@@ -1339,6 +1381,13 @@ async def get_model_status():
|
|||||||
"model_size": "0.6B",
|
"model_size": "0.6B",
|
||||||
"check_loaded": lambda: check_tts_loaded("0.6B"),
|
"check_loaded": lambda: check_tts_loaded("0.6B"),
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
"model_name": "luxtts",
|
||||||
|
"display_name": "LuxTTS (Fast, CPU-friendly)",
|
||||||
|
"hf_repo_id": "YatharthS/LuxTTS",
|
||||||
|
"model_size": "default",
|
||||||
|
"check_loaded": check_luxtts_loaded,
|
||||||
|
},
|
||||||
{
|
{
|
||||||
"model_name": "whisper-base",
|
"model_name": "whisper-base",
|
||||||
"display_name": "Whisper Base",
|
"display_name": "Whisper Base",
|
||||||
@@ -1521,6 +1570,7 @@ async def get_model_status():
|
|||||||
async def trigger_model_download(request: models.ModelDownloadRequest):
|
async def trigger_model_download(request: models.ModelDownloadRequest):
|
||||||
"""Trigger download of a specific model."""
|
"""Trigger download of a specific model."""
|
||||||
import asyncio
|
import asyncio
|
||||||
|
from .backends import get_tts_backend_for_engine
|
||||||
|
|
||||||
task_manager = get_task_manager()
|
task_manager = get_task_manager()
|
||||||
progress_manager = get_progress_manager()
|
progress_manager = get_progress_manager()
|
||||||
@@ -1534,6 +1584,10 @@ async def trigger_model_download(request: models.ModelDownloadRequest):
|
|||||||
"model_size": "0.6B",
|
"model_size": "0.6B",
|
||||||
"load_func": lambda: tts.get_tts_model().load_model("0.6B"),
|
"load_func": lambda: tts.get_tts_model().load_model("0.6B"),
|
||||||
},
|
},
|
||||||
|
"luxtts": {
|
||||||
|
"model_size": "default",
|
||||||
|
"load_func": lambda: get_tts_backend_for_engine("luxtts").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"),
|
||||||
@@ -1646,6 +1700,11 @@ async def delete_model(model_name: str):
|
|||||||
"model_size": "0.6B",
|
"model_size": "0.6B",
|
||||||
"model_type": "tts",
|
"model_type": "tts",
|
||||||
},
|
},
|
||||||
|
"luxtts": {
|
||||||
|
"hf_repo_id": "YatharthS/LuxTTS",
|
||||||
|
"model_size": "default",
|
||||||
|
"model_type": "luxtts",
|
||||||
|
},
|
||||||
"whisper-base": {
|
"whisper-base": {
|
||||||
"hf_repo_id": "openai/whisper-base",
|
"hf_repo_id": "openai/whisper-base",
|
||||||
"model_size": "base",
|
"model_size": "base",
|
||||||
@@ -1680,6 +1739,11 @@ async def delete_model(model_name: str):
|
|||||||
tts_model = tts.get_tts_model()
|
tts_model = tts.get_tts_model()
|
||||||
if tts_model.is_loaded() and tts_model.model_size == config["model_size"]:
|
if tts_model.is_loaded() and tts_model.model_size == config["model_size"]:
|
||||||
tts.unload_tts_model()
|
tts.unload_tts_model()
|
||||||
|
elif config["model_type"] == "luxtts":
|
||||||
|
from .backends import get_tts_backend_for_engine
|
||||||
|
luxtts = get_tts_backend_for_engine("luxtts")
|
||||||
|
if luxtts.is_loaded():
|
||||||
|
luxtts.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"]:
|
||||||
|
|||||||
@@ -57,6 +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)$")
|
||||||
|
|
||||||
|
|
||||||
class GenerationResponse(BaseModel):
|
class GenerationResponse(BaseModel):
|
||||||
|
|||||||
+5
-1
@@ -327,6 +327,7 @@ async def create_voice_prompt_for_profile(
|
|||||||
profile_id: str,
|
profile_id: str,
|
||||||
db: Session,
|
db: Session,
|
||||||
use_cache: bool = True,
|
use_cache: bool = True,
|
||||||
|
engine: str = "qwen",
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""
|
"""
|
||||||
Create a combined voice prompt from all samples in a profile.
|
Create a combined voice prompt from all samples in a profile.
|
||||||
@@ -335,17 +336,20 @@ async def create_voice_prompt_for_profile(
|
|||||||
profile_id: Profile ID
|
profile_id: Profile ID
|
||||||
db: Database session
|
db: Database session
|
||||||
use_cache: Whether to use cached prompts
|
use_cache: Whether to use cached prompts
|
||||||
|
engine: TTS engine to create prompt for ("qwen" or "luxtts")
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Voice prompt dictionary
|
Voice prompt dictionary
|
||||||
"""
|
"""
|
||||||
|
from .backends import get_tts_backend_for_engine
|
||||||
|
|
||||||
# Get all samples for profile
|
# Get all samples for profile
|
||||||
samples = db.query(DBProfileSample).filter_by(profile_id=profile_id).all()
|
samples = db.query(DBProfileSample).filter_by(profile_id=profile_id).all()
|
||||||
|
|
||||||
if not samples:
|
if not samples:
|
||||||
raise ValueError(f"No samples found for profile {profile_id}")
|
raise ValueError(f"No samples found for profile {profile_id}")
|
||||||
|
|
||||||
tts_model = get_tts_model()
|
tts_model = get_tts_backend_for_engine(engine)
|
||||||
|
|
||||||
if len(samples) == 1:
|
if len(samples) == 1:
|
||||||
# Single sample - use directly
|
# Single sample - use directly
|
||||||
|
|||||||
@@ -9,11 +9,19 @@ alembic>=1.13.0
|
|||||||
|
|
||||||
# ML models
|
# ML models
|
||||||
torch>=2.1.0
|
torch>=2.1.0
|
||||||
transformers>=4.36.0
|
transformers>=4.36.0,<=4.57.6
|
||||||
accelerate>=0.26.0
|
accelerate>=0.26.0
|
||||||
huggingface_hub>=0.20.0
|
huggingface_hub>=0.20.0
|
||||||
qwen-tts>=0.0.5
|
qwen-tts>=0.0.5
|
||||||
|
|
||||||
|
# LuxTTS (voice cloning engine)
|
||||||
|
Zipvoice @ git+https://github.com/ysharma3501/LuxTTS.git
|
||||||
|
onnxruntime>=1.16.0
|
||||||
|
piper-phonemize>=1.1.0
|
||||||
|
lhotse>=1.20.0
|
||||||
|
pydub>=0.25.0
|
||||||
|
inflect>=7.0.0
|
||||||
|
|
||||||
# Audio processing
|
# Audio processing
|
||||||
librosa>=0.10.0
|
librosa>=0.10.0
|
||||||
soundfile>=0.12.0
|
soundfile>=0.12.0
|
||||||
|
|||||||
Reference in New Issue
Block a user