Files
doctate/whisper/main.py
T
Brummel 6e6cfd77b7 Add custom faster-whisper service with anti-hallucination params
Off-the-shelf Whisper wrappers hardcode faster-whisper defaults that
cause cascading hallucinations on real dictations (2/13 runs against
ahmetoner v1.9.1). This thin FastAPI service mirrors the /asr interface
but fixes condition_on_previous_text=False, temperature=0.0, and
vad_filter=True so the Axum server can switch by changing WHISPER_URL.
2026-04-13 19:45:52 +02:00

130 lines
3.9 KiB
Python

"""Thin FastAPI wrapper around faster-whisper.
Exposes the same /asr interface as ahmetoner/whisper-asr-webservice so the
Axum server stays unchanged, but hardcodes anti-hallucination parameters
that no off-the-shelf wrapper exposes.
"""
import json
import logging
import os
import tempfile
import time
from typing import Optional
from fastapi import FastAPI, File, Form, Query, UploadFile
from fastapi.responses import JSONResponse, PlainTextResponse
from faster_whisper import WhisperModel
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s %(message)s",
)
log = logging.getLogger("whisper")
MODEL_NAME = os.getenv("WHISPER_MODEL", "large-v3-turbo")
COMPUTE_TYPE = os.getenv("WHISPER_COMPUTE_TYPE", "float16")
DEVICE = os.getenv("WHISPER_DEVICE", "cuda")
MODELS_DIR = os.getenv("WHISPER_MODELS_DIR", "/models")
OFFLINE = os.getenv("WHISPER_OFFLINE", "0") == "1"
log.info("Loading model=%s device=%s compute_type=%s", MODEL_NAME, DEVICE, COMPUTE_TYPE)
model = WhisperModel(
MODEL_NAME,
device=DEVICE,
compute_type=COMPUTE_TYPE,
download_root=MODELS_DIR,
local_files_only=OFFLINE,
)
log.info("Model loaded")
app = FastAPI(title="doctate-whisper", version="0.1.0")
@app.get("/health")
def health():
return {"status": "ok", "model": MODEL_NAME, "device": DEVICE}
@app.get("/info")
def info():
return {
"model": MODEL_NAME,
"device": DEVICE,
"compute_type": COMPUTE_TYPE,
"offline": OFFLINE,
}
@app.post("/asr")
async def asr(
audio_file: UploadFile = File(...),
output: str = Query("txt"),
language: Optional[str] = Query(None),
initial_prompt: Optional[str] = Form(None),
):
suffix = os.path.splitext(audio_file.filename or "")[1] or ".bin"
tmp = tempfile.NamedTemporaryFile(suffix=suffix, delete=False)
try:
tmp.write(await audio_file.read())
tmp.flush()
tmp.close()
t0 = time.monotonic()
segments_gen, info_obj = model.transcribe(
tmp.name,
language=language,
task="transcribe",
beam_size=5,
best_of=5,
# Fixed anti-hallucination params — the whole point of this service.
temperature=0.0,
condition_on_previous_text=False,
vad_filter=True,
vad_parameters={"min_silence_duration_ms": 500},
no_speech_threshold=0.6,
log_prob_threshold=-1.0,
compression_ratio_threshold=2.4,
initial_prompt=initial_prompt,
)
segments = list(segments_gen)
infer_secs = time.monotonic() - t0
text = "".join(s.text for s in segments).strip()
worst_logprob = min((s.avg_logprob for s in segments), default=0.0)
log.info(
json.dumps({
"event": "transcribe",
"duration_audio": info_obj.duration,
"duration_infer": round(infer_secs, 2),
"segments": len(segments),
"worst_logprob": round(worst_logprob, 3),
"language": info_obj.language,
})
)
if output == "json":
return JSONResponse({
"text": text,
"language": info_obj.language,
"duration": info_obj.duration,
"segments": [
{
"start": s.start,
"end": s.end,
"text": s.text,
"avg_logprob": s.avg_logprob,
"compression_ratio": s.compression_ratio,
"no_speech_prob": s.no_speech_prob,
"temperature": s.temperature,
}
for s in segments
],
})
return PlainTextResponse(text)
finally:
try:
os.unlink(tmp.name)
except OSError:
pass