From f1541701fb79938a008180891958e26b6e0e7afe Mon Sep 17 00:00:00 2001 From: Jamie Pine Date: Mon, 16 Mar 2026 22:44:28 -0700 Subject: [PATCH] add model selection and expanded language support to /transcribe endpoint Closes #233 --- app/src/lib/api/client.ts | 10 +++++++++- app/src/lib/api/types.ts | 3 +++ app/src/lib/hooks/useTranscription.ts | 12 ++++++++++-- backend/backends/__init__.py | 1 + backend/backends/mlx_backend.py | 6 ++++-- backend/backends/pytorch_backend.py | 6 ++++-- backend/models.py | 3 ++- backend/routes/transcription.py | 16 +++++++++++++--- 8 files changed, 46 insertions(+), 11 deletions(-) diff --git a/app/src/lib/api/client.ts b/app/src/lib/api/client.ts index c6691ab5..036af3df 100644 --- a/app/src/lib/api/client.ts +++ b/app/src/lib/api/client.ts @@ -32,6 +32,7 @@ import type { TranscriptionResponse, VoiceProfileCreate, VoiceProfileResponse, + WhisperModelSize, } from './types'; class ApiClient { @@ -318,12 +319,19 @@ class ApiClient { } // Transcription - async transcribeAudio(file: File, language?: LanguageCode): Promise { + async transcribeAudio( + file: File, + language?: LanguageCode, + model?: WhisperModelSize, + ): Promise { const formData = new FormData(); formData.append('file', file); if (language) { formData.append('language', language); } + if (model) { + formData.append('model', model); + } const url = `${this.getBaseUrl()}/transcribe`; const response = await fetch(url, { diff --git a/app/src/lib/api/types.ts b/app/src/lib/api/types.ts index 49e90918..daae2a95 100644 --- a/app/src/lib/api/types.ts +++ b/app/src/lib/api/types.ts @@ -99,8 +99,11 @@ export interface HistoryListResponse { total: number; } +export type WhisperModelSize = 'base' | 'small' | 'medium' | 'large' | 'turbo'; + export interface TranscriptionRequest { language?: LanguageCode; + model?: WhisperModelSize; } export interface TranscriptionResponse { diff --git a/app/src/lib/hooks/useTranscription.ts b/app/src/lib/hooks/useTranscription.ts index 0b80722f..641b02df 100644 --- a/app/src/lib/hooks/useTranscription.ts +++ b/app/src/lib/hooks/useTranscription.ts @@ -1,10 +1,18 @@ import { useMutation } from '@tanstack/react-query'; import { apiClient } from '@/lib/api/client'; +import type { WhisperModelSize } from '@/lib/api/types'; import type { LanguageCode } from '@/lib/constants/languages'; export function useTranscription() { return useMutation({ - mutationFn: ({ file, language }: { file: File; language?: LanguageCode }) => - apiClient.transcribeAudio(file, language), + mutationFn: ({ + file, + language, + model, + }: { + file: File; + language?: LanguageCode; + model?: WhisperModelSize; + }) => apiClient.transcribeAudio(file, language, model), }); } diff --git a/backend/backends/__init__.py b/backend/backends/__init__.py index cc35eabe..6f20f3de 100644 --- a/backend/backends/__init__.py +++ b/backend/backends/__init__.py @@ -134,6 +134,7 @@ class STTBackend(Protocol): self, audio_path: str, language: Optional[str] = None, + model_size: Optional[str] = None, ) -> str: """ Transcribe audio to text. diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index e4a1ea97..92405dbf 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -345,18 +345,20 @@ class MLXSTTBackend: self, audio_path: str, language: Optional[str] = None, + model_size: Optional[str] = None, ) -> str: """ Transcribe audio to text. Args: audio_path: Path to audio file - language: Optional language hint (en or zh) + language: Optional language hint + model_size: Optional model size override Returns: Transcribed text """ - await self.load_model_async(None) + await self.load_model_async(model_size) def _transcribe_sync(): """Run synchronous transcription in thread pool.""" diff --git a/backend/backends/pytorch_backend.py b/backend/backends/pytorch_backend.py index 8f4a7a58..8ed4ab9c 100644 --- a/backend/backends/pytorch_backend.py +++ b/backend/backends/pytorch_backend.py @@ -306,18 +306,20 @@ class PyTorchSTTBackend: self, audio_path: str, language: Optional[str] = None, + model_size: Optional[str] = None, ) -> str: """ Transcribe audio to text. Args: audio_path: Path to audio file - language: Optional language hint (en or zh) + language: Optional language hint + model_size: Optional model size override Returns: Transcribed text """ - await self.load_model_async(None) + await self.load_model_async(model_size) def _transcribe_sync(): """Run synchronous transcription in thread pool.""" diff --git a/backend/models.py b/backend/models.py index ef8e196d..3308b3bc 100644 --- a/backend/models.py +++ b/backend/models.py @@ -149,7 +149,8 @@ class HistoryListResponse(BaseModel): class TranscriptionRequest(BaseModel): """Request model for audio transcription.""" - language: Optional[str] = Field(None, pattern="^(en|zh)$") + language: Optional[str] = Field(None, pattern="^(en|zh|ja|ko|de|fr|ru|pt|es|it)$") + model: Optional[str] = Field(None, pattern="^(base|small|medium|large|turbo)$") class TranscriptionResponse(BaseModel): diff --git a/backend/routes/transcription.py b/backend/routes/transcription.py index 90cb1c95..dc949132 100644 --- a/backend/routes/transcription.py +++ b/backend/routes/transcription.py @@ -20,6 +20,7 @@ UPLOAD_CHUNK_SIZE = 1024 * 1024 # 1MB async def transcribe_audio( file: UploadFile = File(...), language: str | None = Form(None), + model: str | None = Form(None), ): """Transcribe audio file to text.""" with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp: @@ -29,14 +30,23 @@ async def transcribe_audio( try: from ..utils.audio import load_audio + from ..backends import WHISPER_HF_REPOS audio, sr = await asyncio.to_thread(load_audio, tmp_path) duration = len(audio) / sr whisper_model = transcribe.get_whisper_model() - model_size = whisper_model.model_size + model_size = model if model else whisper_model.model_size - if not whisper_model.is_loaded() and not whisper_model._is_model_cached(model_size): + valid_sizes = list(WHISPER_HF_REPOS.keys()) + if model_size not in valid_sizes: + raise HTTPException( + status_code=400, + detail=f"Invalid model size '{model_size}'. Must be one of: {', '.join(valid_sizes)}", + ) + + already_loaded = whisper_model.is_loaded() and whisper_model.model_size == model_size + if not already_loaded and not whisper_model._is_model_cached(model_size): progress_model_name = f"whisper-{model_size}" task_manager = get_task_manager() @@ -59,7 +69,7 @@ async def transcribe_audio( }, ) - text = await whisper_model.transcribe(tmp_path, language) + text = await whisper_model.transcribe(tmp_path, language, model_size) return models.TranscriptionResponse( text=text,