Add model size selection to GenerationForm and enhance model management features

This commit is contained in:
Jamie Pine
2026-01-25 03:19:27 -08:00
parent 6429cb6673
commit 6164877f7f
5 changed files with 65 additions and 18 deletions
@@ -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
View File
@@ -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,
+1
View File
@@ -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):