mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -07:00
add model selection and expanded language support to /transcribe endpoint
Closes #233
This commit is contained in:
@@ -32,6 +32,7 @@ import type {
|
|||||||
TranscriptionResponse,
|
TranscriptionResponse,
|
||||||
VoiceProfileCreate,
|
VoiceProfileCreate,
|
||||||
VoiceProfileResponse,
|
VoiceProfileResponse,
|
||||||
|
WhisperModelSize,
|
||||||
} from './types';
|
} from './types';
|
||||||
|
|
||||||
class ApiClient {
|
class ApiClient {
|
||||||
@@ -318,12 +319,19 @@ class ApiClient {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Transcription
|
// Transcription
|
||||||
async transcribeAudio(file: File, language?: LanguageCode): Promise<TranscriptionResponse> {
|
async transcribeAudio(
|
||||||
|
file: File,
|
||||||
|
language?: LanguageCode,
|
||||||
|
model?: WhisperModelSize,
|
||||||
|
): Promise<TranscriptionResponse> {
|
||||||
const formData = new FormData();
|
const formData = new FormData();
|
||||||
formData.append('file', file);
|
formData.append('file', file);
|
||||||
if (language) {
|
if (language) {
|
||||||
formData.append('language', language);
|
formData.append('language', language);
|
||||||
}
|
}
|
||||||
|
if (model) {
|
||||||
|
formData.append('model', model);
|
||||||
|
}
|
||||||
|
|
||||||
const url = `${this.getBaseUrl()}/transcribe`;
|
const url = `${this.getBaseUrl()}/transcribe`;
|
||||||
const response = await fetch(url, {
|
const response = await fetch(url, {
|
||||||
|
|||||||
@@ -99,8 +99,11 @@ export interface HistoryListResponse {
|
|||||||
total: number;
|
total: number;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export type WhisperModelSize = 'base' | 'small' | 'medium' | 'large' | 'turbo';
|
||||||
|
|
||||||
export interface TranscriptionRequest {
|
export interface TranscriptionRequest {
|
||||||
language?: LanguageCode;
|
language?: LanguageCode;
|
||||||
|
model?: WhisperModelSize;
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface TranscriptionResponse {
|
export interface TranscriptionResponse {
|
||||||
|
|||||||
@@ -1,10 +1,18 @@
|
|||||||
import { useMutation } from '@tanstack/react-query';
|
import { useMutation } from '@tanstack/react-query';
|
||||||
import { apiClient } from '@/lib/api/client';
|
import { apiClient } from '@/lib/api/client';
|
||||||
|
import type { WhisperModelSize } from '@/lib/api/types';
|
||||||
import type { LanguageCode } from '@/lib/constants/languages';
|
import type { LanguageCode } from '@/lib/constants/languages';
|
||||||
|
|
||||||
export function useTranscription() {
|
export function useTranscription() {
|
||||||
return useMutation({
|
return useMutation({
|
||||||
mutationFn: ({ file, language }: { file: File; language?: LanguageCode }) =>
|
mutationFn: ({
|
||||||
apiClient.transcribeAudio(file, language),
|
file,
|
||||||
|
language,
|
||||||
|
model,
|
||||||
|
}: {
|
||||||
|
file: File;
|
||||||
|
language?: LanguageCode;
|
||||||
|
model?: WhisperModelSize;
|
||||||
|
}) => apiClient.transcribeAudio(file, language, model),
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -134,6 +134,7 @@ class STTBackend(Protocol):
|
|||||||
self,
|
self,
|
||||||
audio_path: str,
|
audio_path: str,
|
||||||
language: Optional[str] = None,
|
language: Optional[str] = None,
|
||||||
|
model_size: Optional[str] = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
Transcribe audio to text.
|
Transcribe audio to text.
|
||||||
|
|||||||
@@ -345,18 +345,20 @@ class MLXSTTBackend:
|
|||||||
self,
|
self,
|
||||||
audio_path: str,
|
audio_path: str,
|
||||||
language: Optional[str] = None,
|
language: Optional[str] = None,
|
||||||
|
model_size: Optional[str] = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
Transcribe audio to text.
|
Transcribe audio to text.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
audio_path: Path to audio file
|
audio_path: Path to audio file
|
||||||
language: Optional language hint (en or zh)
|
language: Optional language hint
|
||||||
|
model_size: Optional model size override
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Transcribed text
|
Transcribed text
|
||||||
"""
|
"""
|
||||||
await self.load_model_async(None)
|
await self.load_model_async(model_size)
|
||||||
|
|
||||||
def _transcribe_sync():
|
def _transcribe_sync():
|
||||||
"""Run synchronous transcription in thread pool."""
|
"""Run synchronous transcription in thread pool."""
|
||||||
|
|||||||
@@ -306,18 +306,20 @@ class PyTorchSTTBackend:
|
|||||||
self,
|
self,
|
||||||
audio_path: str,
|
audio_path: str,
|
||||||
language: Optional[str] = None,
|
language: Optional[str] = None,
|
||||||
|
model_size: Optional[str] = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
Transcribe audio to text.
|
Transcribe audio to text.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
audio_path: Path to audio file
|
audio_path: Path to audio file
|
||||||
language: Optional language hint (en or zh)
|
language: Optional language hint
|
||||||
|
model_size: Optional model size override
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Transcribed text
|
Transcribed text
|
||||||
"""
|
"""
|
||||||
await self.load_model_async(None)
|
await self.load_model_async(model_size)
|
||||||
|
|
||||||
def _transcribe_sync():
|
def _transcribe_sync():
|
||||||
"""Run synchronous transcription in thread pool."""
|
"""Run synchronous transcription in thread pool."""
|
||||||
|
|||||||
+2
-1
@@ -149,7 +149,8 @@ class HistoryListResponse(BaseModel):
|
|||||||
class TranscriptionRequest(BaseModel):
|
class TranscriptionRequest(BaseModel):
|
||||||
"""Request model for audio transcription."""
|
"""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):
|
class TranscriptionResponse(BaseModel):
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ UPLOAD_CHUNK_SIZE = 1024 * 1024 # 1MB
|
|||||||
async def transcribe_audio(
|
async def transcribe_audio(
|
||||||
file: UploadFile = File(...),
|
file: UploadFile = File(...),
|
||||||
language: str | None = Form(None),
|
language: str | None = Form(None),
|
||||||
|
model: str | None = Form(None),
|
||||||
):
|
):
|
||||||
"""Transcribe audio file to text."""
|
"""Transcribe audio file to text."""
|
||||||
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
|
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
|
||||||
@@ -29,14 +30,23 @@ async def transcribe_audio(
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
from ..utils.audio import load_audio
|
from ..utils.audio import load_audio
|
||||||
|
from ..backends import WHISPER_HF_REPOS
|
||||||
|
|
||||||
audio, sr = await asyncio.to_thread(load_audio, tmp_path)
|
audio, sr = await asyncio.to_thread(load_audio, tmp_path)
|
||||||
duration = len(audio) / sr
|
duration = len(audio) / sr
|
||||||
|
|
||||||
whisper_model = transcribe.get_whisper_model()
|
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}"
|
progress_model_name = f"whisper-{model_size}"
|
||||||
task_manager = get_task_manager()
|
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(
|
return models.TranscriptionResponse(
|
||||||
text=text,
|
text=text,
|
||||||
|
|||||||
Reference in New Issue
Block a user