From 6164877f7fd3680f51b6f6b151905e075e7662ae Mon Sep 17 00:00:00 2001 From: Jamie Pine Date: Sun, 25 Jan 2026 03:19:27 -0800 Subject: [PATCH] Add model size selection to GenerationForm and enhance model management features --- .../components/Generation/GenerationForm.tsx | 28 ++++++++++++++++++- .../ServerSettings/ModelManagement.tsx | 18 ++++++++++-- .../ServerSettings/ServerStatus.tsx | 13 ++------- backend/main.py | 23 +++++++++++++-- backend/models.py | 1 + 5 files changed, 65 insertions(+), 18 deletions(-) diff --git a/app/src/components/Generation/GenerationForm.tsx b/app/src/components/Generation/GenerationForm.tsx index af29dcb9..e4f6cbd9 100644 --- a/app/src/components/Generation/GenerationForm.tsx +++ b/app/src/components/Generation/GenerationForm.tsx @@ -31,6 +31,7 @@ const generationSchema = z.object({ text: z.string().min(1, 'Text is required').max(5000), language: z.enum(['en', 'zh']), seed: z.number().int().optional(), + modelSize: z.enum(['1.7B', '0.6B']).optional(), }); type GenerationFormValues = z.infer; @@ -47,6 +48,7 @@ export function GenerationForm() { text: '', language: 'en', seed: undefined, + modelSize: '1.7B', }, }); @@ -57,6 +59,7 @@ export function GenerationForm() { text: data.text, language: data.language, seed: data.seed, + model_size: data.modelSize, }); toast({ @@ -126,7 +129,7 @@ export function GenerationForm() { )} /> -
+
+ ( + + Model Size + + Larger models produce better quality + + + )} + /> + (null); const { data: modelStatus, isLoading } = useQuery({ queryKey: ['modelStatus'], @@ -18,7 +20,10 @@ export function ModelManagement() { }); const downloadMutation = useMutation({ - mutationFn: (modelName: string) => apiClient.triggerModelDownload(modelName), + mutationFn: (modelName: string) => { + setDownloadingModel(modelName); + return apiClient.triggerModelDownload(modelName); + }, onSuccess: (_, modelName) => { toast({ title: 'Download started', @@ -30,12 +35,19 @@ export function ModelManagement() { }, 1000); }, onError: (error: Error) => { + setDownloadingModel(null); toast({ title: 'Download failed', description: error.message, variant: 'destructive', }); }, + onSettled: () => { + // Clear downloading state after a delay to allow progress to show + setTimeout(() => { + setDownloadingModel(null); + }, 2000); + }, }); const formatSize = (sizeMb?: number): string => { @@ -70,7 +82,7 @@ export function ModelManagement() { key={model.model_name} model={model} onDownload={() => downloadMutation.mutate(model.model_name)} - isDownloading={downloadMutation.isPending} + isDownloading={downloadingModel === model.model_name} formatSize={formatSize} /> ))} @@ -88,7 +100,7 @@ export function ModelManagement() { key={model.model_name} model={model} onDownload={() => downloadMutation.mutate(model.model_name)} - isDownloading={downloadMutation.isPending} + isDownloading={downloadingModel === model.model_name} formatSize={formatSize} /> ))} diff --git a/app/src/components/ServerSettings/ServerStatus.tsx b/app/src/components/ServerSettings/ServerStatus.tsx index e37fa5b3..275dfa07 100644 --- a/app/src/components/ServerSettings/ServerStatus.tsx +++ b/app/src/components/ServerSettings/ServerStatus.tsx @@ -47,18 +47,9 @@ export function ServerStatus() { Connected
- - Model: {health.model_loaded - ? `Loaded${health.model_size ? ` (${health.model_size})` : ''}` - : health.model_downloaded === false - ? 'Not Downloaded' - : 'Not Loaded'} + + {health.model_loaded || health.model_downloaded ? 'Model Ready' : 'No Model'} - {health.model_downloaded === true && !health.model_loaded && ( - - Model Cached (will load on first use) - - )} GPU: {health.gpu_available ? 'Available' : 'Not Available'} diff --git a/backend/main.py b/backend/main.py index 06bd366a..7a56ba30 100644 --- a/backend/main.py +++ b/backend/main.py @@ -64,9 +64,23 @@ async def health(): if gpu_available: vram_used = torch.cuda.memory_allocated() / 1024 / 1024 # MB - # Check if model is loaded - model_loaded = tts_model.is_loaded() - model_size = tts_model.model_size if model_loaded else None + # Check if model is loaded - use the same logic as model status endpoint + model_loaded = False + 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) model_downloaded = None @@ -240,6 +254,9 @@ async def generate_speech( # Generate audio 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( data.text, voice_prompt, diff --git a/backend/models.py b/backend/models.py index 47f7ffc7..14cea39f 100644 --- a/backend/models.py +++ b/backend/models.py @@ -49,6 +49,7 @@ class GenerationRequest(BaseModel): text: str = Field(..., min_length=1, max_length=5000) language: str = Field(default="en", pattern="^(en|zh)$") seed: Optional[int] = Field(None, ge=0) + model_size: Optional[str] = Field(default="1.7B", pattern="^(1\\.7B|0\\.6B)$") class GenerationResponse(BaseModel):