Whisper_AI / backend /transcriber.py
TANMAY-555's picture
Upload folder using huggingface_hub
c85835e verified
Raw
History Blame Contribute Delete
4.64 kB
import json
import os
from functools import lru_cache
from pathlib import Path
from typing import Any, Callable, Optional
from faster_whisper import WhisperModel
from . import config
def cleanup_outputs() -> None:
for child in config.OUTPUT_DIR.iterdir():
if child.is_file():
child.unlink()
def _detect_device() -> tuple[str, str]:
"""Return (device, compute_type). Prefers CUDA when available."""
try:
import ctranslate2
if ctranslate2.get_cuda_device_count() > 0:
return "cuda", "float16"
except Exception:
pass
return "cpu", "int8"
# Detect once at import time so every model load uses the same device
DEVICE, COMPUTE_TYPE = _detect_device()
CPU_THREADS = max(1, min(os.cpu_count() or 4, 8))
@lru_cache(maxsize=8)
def get_model(model_size: str) -> WhisperModel:
kwargs: dict[str, Any] = {"device": DEVICE, "compute_type": COMPUTE_TYPE}
if DEVICE == "cpu":
kwargs["cpu_threads"] = CPU_THREADS
kwargs["num_workers"] = 2
return WhisperModel(model_size, **kwargs)
def device_info() -> dict[str, str]:
return {"device": DEVICE, "compute_type": COMPUTE_TYPE}
def _srt_time(s: float) -> str:
h, rem = divmod(int(s), 3600)
m, sec = divmod(rem, 60)
ms = int((s % 1) * 1000)
return f"{h:02d}:{m:02d}:{sec:02d},{ms:03d}"
def _vtt_time(s: float) -> str:
h, rem = divmod(int(s), 3600)
m, sec = divmod(rem, 60)
ms = int((s % 1) * 1000)
return f"{h:02d}:{m:02d}:{sec:02d}.{ms:03d}"
def format_segments(segments: list[dict], fmt: str) -> str:
if fmt == "txt":
return "\n".join(s["text"] for s in segments)
if fmt == "srt":
lines: list[str] = []
for i, s in enumerate(segments, 1):
lines += [str(i), f"{_srt_time(s['start'])} --> {_srt_time(s['end'])}", s["text"], ""]
return "\n".join(lines)
if fmt == "vtt":
lines = ["WEBVTT", ""]
for s in segments:
lines += [f"{_vtt_time(s['start'])} --> {_vtt_time(s['end'])}", s["text"], ""]
return "\n".join(lines)
if fmt == "json":
return json.dumps({"segments": segments}, indent=2, ensure_ascii=False)
if fmt == "tsv":
lines = ["start\tend\ttext"]
lines += [f"{int(s['start'] * 1000)}\t{int(s['end'] * 1000)}\t{s['text']}" for s in segments]
return "\n".join(lines)
return "\n".join(s["text"] for s in segments)
def transcribe_file(
audio_path: Path,
model_size: str,
language: Optional[str] = None,
task: str = "transcribe",
word_timestamps: bool = False,
no_condition: bool = False,
fmt: str = "srt",
progress_callback: Optional[Callable[[int, str], None]] = None,
) -> dict[str, Any]:
if progress_callback:
progress_callback(5, f"Loading model ({DEVICE.upper()})…")
model = get_model(model_size)
if progress_callback:
progress_callback(12, "Analysing audio…")
segments_iter, info = model.transcribe(
str(audio_path),
language=language,
task=task,
word_timestamps=word_timestamps,
condition_on_previous_text=not no_condition,
vad_filter=True,
# Greedy decoding — 3-5× faster than beam_size=5 with minimal quality loss
beam_size=1,
best_of=1,
temperature=0.0,
)
result_segments: list[dict] = []
total = max(info.duration or 1.0, 1.0)
for segment in segments_iter:
seg: dict[str, Any] = {
"start": round(segment.start, 3),
"end": round(segment.end, 3),
"text": segment.text.strip(),
}
if word_timestamps and segment.words:
seg["words"] = [
{"start": round(w.start, 3), "end": round(w.end, 3), "word": w.word}
for w in segment.words
]
result_segments.append(seg)
if progress_callback:
pct = min(95, 15 + int((segment.end / total) * 80))
progress_callback(pct, f"Transcribing… {int(segment.end)}s / {int(total)}s")
formatted = format_segments(result_segments, fmt)
full_text = " ".join(s["text"] for s in result_segments).strip()
return {
"text": full_text,
"segments": result_segments,
"language": info.language,
"language_probability": info.language_probability,
"duration": round(info.duration, 2),
"format": fmt,
"model_size": model_size,
"task": task,
"formatted": formatted,
}