143 lines
4.8 KiB
Python
143 lines
4.8 KiB
Python
"""Transcription model registry (Stage 1).
|
||
|
||
The supported faster-whisper sizes, their display metadata, validation, and
|
||
local availability checks. Nothing here downloads anything: faster-whisper
|
||
fetches from HuggingFace on first use, so "downloaded" is answered by
|
||
inspecting the HF hub cache, and "available" additionally requires the
|
||
faster-whisper package itself.
|
||
|
||
No silent substitution anywhere: unknown names are rejected, missing
|
||
downloads fail fast with instructions.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import os
|
||
from dataclasses import dataclass
|
||
from pathlib import Path
|
||
|
||
from sqlalchemy import select
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
||
from shonar.db.models import AppSetting
|
||
from shonar.services.ai import ProviderConfigError
|
||
|
||
DEFAULT_MODEL = "base"
|
||
DEFAULT_MODEL_KEY = "transcription.default_model"
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class TranscriptionModelInfo:
|
||
name: str
|
||
display_name: str
|
||
description: str
|
||
params: str
|
||
approx_memory: str
|
||
relative_speed: str
|
||
|
||
|
||
SUPPORTED_TRANSCRIPTION_MODELS: dict[str, TranscriptionModelInfo] = {
|
||
"base": TranscriptionModelInfo(
|
||
name="base",
|
||
display_name="Whisper Base",
|
||
description="The bundled model — works offline out of the box. "
|
||
"Good for clear speech on ordinary computers.",
|
||
params="~74M",
|
||
approx_memory="~1 GB RAM",
|
||
relative_speed="~7x real-time (CPU)",
|
||
),
|
||
"large-v3": TranscriptionModelInfo(
|
||
name="large-v3",
|
||
display_name="Whisper Large v3",
|
||
description="Best accuracy (accents, names, messy audio). Downloads "
|
||
"~3 GB once; slow on CPU — best on machines with 16 GB+ RAM.",
|
||
params="~1.5B",
|
||
approx_memory="~10 GB RAM",
|
||
relative_speed="~0.3–0.4x real-time (CPU)",
|
||
),
|
||
}
|
||
|
||
|
||
def normalize_model_name(raw: str | None) -> str:
|
||
"""Case/whitespace-tolerant normalization. Never maps one model to another."""
|
||
return (raw or "").strip().lower()
|
||
|
||
|
||
def validate_model_name(raw: str | None) -> str:
|
||
"""Return the normalized name, or raise with a helpful message."""
|
||
name = normalize_model_name(raw)
|
||
if name in SUPPORTED_TRANSCRIPTION_MODELS:
|
||
return name
|
||
supported = ", ".join(sorted(SUPPORTED_TRANSCRIPTION_MODELS))
|
||
raise ProviderConfigError(
|
||
f"Unsupported transcription model {raw!r}. Supported models: {supported}. "
|
||
"Check the spelling — a different model is never substituted silently."
|
||
)
|
||
|
||
|
||
def _hf_hub_cache() -> Path:
|
||
try:
|
||
from huggingface_hub.constants import HF_HUB_CACHE
|
||
|
||
return Path(HF_HUB_CACHE)
|
||
except ImportError:
|
||
return Path(
|
||
os.environ.get("HF_HUB_CACHE", str(Path.home() / ".cache" / "huggingface" / "hub"))
|
||
)
|
||
|
||
|
||
def is_model_downloaded(name: str) -> bool:
|
||
"""True when a non-empty faster-whisper snapshot for `name` sits in the
|
||
HuggingFace hub cache (repo Systran/faster-whisper-<name>)."""
|
||
repo_dir = _hf_hub_cache() / f"models--Systran--faster-whisper-{name}"
|
||
snapshots = repo_dir / "snapshots"
|
||
if not snapshots.is_dir():
|
||
return False
|
||
return any(s.is_dir() and any(s.iterdir()) for s in snapshots.iterdir())
|
||
|
||
|
||
def is_faster_whisper_installed() -> bool:
|
||
try:
|
||
import faster_whisper # noqa: F401
|
||
|
||
return True
|
||
except ImportError:
|
||
return False
|
||
|
||
|
||
def download_instructions(name: str) -> str:
|
||
return (
|
||
f'Model "{name}" is not downloaded. Download it with: '
|
||
f"POST /api/v1/models/{name}/download "
|
||
"(needs internet once), or run any transcription with that model selected — "
|
||
"faster-whisper fetches it from HuggingFace automatically."
|
||
)
|
||
|
||
|
||
async def get_global_default_model(session: AsyncSession) -> str:
|
||
row = await session.scalar(select(AppSetting).where(AppSetting.key == DEFAULT_MODEL_KEY))
|
||
if row is None:
|
||
return DEFAULT_MODEL
|
||
model = row.value.get("model") if isinstance(row.value, dict) else None
|
||
return model if model in SUPPORTED_TRANSCRIPTION_MODELS else DEFAULT_MODEL
|
||
|
||
|
||
async def set_global_default_model(session: AsyncSession, raw: str) -> str:
|
||
"""Validate + persist the global default. Affects future recordings only —
|
||
existing rows keep their saved model."""
|
||
name = validate_model_name(raw)
|
||
row = await session.scalar(select(AppSetting).where(AppSetting.key == DEFAULT_MODEL_KEY))
|
||
if row is None:
|
||
row = AppSetting(key=DEFAULT_MODEL_KEY, value={"model": name})
|
||
session.add(row)
|
||
else:
|
||
row.value = {"model": name}
|
||
await session.flush()
|
||
return name
|
||
|
||
|
||
def effective_model(recording_model: str | None, global_default: str) -> str:
|
||
"""Per-recording override wins; otherwise the global default."""
|
||
if recording_model and recording_model in SUPPORTED_TRANSCRIPTION_MODELS:
|
||
return recording_model
|
||
return global_default
|