mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 00:55:14 -07:00
Add model size selection to GenerationForm and enhance model management features
This commit is contained in:
@@ -31,6 +31,7 @@ const generationSchema = z.object({
|
|||||||
text: z.string().min(1, 'Text is required').max(5000),
|
text: z.string().min(1, 'Text is required').max(5000),
|
||||||
language: z.enum(['en', 'zh']),
|
language: z.enum(['en', 'zh']),
|
||||||
seed: z.number().int().optional(),
|
seed: z.number().int().optional(),
|
||||||
|
modelSize: z.enum(['1.7B', '0.6B']).optional(),
|
||||||
});
|
});
|
||||||
|
|
||||||
type GenerationFormValues = z.infer<typeof generationSchema>;
|
type GenerationFormValues = z.infer<typeof generationSchema>;
|
||||||
@@ -47,6 +48,7 @@ export function GenerationForm() {
|
|||||||
text: '',
|
text: '',
|
||||||
language: 'en',
|
language: 'en',
|
||||||
seed: undefined,
|
seed: undefined,
|
||||||
|
modelSize: '1.7B',
|
||||||
},
|
},
|
||||||
});
|
});
|
||||||
|
|
||||||
@@ -57,6 +59,7 @@ export function GenerationForm() {
|
|||||||
text: data.text,
|
text: data.text,
|
||||||
language: data.language,
|
language: data.language,
|
||||||
seed: data.seed,
|
seed: data.seed,
|
||||||
|
model_size: data.modelSize,
|
||||||
});
|
});
|
||||||
|
|
||||||
toast({
|
toast({
|
||||||
@@ -126,7 +129,7 @@ export function GenerationForm() {
|
|||||||
)}
|
)}
|
||||||
/>
|
/>
|
||||||
|
|
||||||
<div className="grid gap-4 md:grid-cols-2">
|
<div className="grid gap-4 md:grid-cols-3">
|
||||||
<FormField
|
<FormField
|
||||||
control={form.control}
|
control={form.control}
|
||||||
name="language"
|
name="language"
|
||||||
@@ -149,6 +152,29 @@ 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"
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import { useQuery, useMutation, useQueryClient } from '@tanstack/react-query';
|
import { useQuery, useMutation, useQueryClient } from '@tanstack/react-query';
|
||||||
|
import { useState } from 'react';
|
||||||
import { apiClient } from '@/lib/api/client';
|
import { apiClient } from '@/lib/api/client';
|
||||||
import { Card, CardContent, CardHeader, CardTitle, CardDescription } from '@/components/ui/card';
|
import { Card, CardContent, CardHeader, CardTitle, CardDescription } from '@/components/ui/card';
|
||||||
import { Button } from '@/components/ui/button';
|
import { Button } from '@/components/ui/button';
|
||||||
@@ -10,6 +11,7 @@ import { useToast } from '@/components/ui/use-toast';
|
|||||||
export function ModelManagement() {
|
export function ModelManagement() {
|
||||||
const { toast } = useToast();
|
const { toast } = useToast();
|
||||||
const queryClient = useQueryClient();
|
const queryClient = useQueryClient();
|
||||||
|
const [downloadingModel, setDownloadingModel] = useState<string | null>(null);
|
||||||
|
|
||||||
const { data: modelStatus, isLoading } = useQuery({
|
const { data: modelStatus, isLoading } = useQuery({
|
||||||
queryKey: ['modelStatus'],
|
queryKey: ['modelStatus'],
|
||||||
@@ -18,7 +20,10 @@ export function ModelManagement() {
|
|||||||
});
|
});
|
||||||
|
|
||||||
const downloadMutation = useMutation({
|
const downloadMutation = useMutation({
|
||||||
mutationFn: (modelName: string) => apiClient.triggerModelDownload(modelName),
|
mutationFn: (modelName: string) => {
|
||||||
|
setDownloadingModel(modelName);
|
||||||
|
return apiClient.triggerModelDownload(modelName);
|
||||||
|
},
|
||||||
onSuccess: (_, modelName) => {
|
onSuccess: (_, modelName) => {
|
||||||
toast({
|
toast({
|
||||||
title: 'Download started',
|
title: 'Download started',
|
||||||
@@ -30,12 +35,19 @@ export function ModelManagement() {
|
|||||||
}, 1000);
|
}, 1000);
|
||||||
},
|
},
|
||||||
onError: (error: Error) => {
|
onError: (error: Error) => {
|
||||||
|
setDownloadingModel(null);
|
||||||
toast({
|
toast({
|
||||||
title: 'Download failed',
|
title: 'Download failed',
|
||||||
description: error.message,
|
description: error.message,
|
||||||
variant: 'destructive',
|
variant: 'destructive',
|
||||||
});
|
});
|
||||||
},
|
},
|
||||||
|
onSettled: () => {
|
||||||
|
// Clear downloading state after a delay to allow progress to show
|
||||||
|
setTimeout(() => {
|
||||||
|
setDownloadingModel(null);
|
||||||
|
}, 2000);
|
||||||
|
},
|
||||||
});
|
});
|
||||||
|
|
||||||
const formatSize = (sizeMb?: number): string => {
|
const formatSize = (sizeMb?: number): string => {
|
||||||
@@ -70,7 +82,7 @@ export function ModelManagement() {
|
|||||||
key={model.model_name}
|
key={model.model_name}
|
||||||
model={model}
|
model={model}
|
||||||
onDownload={() => downloadMutation.mutate(model.model_name)}
|
onDownload={() => downloadMutation.mutate(model.model_name)}
|
||||||
isDownloading={downloadMutation.isPending}
|
isDownloading={downloadingModel === model.model_name}
|
||||||
formatSize={formatSize}
|
formatSize={formatSize}
|
||||||
/>
|
/>
|
||||||
))}
|
))}
|
||||||
@@ -88,7 +100,7 @@ export function ModelManagement() {
|
|||||||
key={model.model_name}
|
key={model.model_name}
|
||||||
model={model}
|
model={model}
|
||||||
onDownload={() => downloadMutation.mutate(model.model_name)}
|
onDownload={() => downloadMutation.mutate(model.model_name)}
|
||||||
isDownloading={downloadMutation.isPending}
|
isDownloading={downloadingModel === model.model_name}
|
||||||
formatSize={formatSize}
|
formatSize={formatSize}
|
||||||
/>
|
/>
|
||||||
))}
|
))}
|
||||||
|
|||||||
@@ -47,18 +47,9 @@ export function ServerStatus() {
|
|||||||
<span className="text-sm">Connected</span>
|
<span className="text-sm">Connected</span>
|
||||||
</div>
|
</div>
|
||||||
<div className="flex flex-wrap gap-2">
|
<div className="flex flex-wrap gap-2">
|
||||||
<Badge variant={health.model_loaded ? 'default' : 'secondary'}>
|
<Badge variant={health.model_loaded || health.model_downloaded ? 'default' : 'secondary'}>
|
||||||
Model: {health.model_loaded
|
{health.model_loaded || health.model_downloaded ? 'Model Ready' : 'No Model'}
|
||||||
? `Loaded${health.model_size ? ` (${health.model_size})` : ''}`
|
|
||||||
: health.model_downloaded === false
|
|
||||||
? 'Not Downloaded'
|
|
||||||
: 'Not Loaded'}
|
|
||||||
</Badge>
|
</Badge>
|
||||||
{health.model_downloaded === true && !health.model_loaded && (
|
|
||||||
<Badge variant="outline">
|
|
||||||
Model Cached (will load on first use)
|
|
||||||
</Badge>
|
|
||||||
)}
|
|
||||||
<Badge variant={health.gpu_available ? 'default' : 'secondary'}>
|
<Badge variant={health.gpu_available ? 'default' : 'secondary'}>
|
||||||
GPU: {health.gpu_available ? 'Available' : 'Not Available'}
|
GPU: {health.gpu_available ? 'Available' : 'Not Available'}
|
||||||
</Badge>
|
</Badge>
|
||||||
|
|||||||
+20
-3
@@ -64,9 +64,23 @@ async def health():
|
|||||||
if gpu_available:
|
if gpu_available:
|
||||||
vram_used = torch.cuda.memory_allocated() / 1024 / 1024 # MB
|
vram_used = torch.cuda.memory_allocated() / 1024 / 1024 # MB
|
||||||
|
|
||||||
# Check if model is loaded
|
# Check if model is loaded - use the same logic as model status endpoint
|
||||||
model_loaded = tts_model.is_loaded()
|
model_loaded = False
|
||||||
model_size = tts_model.model_size if model_loaded else None
|
model_size = None
|
||||||
|
try:
|
||||||
|
# Use the same check as model status endpoint
|
||||||
|
if tts_model.is_loaded():
|
||||||
|
model_loaded = True
|
||||||
|
# Get the actual loaded model size
|
||||||
|
# Check _current_model_size first (more reliable for actually loaded models)
|
||||||
|
model_size = getattr(tts_model, '_current_model_size', None)
|
||||||
|
if not model_size:
|
||||||
|
# Fallback to model_size attribute (which should be set when model loads)
|
||||||
|
model_size = getattr(tts_model, 'model_size', None)
|
||||||
|
except Exception:
|
||||||
|
# If there's an error checking, assume not loaded
|
||||||
|
model_loaded = False
|
||||||
|
model_size = None
|
||||||
|
|
||||||
# Check if default model is downloaded (cached)
|
# Check if default model is downloaded (cached)
|
||||||
model_downloaded = None
|
model_downloaded = None
|
||||||
@@ -240,6 +254,9 @@ async def generate_speech(
|
|||||||
|
|
||||||
# Generate audio
|
# Generate audio
|
||||||
tts_model = tts.get_tts_model()
|
tts_model = tts.get_tts_model()
|
||||||
|
# Load the requested model size if different from current
|
||||||
|
model_size = data.model_size or "1.7B"
|
||||||
|
tts_model.load_model(model_size)
|
||||||
audio, sample_rate = await tts_model.generate(
|
audio, sample_rate = await tts_model.generate(
|
||||||
data.text,
|
data.text,
|
||||||
voice_prompt,
|
voice_prompt,
|
||||||
|
|||||||
@@ -49,6 +49,7 @@ class GenerationRequest(BaseModel):
|
|||||||
text: str = Field(..., min_length=1, max_length=5000)
|
text: str = Field(..., min_length=1, max_length=5000)
|
||||||
language: str = Field(default="en", pattern="^(en|zh)$")
|
language: str = Field(default="en", pattern="^(en|zh)$")
|
||||||
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)$")
|
||||||
|
|
||||||
|
|
||||||
class GenerationResponse(BaseModel):
|
class GenerationResponse(BaseModel):
|
||||||
|
|||||||
Reference in New Issue
Block a user